video-chat / utils /vision_analyzer.py
corl0s's picture
Intial deploy
e417aa9
Raw
History Blame Contribute Delete
3.18 kB
"""Frame analysis using NVIDIA NIM vision API (llama-3.2-11b-vision-instruct)."""
import logging
import os
import time
from openai import OpenAI, RateLimitError
logger = logging.getLogger(__name__)
_NIM_BASE_URL = "https://integrate.api.nvidia.com/v1"
_VISION_MODEL = "meta/llama-3.2-11b-vision-instruct"
_FRAME_PROMPT = (
"Describe what is happening in this video frame in 2-3 sentences. "
"Focus on: people, objects, actions, text visible, scene/setting, and any notable events."
)
def _get_client() -> OpenAI:
api_key = os.environ.get("NVIDIA_API_KEY")
if not api_key:
raise EnvironmentError("NVIDIA_API_KEY environment variable is not set.")
return OpenAI(base_url=_NIM_BASE_URL, api_key=api_key)
def analyze_frame(base64_image: str, frame_index: int, timestamp: float) -> dict:
"""Send a single base64 JPEG frame to the NIM vision API and return its description.
Retries up to 3 times on rate-limit errors with a 2s delay between attempts.
Returns a dict with keys: frame_index, timestamp, description, model_used.
"""
client = _get_client()
data_uri = f"data:image/jpeg;base64,{base64_image}"
for attempt in range(1, 4):
try:
response = client.chat.completions.create(
model=_VISION_MODEL,
messages=[
{
"role": "user",
"content": [
{"type": "image_url", "image_url": {"url": data_uri}},
{"type": "text", "text": _FRAME_PROMPT},
],
}
],
max_tokens=256,
)
description = response.choices[0].message.content.strip()
logger.info("Frame %d (%.1fs) analyzed", frame_index, timestamp)
return {
"frame_index": frame_index,
"timestamp": timestamp,
"description": description,
"model_used": _VISION_MODEL,
}
except RateLimitError:
if attempt < 3:
logger.warning("Rate limit on frame %d, retrying in 2s (attempt %d/3)", frame_index, attempt)
time.sleep(2)
else:
raise
except Exception as exc:
logger.error("Vision API error on frame %d: %s", frame_index, exc)
raise
def analyze_frames_batch(frames_list: list[dict]) -> list[dict]:
"""Analyze a list of frame dicts, returning analysis results for each.
Skips frames that fail and continues processing the rest.
"""
results = []
for frame in frames_list:
try:
result = analyze_frame(
frame["base64_image"],
frame["frame_index"],
frame["timestamp_seconds"],
)
results.append(result)
except Exception as exc:
logger.warning(
"Skipping frame %d (%.1fs) due to error: %s",
frame["frame_index"],
frame["timestamp_seconds"],
exc,
)
time.sleep(0.5)
return results