5d0c771a2b
- trace_vlm_caption.py: use MOMENTRY_LLM_VISION_URL and MOMENTRY_LLM_VISION_MODEL - scene_vlm_caption.py: use MOMENTRY_LLM_VISION_URL and MOMENTRY_LLM_VISION_MODEL - Changed from Ollama /api/generate to OpenAI-compatible /v1/chat/completions format - Added embedding server environment variables
358 lines
14 KiB
Python
Executable File
358 lines
14 KiB
Python
Executable File
#!/opt/homebrew/bin/python3.11
|
|
"""
|
|
Trace VLM Caption - Generate VLM descriptions for face traces
|
|
|
|
Analyzes key_face.jpg or key_frame.jpg using VLM and updates trace_profile.json.
|
|
|
|
Usage:
|
|
python trace_vlm_caption.py --trace-dir /path/to/output/{uuid}/trace_0
|
|
python trace_vlm_caption.py --file-uuid abc123 --trace-id 0 --output-dir /path/to/output
|
|
|
|
Output (13 fields):
|
|
Person: vlm_description, vlm_clothing, vlm_tags, vlm_hand_objects
|
|
Environment: vlm_lighting, vlm_location, vlm_weather, vlm_setting, vlm_transportation
|
|
Nature: vlm_has_plants, vlm_plants, vlm_has_animals, vlm_animals
|
|
Context: vlm_background, vlm_bg_tags
|
|
|
|
Environment Variables:
|
|
MOMENTRY_LLM_VISION_URL - VLM endpoint (default: http://localhost:8091/v1/chat/completions)
|
|
MOMENTRY_LLM_VISION_MODEL - VLM model (default: llava-v1.6-vicuna-13b)
|
|
"""
|
|
|
|
import argparse
|
|
import base64
|
|
import json
|
|
import os
|
|
import sys
|
|
from pathlib import Path
|
|
|
|
# VLM configuration from environment variables
|
|
VLM_URL = os.environ.get("MOMENTRY_LLM_VISION_URL", "http://localhost:8091/v1/chat/completions")
|
|
VLM_MODEL = os.environ.get("MOMENTRY_LLM_VISION_MODEL", "llava-v1.6-vicuna-13b")
|
|
|
|
try:
|
|
import requests
|
|
except ImportError:
|
|
print("requests not installed: pip install requests", file=sys.stderr)
|
|
sys.exit(1)
|
|
|
|
|
|
def encode_image(image_path: str) -> str:
|
|
"""Encode image to base64."""
|
|
with open(image_path, "rb") as f:
|
|
return base64.b64encode(f.read()).decode("utf-8")
|
|
|
|
|
|
def call_vlm(image_path: str, prompt: str) -> str:
|
|
"""Call VLM API using OpenAI-compatible format."""
|
|
image_b64 = encode_image(image_path)
|
|
|
|
payload = {
|
|
"model": VLM_MODEL,
|
|
"messages": [
|
|
{"role": "user", "content": [
|
|
{"type": "text", "text": prompt},
|
|
{"type": "image_url", "image_url": {"url": f"data:image/jpeg;base64,{image_b64}"}}
|
|
]}
|
|
],
|
|
"max_tokens": 100,
|
|
}
|
|
|
|
try:
|
|
resp = requests.post(VLM_URL, json=payload, timeout=30)
|
|
resp.raise_for_status()
|
|
data = resp.json()
|
|
return data.get("choices", [{}])[0].get("message", {}).get("content", "").strip()
|
|
except Exception as e:
|
|
print(f"VLM API error: {e}", file=sys.stderr)
|
|
return ""
|
|
|
|
|
|
def get_embedding(text: str) -> list:
|
|
"""Get embedding from embedding server."""
|
|
embed_url = os.environ.get("MOMENTRY_EMBEDDING_URL", "http://localhost:11436/v1/embeddings")
|
|
embed_model = os.environ.get("MOMENTRY_EMBEDDING_MODEL", "embeddinggemma-300m")
|
|
try:
|
|
resp = requests.post(
|
|
embed_url,
|
|
json={"model": embed_model, "input": text},
|
|
timeout=30,
|
|
)
|
|
resp.raise_for_status()
|
|
data = resp.json()
|
|
return data.get("data", [{}])[0].get("embedding", [])
|
|
except Exception as e:
|
|
print(f"[vlm] Embedding error: {e}", file=sys.stderr)
|
|
return []
|
|
|
|
|
|
def store_to_qdrant(profile: dict, file_uuid: str, trace_id: int, qdrant_url: str = "http://localhost:6333", qdrant_api_key: str = None) -> bool:
|
|
"""Store VLM results to Qdrant _vlm collection."""
|
|
description = profile.get("vlm_description", "")
|
|
if not description:
|
|
return False
|
|
|
|
# Get embedding
|
|
embedding = get_embedding(description)
|
|
if not embedding:
|
|
print(f"[vlm] Failed to get embedding for trace_{trace_id}", file=sys.stderr)
|
|
return False
|
|
|
|
# Generate point ID from file_uuid + trace_id
|
|
import hashlib
|
|
point_id = int(hashlib.md5(f"{file_uuid}_trace_{trace_id}".encode()).hexdigest()[:16], 16)
|
|
|
|
# Build payload
|
|
payload = {
|
|
"type": "trace",
|
|
"file_uuid": file_uuid,
|
|
"trace_id": trace_id,
|
|
"vlm_description": profile.get("vlm_description", ""),
|
|
"vlm_clothing": profile.get("vlm_clothing", ""),
|
|
"vlm_tags": profile.get("vlm_tags", []),
|
|
"vlm_hand_objects": profile.get("vlm_hand_objects", ""),
|
|
"vlm_lighting": profile.get("vlm_lighting", "unknown"),
|
|
"vlm_location": profile.get("vlm_location", "unknown"),
|
|
"vlm_weather": profile.get("vlm_weather", "unknown"),
|
|
"vlm_setting": profile.get("vlm_setting", "unknown"),
|
|
"vlm_transportation": profile.get("vlm_transportation", "unknown"),
|
|
"vlm_has_plants": profile.get("vlm_has_plants", False),
|
|
"vlm_plants": profile.get("vlm_plants", []),
|
|
"vlm_has_animals": profile.get("vlm_has_animals", False),
|
|
"vlm_animals": profile.get("vlm_animals", []),
|
|
"vlm_background": profile.get("vlm_background", ""),
|
|
"vlm_bg_tags": profile.get("vlm_bg_tags", []),
|
|
"vlm_model": profile.get("vlm_model", ""),
|
|
}
|
|
|
|
# Upsert to Qdrant
|
|
try:
|
|
headers = {}
|
|
if qdrant_api_key:
|
|
headers["api-key"] = qdrant_api_key
|
|
|
|
resp = requests.put(
|
|
f"{qdrant_url}/collections/_vlm/points?wait=true",
|
|
json={
|
|
"points": [{
|
|
"id": point_id,
|
|
"vector": embedding,
|
|
"payload": payload,
|
|
}]
|
|
},
|
|
headers=headers,
|
|
timeout=30,
|
|
)
|
|
resp.raise_for_status()
|
|
print(f"[vlm] Stored to Qdrant: trace_{trace_id}")
|
|
return True
|
|
except Exception as e:
|
|
print(f"[vlm] Qdrant error: {e}", file=sys.stderr)
|
|
return False
|
|
|
|
|
|
def analyze_trace(trace_dir: str, model: str = "llava:7b", store_qdrant: bool = True) -> dict:
|
|
"""
|
|
Analyze a face trace with VLM.
|
|
|
|
Returns:
|
|
Dict with VLM analysis results
|
|
"""
|
|
trace_path = Path(trace_dir)
|
|
profile_path = trace_path / "trace_profile.json"
|
|
|
|
if not profile_path.exists():
|
|
print(f"[vlm] No trace_profile.json in {trace_dir}", file=sys.stderr)
|
|
return {}
|
|
|
|
# Load existing profile
|
|
with open(profile_path, "r") as f:
|
|
profile = json.load(f)
|
|
|
|
# Find image to analyze (prefer key_frame.jpg for full context)
|
|
key_frame = trace_path / "key_frame.jpg"
|
|
key_face = trace_path / "key_face.jpg"
|
|
|
|
if not key_frame.exists() and not key_face.exists():
|
|
print(f"[vlm] No key_frame.jpg or key_face.jpg in {trace_dir}", file=sys.stderr)
|
|
return profile
|
|
|
|
# Use key_frame for clothing/background analysis (full body context)
|
|
image_to_analyze = str(key_frame) if key_frame.exists() else str(key_face)
|
|
|
|
print(f"[vlm] Analyzing {trace_path.name}...")
|
|
|
|
# Prompt 1: Person description
|
|
desc_prompt = "Describe this person briefly. Include: gender, age range, hair, visible clothing. If uncertain, say 'unknown'. Be concise and honest."
|
|
description = call_vlm(image_to_analyze, desc_prompt, model)
|
|
|
|
# Prompt 2: Clothing details
|
|
clothing_prompt = "Describe this person's clothing in detail. Include colors, type of clothing, any visible text or logos. If unclear, say 'unclear' or 'partially visible'. Do not guess."
|
|
clothing = call_vlm(image_to_analyze, clothing_prompt, model)
|
|
|
|
# Prompt 3: Tags
|
|
tags_prompt = "List 5 tags describing this person's appearance, comma-separated. Only include what you can clearly see. Examples: man, glasses, red-shirt, formal, casual."
|
|
tags_raw = call_vlm(image_to_analyze, tags_prompt, model)
|
|
tags = [t.strip() for t in tags_raw.replace(",", " ").split() if t.strip()][:5]
|
|
|
|
# Prompt 5: Objects in hand
|
|
hand_prompt = "What is this person holding in their hands? Answer: object names if clearly visible, or 'nothing visible', or 'unclear'. Do not guess."
|
|
hand_objects = call_vlm(image_to_analyze, hand_prompt, model)
|
|
|
|
# Prompt 6: Lighting (day/night)
|
|
light_prompt = "What is the lighting condition? Answer one word: day, night, indoor-light, mixed, or unknown. If uncertain, answer 'unknown'."
|
|
lighting = call_vlm(image_to_analyze, light_prompt, model).lower().strip()
|
|
|
|
# Prompt 7: Scene classification
|
|
scene_prompt = "Classify the scene. Answer in JSON: {\"location\": \"indoor/outdoor/unknown\", \"weather\": \"sunny/cloudy/rainy/night/unknown\", \"setting\": \"office/street/home/nature/studio/unknown\", \"transportation\": \"car/train/bus/none/unknown\"}. Use 'unknown' if uncertain."
|
|
scene_raw = call_vlm(image_to_analyze, scene_prompt, model)
|
|
|
|
# Parse scene JSON (handle markdown code blocks)
|
|
scene_data = {}
|
|
try:
|
|
# Remove markdown code blocks if present
|
|
scene_clean = scene_raw.replace("```json", "").replace("```", "").strip()
|
|
scene_data = json.loads(scene_clean)
|
|
except:
|
|
scene_data = {}
|
|
|
|
# Prompt 8: Background description
|
|
bg_prompt = "Describe the background and environment briefly. Include only what is clearly visible. If uncertain about details, say 'unclear' or 'partially visible'. Do not guess or imagine."
|
|
background = call_vlm(image_to_analyze, bg_prompt, model)
|
|
|
|
# Prompt 9: Plants detection
|
|
plants_prompt = "What plants, trees, or flowers are clearly visible? Answer in JSON: {\"has_plants\": true/false, \"plants\": [\"list recognizable plants by name. If not recognizable, describe briefly what you see. Use empty list if none or uncertain.\"]}"
|
|
plants_raw = call_vlm(image_to_analyze, plants_prompt, model)
|
|
|
|
# Parse plants JSON
|
|
plants_data = {}
|
|
try:
|
|
plants_clean = plants_raw.replace("```json", "").replace("```", "").strip()
|
|
plants_data = json.loads(plants_clean)
|
|
except:
|
|
plants_data = {"has_plants": False, "plants": []}
|
|
|
|
# Prompt 10: Animals detection
|
|
animals_prompt = "What animals are clearly visible? Answer in JSON: {\"has_animals\": true/false, \"animals\": [\"list recognizable animals by name. If not recognizable, describe briefly what you see. Use empty list if none or uncertain.\"]}"
|
|
animals_raw = call_vlm(image_to_analyze, animals_prompt, model)
|
|
|
|
# Parse animals JSON
|
|
animals_data = {}
|
|
try:
|
|
animals_clean = animals_raw.replace("```json", "").replace("```", "").strip()
|
|
animals_data = json.loads(animals_clean)
|
|
except:
|
|
animals_data = {"has_animals": False, "animals": []}
|
|
|
|
# Prompt 11: Background tags
|
|
bg_tags_prompt = "List 5 tags for the background/scene, comma-separated. Examples: office, street, sunny, building, car, trees."
|
|
bg_tags_raw = call_vlm(image_to_analyze, bg_tags_prompt, model)
|
|
bg_tags = [t.strip() for t in bg_tags_raw.replace(",", " ").split() if t.strip()][:5]
|
|
|
|
# Update profile
|
|
profile["vlm_description"] = description
|
|
profile["vlm_clothing"] = clothing
|
|
profile["vlm_tags"] = tags
|
|
profile["vlm_hand_objects"] = hand_objects
|
|
profile["vlm_lighting"] = lighting
|
|
profile["vlm_location"] = scene_data.get("location", "unknown")
|
|
profile["vlm_weather"] = scene_data.get("weather", "unknown")
|
|
profile["vlm_setting"] = scene_data.get("setting", "unknown")
|
|
profile["vlm_transportation"] = scene_data.get("transportation", "unknown")
|
|
profile["vlm_has_plants"] = plants_data.get("has_plants", False)
|
|
profile["vlm_plants"] = plants_data.get("plants", [])
|
|
profile["vlm_has_animals"] = animals_data.get("has_animals", False)
|
|
profile["vlm_animals"] = animals_data.get("animals", [])
|
|
profile["vlm_background"] = background
|
|
profile["vlm_bg_tags"] = bg_tags
|
|
profile["vlm_model"] = model
|
|
|
|
# Save updated profile
|
|
with open(profile_path, "w") as f:
|
|
json.dump(profile, f, indent=2)
|
|
|
|
# Store to Qdrant
|
|
if store_qdrant:
|
|
file_uuid = profile.get("file_uuid", "")
|
|
trace_id = profile.get("trace_id", 0)
|
|
if file_uuid:
|
|
qdrant_api_key = os.environ.get("QDRANT_API_KEY")
|
|
store_to_qdrant(profile, file_uuid, trace_id, qdrant_api_key=qdrant_api_key)
|
|
|
|
print(f"[vlm] Updated {trace_path.name}: {description[:30]}... | Loc: {scene_data.get('location', '?')} | Light: {lighting} | Hand: {hand_objects[:20]}...")
|
|
|
|
return profile
|
|
|
|
|
|
def analyze_all_traces(file_uuid: str, output_dir: str, model: str = "llava:7b") -> dict:
|
|
"""
|
|
Analyze all traces for a file.
|
|
|
|
Returns:
|
|
Summary dict
|
|
"""
|
|
file_dir = Path(output_dir) / file_uuid
|
|
|
|
if not file_dir.exists():
|
|
print(f"No trace directory: {file_dir}", file=sys.stderr)
|
|
return {"error": "No trace directory"}
|
|
|
|
trace_dirs = sorted(file_dir.glob("trace_*"))
|
|
|
|
if not trace_dirs:
|
|
print(f"No traces found in {file_dir}", file=sys.stderr)
|
|
return {"error": "No traces"}
|
|
|
|
results = []
|
|
for trace_dir in trace_dirs:
|
|
profile = analyze_trace(str(trace_dir), model)
|
|
if profile:
|
|
results.append({
|
|
"trace_id": profile.get("trace_id"),
|
|
"vlm_description": profile.get("vlm_description", "")[:50] + "...",
|
|
"vlm_tags": profile.get("vlm_tags", []),
|
|
})
|
|
|
|
return {
|
|
"file_uuid": file_uuid,
|
|
"total_traces": len(trace_dirs),
|
|
"analyzed": len(results),
|
|
"traces": results,
|
|
}
|
|
|
|
|
|
def main():
|
|
parser = argparse.ArgumentParser(description="VLM caption generation for face traces")
|
|
parser.add_argument("--trace-dir", "-t", help="Single trace directory")
|
|
parser.add_argument("--file-uuid", "-u", help="File UUID (analyze all traces)")
|
|
parser.add_argument("--trace-id", type=int, help="Single trace ID (requires --file-uuid)")
|
|
parser.add_argument("--output-dir", "-o", default="/Users/accusys/momentry/output", help="Output directory")
|
|
parser.add_argument("--model", "-m", default="llava:7b", help="VLM model name")
|
|
parser.add_argument("--json", "-j", action="store_true", help="Output as JSON")
|
|
args = parser.parse_args()
|
|
|
|
if args.trace_dir:
|
|
# Single trace directory
|
|
result = analyze_trace(args.trace_dir, args.model)
|
|
elif args.file_uuid:
|
|
if args.trace_id is not None:
|
|
# Single trace
|
|
trace_dir = Path(args.output_dir) / args.file_uuid / f"trace_{args.trace_id}"
|
|
result = analyze_trace(str(trace_dir), args.model)
|
|
else:
|
|
# All traces for file
|
|
result = analyze_all_traces(args.file_uuid, args.output_dir, args.model)
|
|
else:
|
|
parser.error("Requires --trace-dir or --file-uuid")
|
|
|
|
if args.json:
|
|
print(json.dumps(result, indent=2))
|
|
else:
|
|
if "vlm_description" in result:
|
|
print(f"Description: {result['vlm_description']}")
|
|
print(f"Clothing: {result.get('vlm_clothing', '')}")
|
|
print(f"Tags: {result.get('vlm_tags', [])}")
|
|
|
|
|
|
if __name__ == "__main__":
|
|
main() |