from __future__ import annotations import contextlib import io import json import logging import subprocess import tempfile import warnings from pathlib import Path from typing import Any import torch from transformers import AutoModelForMultimodalLM, AutoProcessor SCRIPT_DIR = Path(__file__).resolve().parent MODEL_ROOT = SCRIPT_DIR from caption_model_runtime import ( DEFAULT_PARAKEET_MODEL_ID, CAPTION_LENGTH_LABELS, CAPTION_SETTING_FIELD_CHOICES, DEFAULT_CAPTION_SETTING_VALUES, format_caption_settings_prompt, gemma_core, install_parakeet_audio_bridge, load_state_file, replace_batch_audio_features, ) # Edit these. VIDEO_PATH = "/workspace/test7.mp4" # Options: path to the video you want to caption. MODEL_PATH = str(MODEL_ROOT / "model") # Options: merged model path. PROCESSOR_PATH = str(MODEL_ROOT / "processor") # Options: processor path. AUDIO_PROJECTOR_PATH = str(MODEL_ROOT / "model" / "embed_audio.safetensors") # Options: trained audio projector path. # Parakeet hybrid audio bridge settings. These should normally match the packaged model. PARAKEET_MODEL_ID = str(MODEL_ROOT / "parakeet") # Options: "nvidia/parakeet-tdt-0.6b-v3" or compatible local/HF path. PARAKEET_BRIDGE_MODE = "tdt_token_embeddings_with_encoder_context" # Options: "encoder", "tdt_tokens", "tdt_token_embeddings", "encoder_soft_tdt_token_embeddings", "tdt_token_embeddings_with_encoder_context". PARAKEET_NATIVE_FEATURES = True # Options: True to replace Gemma audio features with Parakeet features, False for debugging only. PARAKEET_TDT_FILTER_BLANK_TOKENS = True # Options: True or False. PARAKEET_TDT_FILTER_SPECIAL_TOKEN_IDS = True # Options: True or False. PROJECTOR_INTERMEDIATE_SIZE = 4096 # Options: integer; the packaged model uses 4096. PROJECTOR_DROPOUT = 0.0 # Options: float; inference should normally be 0.0. HYBRID_ENCODER_GATE_INIT = 0.0 # Options: float; saved audio projector weights override the initial gate. # Prompt settings. Empty PROMPT_OVERRIDE builds the standard dynamic prompt. PROMPT_OVERRIDE = "" # Options: "" or any full custom prompt string. CAPTION_SETTINGS_JSON_PATH = "" # Options: "" or a JSON path under /workspace/dataset_jsons to override the settings below. CAPTION_LENGTH = "very large" # Options: "very small", "small", "medium", "large", "very large". INCLUDE_WATERMARK_INFO = False # Options: True or False. VULGARITY = "low" # Options: "none", "low", "medium", "high". UNCERTAINTY = "low" # Options: "none", "low", "medium", "high". CHARACTER_NAMES = "none" # Options: "none", "ambiguous", "single", "multiple". FLUFF = "none" # Options: "none", "low", "medium", "high". HAS_REPETITION = False # Options: True or False. SPECULATION = "low" # Options: "none", "low", "medium", "high". TEMPORAL_DETAIL = "medium" # Options: "static", "low", "medium", "high". VISUAL_SPECIFICITY = "moderate" # Options: "generic", "moderate", "detailed", "excessive". CAMERA_DETAIL = "medium" # Options: "none", "low", "medium", "high". CAPTION_STYLE = "plain" # Options: "plain", "verbose", "ornate", "robotic". HAS_THINKING = True # Options: True to request thought JSON plus final caption, False to request only the final caption. # Media settings. Training used separate sidecar audio and random frame counts. NUM_FRAMES = 12 # Options: None for processor default, or an integer frame count. FPS = None # Options: None for processor default, or a float such as 1.0. SAMPLING_RATE = 16_000 # Options: normally 16000. AUDIO_MAX_LENGTH_SAMPLES = 0 # Options: 0 keeps full audio; positive integer truncates Parakeet audio. MAX_AUDIO_SECONDS = 0.0 # Options: 0.0 keeps full audio; positive float caps extracted sidecar audio. # Generation settings. MAX_NEW_TOKENS = 1200 # Options: positive integer token cap. TEMPERATURE = 0.0 # Options: 0.0 for greedy decoding, >0.0 for sampling. TOP_P = 0.9 # Options: float in (0, 1], used only when TEMPERATURE > 0. REPETITION_PENALTY = 1.1 # Options: 1.0 disables the penalty, >1.0 penalizes repetition. PRINT_INPUT_STATS = False # Options: True or False. QUIET_MODEL_LOAD = True # Options: True hides noisy missing-key load reports; False prints full loader output. # Usually leave these alone. LOCAL_FILES_ONLY = True # Options: True to use cached files only, False to allow downloads. DTYPE = torch.bfloat16 # Options: torch.bfloat16, torch.float16, torch.float32. DEVICE_MAP = "auto" # Options: "auto", "cuda", or another Transformers device_map value. ATTN_IMPLEMENTATION = "sdpa" # Options: "sdpa", "flash_attention_2", None. warnings.filterwarnings( "ignore", message=r"RNN module weights are not part of single contiguous chunk of memory.*", category=UserWarning, ) @contextlib.contextmanager def quiet_model_load() -> Any: if not QUIET_MODEL_LOAD: yield return load_report_logger = logging.getLogger("transformers.utils.loading_report") old_level = load_report_logger.level load_report_logger.setLevel(logging.ERROR) patched_modules: list[tuple[Any, Any]] = [] try: import transformers.modeling_utils as modeling_utils import transformers.utils.loading_report as loading_report original_report = modeling_utils.log_state_dict_report def quiet_report( model: Any, pretrained_model_name_or_path: str, ignore_mismatched_sizes: bool, loading_info: Any, logger: logging.Logger | None = None, ) -> None: has_fatal_issue = bool(getattr(loading_info, "error_msgs", None)) or bool( getattr(loading_info, "conversion_errors", None) ) if not ignore_mismatched_sizes and bool(getattr(loading_info, "mismatched_keys", None)): has_fatal_issue = True if has_fatal_issue: original_report( model, pretrained_model_name_or_path, ignore_mismatched_sizes, loading_info, logger=logger, ) for module in (loading_report, modeling_utils): patched_modules.append((module, module.log_state_dict_report)) module.log_state_dict_report = quiet_report except Exception: patched_modules = [] try: with contextlib.redirect_stdout(io.StringIO()), contextlib.redirect_stderr(io.StringIO()): yield finally: for module, original in patched_modules: module.log_state_dict_report = original load_report_logger.setLevel(old_level) def load_parakeet_projector(model: torch.nn.Module) -> None: state_path = Path(AUDIO_PROJECTOR_PATH) state = load_state_file(state_path) module = gemma_core(model).embed_audio module.load_state_dict(state, strict=True) def video_has_audio_stream(video_path: Path) -> bool: cmd = [ "ffprobe", "-v", "error", "-select_streams", "a:0", "-show_entries", "stream=index", "-of", "csv=p=0", str(video_path), ] result = subprocess.run(cmd, check=True, capture_output=True, text=True) return bool(result.stdout.strip()) def probe_video_duration_seconds(video_path: Path) -> float: cmd = [ "ffprobe", "-v", "error", "-show_entries", "format=duration", "-of", "default=noprint_wrappers=1:nokey=1", str(video_path), ] result = subprocess.run(cmd, check=True, capture_output=True, text=True) duration = float(result.stdout.strip()) if duration <= 0: raise RuntimeError(f"Video duration must be positive: {video_path}") return duration def extract_audio(video_path: Path, audio_path: Path) -> None: cmd = [ "ffmpeg", "-hide_banner", "-loglevel", "error", "-y", "-i", str(video_path), "-vn", "-ac", "1", "-ar", str(SAMPLING_RATE), ] if MAX_AUDIO_SECONDS > 0: cmd.extend(["-t", f"{MAX_AUDIO_SECONDS:.6f}"]) cmd.extend(["-c:a", "pcm_s16le", str(audio_path)]) subprocess.run(cmd, check=True) def create_silent_audio(audio_path: Path, duration_seconds: float) -> None: if MAX_AUDIO_SECONDS > 0: duration_seconds = min(duration_seconds, MAX_AUDIO_SECONDS) cmd = [ "ffmpeg", "-hide_banner", "-loglevel", "error", "-y", "-f", "lavfi", "-i", f"anullsrc=channel_layout=mono:sample_rate={SAMPLING_RATE}", "-t", f"{duration_seconds:.6f}", "-ac", "1", "-ar", str(SAMPLING_RATE), "-c:a", "pcm_s16le", str(audio_path), ] subprocess.run(cmd, check=True) def prepare_sidecar_audio(video_path: Path, tmpdir: Path) -> tuple[Path, bool]: audio_path = tmpdir / "sidecar_audio.wav" if video_has_audio_stream(video_path): extract_audio(video_path, audio_path) return audio_path, True duration_seconds = probe_video_duration_seconds(video_path) print(f"input_video_has_audio=false; creating_silent_sidecar_audio duration_seconds={duration_seconds:.3f}", flush=True) create_silent_audio(audio_path, duration_seconds) return audio_path, False def bool_from_json(value: Any, field_name: str) -> bool: if isinstance(value, bool): return value if isinstance(value, str): lowered = value.strip().lower() if lowered in {"1", "true", "yes", "y", "on"}: return True if lowered in {"0", "false", "no", "n", "off"}: return False raise ValueError(f"{field_name} must be boolean-like, got {value!r}") def load_prompt_settings_json() -> dict[str, Any]: if not CAPTION_SETTINGS_JSON_PATH.strip(): return {} path = Path(CAPTION_SETTINGS_JSON_PATH) data = json.loads(path.read_text(encoding="utf-8")) if not isinstance(data, dict): raise ValueError(f"CAPTION_SETTINGS_JSON_PATH must point to a JSON object: {path}") return data def build_prompt() -> str: if PROMPT_OVERRIDE.strip(): return PROMPT_OVERRIDE.strip() settings: dict[str, Any] = { "caption_length": CAPTION_LENGTH, "include_watermark_info": INCLUDE_WATERMARK_INFO, **DEFAULT_CAPTION_SETTING_VALUES, "vulgarity": VULGARITY, "uncertainty": UNCERTAINTY, "character_names": CHARACTER_NAMES, "fluff": FLUFF, "has_repetition": HAS_REPETITION, "speculation": SPECULATION, "temporal_detail": TEMPORAL_DETAIL, "visual_specificity": VISUAL_SPECIFICITY, "camera_detail": CAMERA_DETAIL, "caption_style": CAPTION_STYLE, "has_thinking": HAS_THINKING, } settings.update(load_prompt_settings_json()) settings["has_thinking"] = HAS_THINKING settings["caption_length"] = str(settings["caption_length"]).strip().lower() if settings["caption_length"] not in CAPTION_LENGTH_LABELS: raise ValueError(f"caption_length must be one of {CAPTION_LENGTH_LABELS}, got {settings['caption_length']!r}") settings["include_watermark_info"] = bool_from_json(settings["include_watermark_info"], "include_watermark_info") settings["has_repetition"] = bool_from_json(settings["has_repetition"], "has_repetition") settings["has_thinking"] = bool_from_json(settings["has_thinking"], "has_thinking") for field_name, allowed in CAPTION_SETTING_FIELD_CHOICES.items(): value = str(settings[field_name]).strip().lower() if value not in allowed: raise ValueError(f"{field_name} must be one of {allowed}, got {value!r}") settings[field_name] = value return format_caption_settings_prompt(settings) def build_messages(video_path: Path, audio_path: Path, prompt: str) -> list[dict[str, Any]]: return [ { "role": "user", "content": [ {"type": "video", "path": str(video_path)}, {"type": "text", "text": prompt}, {"type": "audio", "path": str(audio_path)}, ], }, {"role": "assistant", "content": [{"type": "text", "text": ""}]}, ] def trim_empty_assistant_terminator(inputs: dict[str, torch.Tensor], processor: Any) -> dict[str, torch.Tensor]: eos_tail = processor.tokenizer.encode("\n", add_special_tokens=False) if not eos_tail: return inputs tail_len = len(eos_tail) input_ids = inputs["input_ids"][0] if input_ids[-tail_len:].tolist() != eos_tail: return inputs trimmed = {} for key, value in inputs.items(): if isinstance(value, torch.Tensor) and value.ndim >= 2 and value.shape[1] == input_ids.shape[0]: trimmed[key] = value[:, :-tail_len] else: trimmed[key] = value return trimmed def tensor_stats(tensor: torch.Tensor | None, mask: torch.Tensor | None = None) -> dict[str, Any]: if tensor is None: return {"present": False} stats_tensor = tensor.detach().float().cpu() result: dict[str, Any] = { "present": True, "shape": list(tensor.shape), "mean": round(float(stats_tensor.mean().item()), 8), "std": round(float(stats_tensor.std().item()), 8), "abs_mean": round(float(stats_tensor.abs().mean().item()), 8), } if mask is not None: result["mask_shape"] = list(mask.shape) result["mask_sum"] = int(mask.detach().cpu().sum().item()) return result def print_input_stats(inputs: dict[str, torch.Tensor], label: str) -> None: stats = { "input_ids_shape": list(inputs["input_ids"].shape), "input_features": tensor_stats(inputs.get("input_features"), inputs.get("input_features_mask")), "keys": sorted(inputs.keys()), } print(f"{label}=" + json.dumps(stats, sort_keys=True), flush=True) def move_inputs_to_model_device(inputs: dict[str, Any], model: torch.nn.Module) -> dict[str, Any]: device = getattr(model, "device", None) if device is None: try: device = next(model.parameters()).device except StopIteration: device = torch.device("cuda" if torch.cuda.is_available() else "cpu") moved = {} for key, value in inputs.items(): moved[key] = value.to(device) if isinstance(value, torch.Tensor) else value return moved def generation_kwargs() -> dict[str, Any]: kwargs: dict[str, Any] = { "max_new_tokens": MAX_NEW_TOKENS, "do_sample": TEMPERATURE > 0, "repetition_penalty": REPETITION_PENALTY, "use_cache": True, } if TEMPERATURE > 0: kwargs["temperature"] = TEMPERATURE kwargs["top_p"] = TOP_P return kwargs def prepare_inputs( processor: Any, parakeet_processor: Any, video_path: Path, audio_path: Path, prompt: str, ) -> dict[str, torch.Tensor]: processor_kwargs: dict[str, Any] = { "padding": True, "truncation": False, "sampling_rate": SAMPLING_RATE, } if NUM_FRAMES is not None: processor_kwargs["num_frames"] = NUM_FRAMES if FPS is not None: processor_kwargs["fps"] = FPS inputs = processor.apply_chat_template( build_messages(video_path, audio_path, prompt), tokenize=True, return_dict=True, return_tensors="pt", load_audio_from_video=False, processor_kwargs=processor_kwargs, ) inputs = trim_empty_assistant_terminator(inputs, processor) if PRINT_INPUT_STATS: print_input_stats(inputs, "input_stats_before_parakeet_swap") if PARAKEET_NATIVE_FEATURES: replace_batch_audio_features( inputs, audio_paths=[str(audio_path)], parakeet_processor=parakeet_processor, sampling_rate=SAMPLING_RATE, max_length_samples=AUDIO_MAX_LENGTH_SAMPLES, ) if PRINT_INPUT_STATS: print_input_stats(inputs, "input_stats_after_parakeet_swap") return inputs def load_model(processor_path: Path) -> tuple[Any, torch.nn.Module, Any]: processor = AutoProcessor.from_pretrained(str(processor_path), local_files_only=LOCAL_FILES_ONLY) parakeet_processor = AutoProcessor.from_pretrained(PARAKEET_MODEL_ID, local_files_only=LOCAL_FILES_ONLY) model_kwargs: dict[str, Any] = { "local_files_only": LOCAL_FILES_ONLY, "dtype": DTYPE, "low_cpu_mem_usage": True, "device_map": DEVICE_MAP, } if ATTN_IMPLEMENTATION: model_kwargs["attn_implementation"] = ATTN_IMPLEMENTATION with quiet_model_load(): model = AutoModelForMultimodalLM.from_pretrained(MODEL_PATH, **model_kwargs) bridge_args = type( "BridgeArgs", (), { "parakeet_model_id": PARAKEET_MODEL_ID, "parakeet_bridge_mode": PARAKEET_BRIDGE_MODE, "local_files_only": LOCAL_FILES_ONLY, "projector_intermediate_size": PROJECTOR_INTERMEDIATE_SIZE, "projector_dropout": PROJECTOR_DROPOUT, "hybrid_encoder_gate_init": HYBRID_ENCODER_GATE_INIT, "parakeet_tdt_filter_blank_tokens": PARAKEET_TDT_FILTER_BLANK_TOKENS, "parakeet_tdt_filter_special_token_ids": PARAKEET_TDT_FILTER_SPECIAL_TOKEN_IDS, }, )() with quiet_model_load(): install_parakeet_audio_bridge(model, bridge_args) gemma_core(model) load_parakeet_projector(model) model.eval() return processor, model, parakeet_processor def main() -> None: video_path = Path(VIDEO_PATH) model_path = Path(MODEL_PATH) processor_path = Path(PROCESSOR_PATH) audio_projector_path = Path(AUDIO_PROJECTOR_PATH) parakeet_path = Path(PARAKEET_MODEL_ID) for label, path in ( ("VIDEO_PATH", video_path), ("MODEL_PATH", model_path), ("processor", processor_path), ("audio_projector", audio_projector_path), ("parakeet", parakeet_path), ): if not path.exists(): raise FileNotFoundError(f"{label} does not exist: {path}") prompt = build_prompt() processor, model, parakeet_processor = load_model(processor_path) with tempfile.TemporaryDirectory(prefix="gemma4_caption_inference_") as tmpdir_raw: tmpdir = Path(tmpdir_raw) audio_path, _had_audio = prepare_sidecar_audio(video_path, tmpdir) inputs = prepare_inputs(processor, parakeet_processor, video_path, audio_path, prompt) moved_inputs = move_inputs_to_model_device(inputs, model) input_len = moved_inputs["input_ids"].shape[-1] with torch.inference_mode(): output_ids = model.generate(**moved_inputs, **generation_kwargs()) new_tokens = output_ids[0, input_len:] response = processor.decode(new_tokens, skip_special_tokens=True).strip() print(response, flush=True) if __name__ == "__main__": main()