""" RGB-to-all chained generation pipeline for the Gradio Space. Adapted from cvpr_fvd_videos_no_poses_raw_rgb.py: same model architecture, same chained-generation logic, same raw-video tokenization path. What's different: no ground-truth comparison (Space users don't have GT for their own uploads), no CLI/argparse, models are loaded once and reused across requests, and checkpoints come from the Hugging Face Hub instead of cluster scratch paths. human_poses is intentionally excluded from the generation target chain (it needs TokHMR + SMPL + an EGL renderer, which we're not bundling in v1), but it stays in the model's modality_info/data_config so the architecture still matches the trained checkpoint's state_dict. """ import json import os import sys from pathlib import Path from types import SimpleNamespace import numpy as np import torch import torch.backends.cudnn as cudnn import yaml from tokenizers import Tokenizer from einops import rearrange from decord import VideoReader, cpu import torchvision.transforms.functional as TF import torchvision.transforms as T from huggingface_hub import hf_hub_download import fourm.utils as utils from fourm.data.modality_info import MODALITY_INFO from fourm.utils import create_model from fourm.models import fm # noqa: F401 -- import side effect: registers model architectures (e.g. fm_large_24e_24d_swiglu_qknorm_nobias) from fourm.models.generate import ( GenerationSampler, build_chained_generation_schedules, create_frame_ids, init_empty_target_modality, ) from fourm.data.modality_transforms import VideoDetectionTransform from fourm.utils.plotting_utils import visualize_temporal_bboxes, decode_dict_text_modalities from fourm.vq.vqvae import VQVAE sys.path.append(os.path.dirname(__file__)) sys.path.append(os.path.join(os.path.dirname(__file__), "my_VidTok")) from scripts.inference_evaluate import load_model_from_config from my_VidTok.scripts.inference_evaluate_m import load_model_from_config_m from helper_functions import ( tensor_to_uint8, resize_and_duplicate, FeatureToPCAConverter, image_mask_first_frame_conditional, pop_conditioning_domain, load_and_decode_tokens, convert_raw_optical_flow, concat_videos_horizontally, save_video_with_imageio, merge_tokens_with_frames, merge_detection_tokens_with_sentinel_tokens, transform_tensor_with_markers, ) ROOT = Path(__file__).parent VJEPA_TOKEN_COUNT = 1024 FPS = 4.0 # Two separate Hub repos: the model checkpoint changes often as training # progresses, the tokenizers are stable, so they're versioned independently. # Set via env vars so this can be pointed elsewhere without editing code. MODEL_REPO = os.environ.get("FOURM_MODEL_REPO", "EPFL-VILAB/Video-4M-models") TOKENIZERS_REPO = os.environ.get("FOURM_TOKENIZERS_REPO", "EPFL-VILAB/Video-4M-tokenizers") # Filenames expected inside each repo -- one folder per modality. Keep in # sync with scripts/upload_checkpoints_to_hub.py. MODEL_WEIGHT_FILES = { "main_model": "main_model/checkpoint.pth", } TOKENIZER_WEIGHT_FILES = { "vidtok_rgb": "rgb/ckpt.ckpt", "vidtok_normal": "surface-normals/ckpt.ckpt", "vidtok_depth": "depth/ckpt.ckpt", "vidtok_opticalflow": "opticalflow/ckpt.ckpt", "vjepa": "v-jepa-2/ckpt.ckpt", "dinov2": "dinov2/ckpt.ckpt", "siglipv2": "siglip-2/ckpt.ckpt", } VIDTOK_CFG_PATH = str(ROOT / "configs" / "vidtok_fsq_causal_488_32768.yaml") VJEPA_CFG_PATH = str(ROOT / "my_VidTok" / "configs" / "vjepa_l1_176_211_16807.yaml") DINOV2_CFG_PATH = str(ROOT / "my_VidTok" / "configs" / "dinov2_l1_176_411_16807.yaml") SIGLIPV2_CFG_PATH = str(ROOT / "my_VidTok" / "configs" / "siglip_ens_176_411_16807.yaml") MODEL_CONFIG_PATH = str(ROOT / "configs" / "model_config.yaml") # The 9 modalities we generate from an RGB video, excluding human_poses. # tuple layout matches the CONFIGS dict in the original scripts: # (modality_key, autoregression_scheme, decoding_steps, token_decoding_schedule, # tokens_per_target, temperature, cfg_scale) DETECTION_SIZE = 750 CONFIGS = { "siglip": ("tok_video_siglipv2@224", "roar", 50, "linear", 980, 1.0, 2.0), "dinov2": ("tok_video_dinov2@224", "roar", 20, "linear", 1280, 3.0, 3.0), "vjepa": ("tok_video_vjepa@224", "roar", 20, "linear", VJEPA_TOKEN_COUNT, 5.0, 1.0), "caption": ("caption", "autoregressive", None, None, 256, 3.0, 1.0), "transcription": ("transcription", "autoregressive", None, None, 450, 3.0, 1.0), "det": ("det", "autoregressive", None, None, DETECTION_SIZE, 2.0, 1.0), "depth": ("tok_video_depth@128", "roar", 50, "linear", 1280, 0.5, 2.0), "normal": ("tok_video_normal@128", "roar", 50, "linear", 1280, 0.5, 2.0), "opticalflow": ("tok_video_opticalflow@128", "roar", 30, "linear", 1280, 0.2, 2.0), "rgb": ("tok_video_rgb@128", "roar", 100, "linear", 1280, 0.1, 2.0), } DEFAULT_CHAIN = ["siglip", "transcription", "vjepa", "dinov2", "caption", "det", "depth", "normal", "opticalflow"] # Reverse lookup: internal domain string (e.g. "tok_video_rgb@128") -> short UI key. _DOMAIN_TO_KEY = {cfg[0]: key for key, cfg in CONFIGS.items()} # Two paired (input_modality, chain, hyperparameters) presets for the UI's # "quick start" buttons -- dense (rgb) conditioning needs different # temperature/CFG/decoding-steps tuning than sparse (text) conditioning # driving video generation, even though the structural CONFIGS above # (scheme, token budget, decoding schedule) stay the same either way. CHAIN_PRESETS = { "rgb_to_others": { "input_modality": "rgb", "chain": ["depth", "normal", "siglip", "det", "caption", "transcription", "dinov2", "vjepa", "opticalflow"], }, "text_to_rgb": { "input_modality": "caption", "chain": ["transcription", "siglip", "det", "depth", "rgb"], }, } # modality -> {temp, cfg, decoding_steps (only for 'roar'-scheme modalities)} HYPERPARAM_PRESETS = { "rgb_to_others": { "siglip": {"decoding_steps": 50, "temp": 0.01, "cfg": 2.0}, "dinov2": {"decoding_steps": 50, "temp": 0.01, "cfg": 2.0}, "vjepa": {"decoding_steps": 50, "temp": 0.01, "cfg": 2.0}, "caption": {"temp": 0.1, "cfg": 1.0}, "transcription": {"temp": 0.1, "cfg": 1.0}, "det": {"temp": 0.7, "cfg": 1.0}, "depth": {"decoding_steps": 100, "temp": 0.01, "cfg": 2.0}, "normal": {"decoding_steps": 100, "temp": 0.01, "cfg": 2.0}, "opticalflow": {"decoding_steps": 10, "temp": 0.1, "cfg": 2.0}, "rgb": {"decoding_steps": 50, "temp": 0.01, "cfg": 2.0}, }, "text_to_rgb": { "siglip": {"decoding_steps": 50, "temp": 1.0, "cfg": 2.0}, "dinov2": {"decoding_steps": 20, "temp": 3.0, "cfg": 3.0}, "vjepa": {"decoding_steps": 20, "temp": 5.0, "cfg": 1.0}, "caption": {"temp": 3.0, "cfg": 1.0}, "transcription": {"temp": 3.0, "cfg": 1.0}, "det": {"temp": 2.0, "cfg": 1.0}, "depth": {"decoding_steps": 30, "temp": 0.5, "cfg": 2.0}, "normal": {"decoding_steps": 30, "temp": 0.5, "cfg": 2.0}, "opticalflow": {"decoding_steps": 30, "temp": 0.2, "cfg": 2.0}, "rgb": {"decoding_steps": 100, "temp": 0.1, "cfg": 2.0}, }, } # Full RGB conditioning: all 1280 tokens are given as input, none need to be # completed (matches cvpr_fvd_videos_no_poses.py's `not partial_conditioning` # path, as opposed to cvpr_fvd_videos_no_poses_raw_rgb.py's 512-seed-token # partial-completion scheme). RGB_SEED_TOKENS = 1280 RGB_TARGET_TOKENS = 0 _STATE = {"loaded": False} # For local/cluster testing before the Hub repos are usable (e.g. private # storage quota not sorted out yet): set these env vars to point straight at # the checkpoints on disk, skipping hf_hub_download entirely. Leave unset to # download from the Hub as normal (that's the path a real Space will use). _LOCAL_ENV_VARS = { "main_model": "FOURM_LOCAL_MODEL_PATH", "vidtok_rgb": "FOURM_LOCAL_VIDTOK_RGB", "vidtok_normal": "FOURM_LOCAL_VIDTOK_NORMAL", "vidtok_depth": "FOURM_LOCAL_VIDTOK_DEPTH", "vidtok_opticalflow": "FOURM_LOCAL_VIDTOK_OPTICALFLOW", "vjepa": "FOURM_LOCAL_VJEPA", "dinov2": "FOURM_LOCAL_DINOV2", "siglipv2": "FOURM_LOCAL_SIGLIPV2", } def _download(name): local_path = os.environ.get(_LOCAL_ENV_VARS[name]) if local_path: return local_path if name in MODEL_WEIGHT_FILES: return hf_hub_download(repo_id=MODEL_REPO, filename=MODEL_WEIGHT_FILES[name]) return hf_hub_download(repo_id=TOKENIZERS_REPO, filename=TOKENIZER_WEIGHT_FILES[name]) # --- Any-to-any curated examples --------------------------------------------- # One HF *dataset* repo, uploaded by scripts/upload_examples_to_hub.py, laid # out exactly like the source folder on the cluster: one subfolder per # modality, files named by a shared "stem" per clip. EXAMPLES_REPO = os.environ.get("FOURM_EXAMPLES_REPO", "EPFL-VILAB/Video-4M-examples") # For local/cluster testing: point straight at a local copy of that folder # layout instead of downloading from the Hub (mirrors _LOCAL_ENV_VARS above). _LOCAL_EXAMPLES_DIR = os.environ.get("FOURM_LOCAL_EXAMPLES_DIR") _EXAMPLE_SUBFOLDERS = { "rgb": ("tok_video_rgb@128", ".npy"), "depth": ("tok_video_depth@128", ".npy"), "normal": ("tok_video_normal@128", ".npy"), "opticalflow": ("tok_video_opticalflow@128", ".npy"), "siglip": ("tok_video_siglipv2@224", ".npy"), "dinov2": ("tok_video_dinov2@224", ".npy"), "vjepa": ("tok_video_vjepa@224", ".npy"), "det": ("det", ".json"), "caption": ("caption", ".json"), "transcription": ("transcription", ".json"), "crop_settings": ("crop_settings", ".npy"), } def _example_file(key, stem): subfolder, ext = _EXAMPLE_SUBFOLDERS[key] path_in_repo = f"{subfolder}/{stem}{ext}" if _LOCAL_EXAMPLES_DIR: return os.path.join(_LOCAL_EXAMPLES_DIR, path_in_repo) return hf_hub_download(repo_id=EXAMPLES_REPO, filename=path_in_repo, repo_type="dataset") def _list_examples_from_manifest(manifest_filename): if _LOCAL_EXAMPLES_DIR: manifest_path = os.path.join(_LOCAL_EXAMPLES_DIR, manifest_filename) else: manifest_path = hf_hub_download(repo_id=EXAMPLES_REPO, filename=manifest_filename, repo_type="dataset") with open(manifest_path) as f: return json.load(f)["examples"] def list_examples(): """Returns the list of curated example stems available for the any-to-any tab.""" return _list_examples_from_manifest("examples.json") # Separate, hand-picked set of stems specifically suited for the Future # Prediction tab (e.g. clips with clear, consistent motion), uploaded via # scripts/upload_examples_to_hub.py --stems_file ... --manifest_repo_filename # future_examples.json. Lives in the same EXAMPLES_REPO/local dir as the # any-to-any set, just under a different manifest filename so the two don't collide. FUTURE_EXAMPLES_MANIFEST_FILENAME = os.environ.get("FOURM_FUTURE_EXAMPLES_MANIFEST", "future_examples.json") def list_future_examples(): """Returns the list of curated example stems available for the Future Prediction tab.""" return _list_examples_from_manifest(FUTURE_EXAMPLES_MANIFEST_FILENAME) def _load_example_raw(stem): """Loads one curated example's pre-tokenized data for every modality.""" return { "rgb": np.load(_example_file("rgb", stem))[0], "depth": np.load(_example_file("depth", stem))[0], "normal": np.load(_example_file("normal", stem))[0], "opticalflow": np.load(_example_file("opticalflow", stem))[0], "siglip": np.load(_example_file("siglip", stem))[0], "dinov2": np.load(_example_file("dinov2", stem))[0], "vjepa": np.load(_example_file("vjepa", stem))[0], "crop_settings": np.load(_example_file("crop_settings", stem)), "detection_dict": json.load(open(_example_file("det", stem), "rb")), "caption": json.load(open(_example_file("caption", stem), "rb"))[0], "transcription": json.load(open(_example_file("transcription", stem), "rb"))[0], } # Pre-rendered (at curation time, via scripts/upload_examples_to_hub.py + # visualize_multimodal_pretraining_data_13_modalities.py) previews, so the UI # can show what an example/modality looks like without running the model. PREVIEW_VISUAL_KEYS = ["rgb", "depth", "normal", "opticalflow", "dinov2", "siglip", "vjepa", "det"] def get_example_preview(stem, modality_key): """Returns a local mp4 path for visual modalities, or a text string for caption/transcription. Returns None if no preview exists yet (e.g. the preview upload phase hasn't been run for this example). """ if modality_key == "caption": return json.load(open(_example_file("caption", stem), "rb"))[0] if modality_key == "transcription": return " ".join(json.load(open(_example_file("transcription", stem), "rb"))[0]) if modality_key not in PREVIEW_VISUAL_KEYS: return None path_in_repo = f"preview/{modality_key}/{stem}.mp4" try: if _LOCAL_EXAMPLES_DIR: local_path = os.path.join(_LOCAL_EXAMPLES_DIR, path_in_repo) return local_path if os.path.exists(local_path) else None return hf_hub_download(repo_id=EXAMPLES_REPO, filename=path_in_repo, repo_type="dataset") except Exception: return None SPECIAL_TOKENS_DET_NAMES = [f"[FRAME_{i}]" for i in range(1, 18)] SPECIAL_TOKENS_TRANSCRIPTION_NAMES = ["[SEC_1]", "[SEC_2]", "[SEC_3]", "[SEC_4]"] def _prepare_detections_text(detection_dict, crop_settings): """Builds the same tokenized-text representation of detections used as both conditioning input and GT-visualization input for the det modality. """ detection_transforms = VideoDetectionTransform( det_threshold=0.2, det_max_instances=None, bbox_order="dist_to_orig", coord_bins=1000, min_visibility=0.0 ) starting_time, ending_time, i, j, h, w, h_flip = crop_settings[0] processed_detections = detection_transforms.image_augment( detection_dict, (starting_time, ending_time, i, j, h, w), False, None, (None, None), None, None ) gt_processed_detections, per_frame_instance_counts = detection_transforms.postprocess(processed_detections) return merge_tokens_with_frames(gt_processed_detections, SPECIAL_TOKENS_DET_NAMES, per_frame_instance_counts) def _prepare_transcription_text(transcription_list): return " ".join(f"{token} {text} [EOS]" for token, text in zip(SPECIAL_TOKENS_TRANSCRIPTION_NAMES, transcription_list)) def get_example_transcription_tagged(stem): """Returns the transcription in its fully tagged, model-ready format ([SEC_1] ... [EOS] [SEC_2] ... [EOS] ...) -- unlike get_example_preview's plain-text version (used for read-only display elsewhere), this is what the Future Prediction tab shows for editing: the tags are structurally required by the model and must not be removed. """ transcription = json.load(open(_example_file("transcription", stem), "rb"))[0] return _prepare_transcription_text(transcription) def _set_full_conditioning( batched_sample, modality_key, example, text_tokenizer, device, override_caption_text=None, override_transcription_text=None, ): """Fills batched_sample with full-conditioning tensors for modality_key (the single input modality for any-to-any generation). Ported 1:1 from the `not partial_conditioning` branches in cvpr_fvd_videos_no_poses.py's per-modality conditioning setup -- human_poses/class_condition dropped (out of scope for this Space), and the vjepa branch's tokens_dinov2->tokens_vjepa copy-paste bug fixed. override_caption_text / override_transcription_text: optional user-edited text to condition on instead of the example's own caption/transcription (used by the Future Prediction tab's editable extra-conditioning box). override_transcription_text is expected already in the fully tagged format (see get_example_transcription_tagged) -- it's tokenized as-is, not re-wrapped by _prepare_transcription_text. """ eos_id = text_tokenizer.token_to_id("[EOS]") if modality_key == "caption": caption_text = override_caption_text if override_caption_text is not None else example["caption"] seq_ids = text_tokenizer.encode(caption_text).ids + [eos_id] tensor = torch.tensor(seq_ids).unsqueeze(0).to(device) input_mask = torch.zeros(len(seq_ids), dtype=torch.bool).unsqueeze(0).to(device) target_mask = torch.ones(len(seq_ids), dtype=torch.bool).unsqueeze(0).to(device) batched_sample["caption"] = { "tensor": tensor, "input_mask": input_mask, "target_mask": target_mask, "decoder_attention_mask": input_mask, } elif modality_key == "transcription": transcription_text = ( override_transcription_text if override_transcription_text is not None else _prepare_transcription_text(example["transcription"]) ) seq_ids = text_tokenizer.encode(transcription_text).ids tensor = torch.tensor(seq_ids).unsqueeze(0).to(device) input_mask = torch.zeros(len(seq_ids), dtype=torch.bool).unsqueeze(0).to(device) target_mask = torch.ones(len(seq_ids), dtype=torch.bool).unsqueeze(0).to(device) batched_sample["transcription"] = {"tensor": tensor, "input_mask": input_mask, "target_mask": target_mask, "decoder_attention_mask": input_mask} elif modality_key == "det": sequence = _prepare_detections_text(example["detection_dict"], example["crop_settings"]) seq_ids = text_tokenizer.encode(sequence).ids tensor_batched = torch.tensor(seq_ids).unsqueeze(0).to(device) input_mask = torch.zeros(len(seq_ids), dtype=torch.bool).unsqueeze(0).to(device) target_mask = torch.ones(len(seq_ids), dtype=torch.bool).unsqueeze(0).to(device) batched_sample["det"] = { "tensor": tensor_batched, "input_mask": input_mask, "tensor_frame_ids": create_frame_ids(tensor_batched), "target_mask": target_mask, "decoder_attention_mask": torch.zeros(target_mask.shape, dtype=torch.bool, device=device), } elif modality_key in ("rgb", "depth", "normal", "opticalflow", "dinov2"): domain = CONFIGS[modality_key][0] batched_sample[domain] = image_mask_first_frame_conditional(example[modality_key], 1280, 1280, 0) elif modality_key == "siglip": domain = CONFIGS[modality_key][0] batched_sample[domain] = image_mask_first_frame_conditional(example["siglip"], 980, 980, 0) elif modality_key == "vjepa": domain = CONFIGS[modality_key][0] batched_sample[domain] = image_mask_first_frame_conditional( example["vjepa"], VJEPA_TOKEN_COUNT, VJEPA_TOKEN_COUNT, 0 ) else: raise ValueError(f"Unsupported input modality: {modality_key}") def _build_args(): with open(MODEL_CONFIG_PATH) as f: cfg = yaml.safe_load(f) args = SimpleNamespace(**cfg) args.num_register_tokens = 0 args.min_input_tokens = args.num_input_tokens args.min_target_tokens = args.num_target_tokens data_config_path = args.data_config if not os.path.isabs(data_config_path): data_config_path = str(ROOT / data_config_path) with open(data_config_path) as f: data_config = yaml.safe_load(f) train_config = data_config["train"]["datasets"] args.in_domains = sorted(set.union(*[set(c["in_domains"].split("-")) for c in train_config.values()])) args.out_domains = sorted(set.union(*[set(c["out_domains"].split("-")) for c in train_config.values()])) args.all_domains = sorted(set(args.in_domains) | set(args.out_domains)) return args def _setup_modality_info(args): modality_info = {mod: MODALITY_INFO[mod] for mod in args.all_domains} for mod in modality_info: image_size = modality_info[mod].get("input_size", args.input_size) patch_size = modality_info[mod].get("patch_size", args.patch_size) num_patches = (image_size // patch_size) ** 2 if modality_info[mod]["type"] == "img": if "tok_video_vjepa" in mod: modality_info[mod]["max_tokens"] = VJEPA_TOKEN_COUNT elif "tok_video" in mod: assert args.total_frames_to_be_extracted % 2 != 0 modality_info[mod]["max_tokens"] = num_patches * ((args.total_frames_to_be_extracted // 4) + 1) elif "video_rgb" in mod: modality_info[mod]["max_tokens"] = num_patches * (args.total_frames_to_be_extracted // 4) return modality_info def _build_model(args, modality_info): encoder_embeddings = {} for mod in args.in_domains: info = modality_info[mod] if info.get("encoder_embedding") is not None: if info["type"] == "img": image_size = info.get("input_size", args.input_size) patch_size = info.get("patch_size", args.patch_size) encoder_embeddings[mod] = info["encoder_embedding"](patch_size=patch_size, image_size=image_size) else: encoder_embeddings[mod] = info["encoder_embedding"]() decoder_embeddings = {} for mod in args.out_domains: info = modality_info[mod] if info.get("decoder_embedding") is not None: if info["type"] == "img": image_size = info.get("input_size", args.input_size) patch_size = info.get("patch_size", args.patch_size) decoder_embeddings[mod] = info["decoder_embedding"](patch_size=patch_size, image_size=image_size) else: decoder_embeddings[mod] = info["decoder_embedding"]() return create_model( args.model, encoder_embeddings=encoder_embeddings, decoder_embeddings=decoder_embeddings, modality_info=modality_info, num_register_tokens=args.num_register_tokens, ) def load_pipeline(device="cuda"): """Loads the main 4M model plus all tokenizers once and caches them in _STATE. Safe to call repeatedly -- only does the (slow) loading work on the first call. Must be called from within GPU-provisioned code on ZeroGPU Spaces. """ if _STATE["loaded"]: return _STATE cudnn.benchmark = True args = _build_args() text_tokenizer = Tokenizer.from_file(str(ROOT / args.tokenizer_path)) modality_info = _setup_modality_info(args) model = _build_model(args, modality_info).to(device) state_dict = torch.load(_download("main_model"), map_location="cpu", weights_only=False)["model"] state_dict = {k.replace("_orig_mod.", ""): v for k, v in state_dict.items()} model.load_state_dict(state_dict, strict=True) model.eval() human_pose_tokenizer = VQVAE.from_pretrained("EPFL-VILAB/4M_tokenizers_human-poses_1k_8").eval().to(device) model_tokenizer_rgb = load_model_from_config(VIDTOK_CFG_PATH, _download("vidtok_rgb")).to(device).eval().float() model_tokenizer_normal = load_model_from_config(VIDTOK_CFG_PATH, _download("vidtok_normal")).to(device).eval().float() model_tokenizer_opticalflow = load_model_from_config(VIDTOK_CFG_PATH, _download("vidtok_opticalflow")).to(device).eval().float() model_tokenizer_depth = load_model_from_config(VIDTOK_CFG_PATH, _download("vidtok_depth")).to(device).eval().float() tokenizer_model_vjepa = load_model_from_config_m(VJEPA_CFG_PATH, _download("vjepa")).to(device).eval().float() tokenizer_model_dinov2 = load_model_from_config_m(DINOV2_CFG_PATH, _download("dinov2")).to(device).eval().float() tokenizer_model_siglipv2 = load_model_from_config_m(SIGLIPV2_CFG_PATH, _download("siglipv2")).to(device).eval().float() _STATE.update( loaded=True, device=device, args=args, text_tokenizer=text_tokenizer, model=model, sampler=GenerationSampler(model), human_pose_tokenizer=human_pose_tokenizer, tokenizer_rgb=model_tokenizer_rgb, tokenizer_normal=model_tokenizer_normal, tokenizer_opticalflow=model_tokenizer_opticalflow, tokenizer_depth=model_tokenizer_depth, tokenizer_vjepa=tokenizer_model_vjepa, tokenizer_dinov2=tokenizer_model_dinov2, tokenizer_siglipv2=tokenizer_model_siglipv2, ) return _STATE def tokenize_raw_video(video_path, model_tokenizer_rgb, device="cuda"): """Load a raw video, center-crop to square, resize to 128x128, tokenize with the RGB VQ encoder. Returns: gt_rgb_tokens: flat numpy int array of VQ indices, shape [1280] original_rgb_detokenized: uint8 tensor [T H W C] of the decoded reconstruction single_crop_video_rgb: uint8 tensor [T H W C] of the resized raw frames """ vr = VideoReader(video_path, ctx=cpu()) starting_time, ending_time = 0, 4 original_video_fps = vr.get_avg_fps() selected_frame_indices = np.linspace( starting_time * original_video_fps, ending_time * original_video_fps, 17, dtype=np.int32 ) selected_frame_indices[-1] = selected_frame_indices[-1] - 1 raw_frames = torch.tensor(vr.get_batch(selected_frame_indices).asnumpy()).permute(0, 3, 1, 2) h, w = raw_frames.shape[-2:] min_dim = min(h, w) cropped_frames = torch.stack([TF.center_crop(frame, (min_dim, min_dim)) for frame in raw_frames]) resized_frames = torch.stack([TF.resize(frame, (128, 128), antialias=True) for frame in cropped_frames]) single_crop_video_rgb = resized_frames.permute(0, 2, 3, 1) # T H W C uint8 transform_norm = T.Normalize(mean=[0.5] * 3, std=[0.5] * 3) x_input_rgb = transform_norm(resized_frames.float() / 255.0).unsqueeze(0).permute(0, 2, 1, 3, 4).to(device) with torch.no_grad(), torch.autocast(device_type="cuda"): z, reg_log = model_tokenizer_rgb.encode(x_input_rgb, return_reg_log=True) latents = model_tokenizer_rgb.regularization.indices_to_codes(reg_log["indices"]) recon_original_video = model_tokenizer_rgb.decoder(latents) reshaped = rearrange(recon_original_video.squeeze(0), "c t h w -> t c h w") original_rgb_detokenized = torch.tensor(tensor_to_uint8(reshaped)).permute(0, 2, 3, 1) gt_rgb_tokens = reg_log["indices"].reshape(-1).cpu().numpy() return gt_rgb_tokens, original_rgb_detokenized, single_crop_video_rgb def _build_schedule(chain, cond_domains_user, overrides=None): """overrides: optional {modality_key: {"temp": float, "cfg": float, "decoding_steps": int}}. Only the 3 params callers are allowed to tune; anything not present in overrides[modality_key] falls back to the CONFIGS default. decoding_steps overrides are ignored for autoregressive modalities (caption/transcription/det), since that scheme doesn't use a decoding-steps schedule at all. cond_domains_user: modality keys given as full conditioning input (for both current use cases this is a single-element list, e.g. ["rgb"]). """ overrides = overrides or {} cond_domains = [CONFIGS[k][0] for k in cond_domains_user] # Conditioning modalities are given in full (see _set_full_conditioning), # so they aren't generation targets -- prepend them here only so # pop_conditioning_domain (below) can find and remove them from # target_domains, matching cvpr_fvd_videos_no_poses.py's # `not partial_conditioning` path. chain = list(dict.fromkeys(cond_domains_user + chain)) target_domains = [CONFIGS[k][0] for k in chain] autoregression_schemes = [CONFIGS[k][1] for k in chain] decoding_steps = [ overrides.get(k, {}).get("decoding_steps", CONFIGS[k][2]) if CONFIGS[k][2] is not None else None for k in chain ] token_decoding_schedules = [CONFIGS[k][3] for k in chain] tokens_per_target = [CONFIGS[k][4] for k in chain] # build_chained_generation_schedules only special-cases cfg_scale when # it's exactly a native float or list -- anything else (e.g. an int) # leaves cfg_schedule unassigned and crashes. Force native float here so # slider values (or CONFIGS defaults) can never trip that. temps = [float(overrides.get(k, {}).get("temp", CONFIGS[k][5])) for k in chain] cfg_scales = [float(overrides.get(k, {}).get("cfg", CONFIGS[k][6])) for k in chain] temp_schedules = ["constant"] * len(target_domains) cfg_schedules = ["constant"] * len(target_domains) # partial_conditioning_tokens=[None]*n + complete=False -> full conditioning: # cond_domains are fully given, so pop_conditioning_domain removes them # from target_domains entirely (nothing left to complete/generate for them). (target_domains, autoregression_schemes, decoding_steps, token_decoding_schedules, temps, temp_schedules, cfg_scales, cfg_schedules, tokens_per_target, _) = pop_conditioning_domain( cond_domains, target_domains, [None] * len(cond_domains_user), False, autoregression_schemes, decoding_steps, token_decoding_schedules, temps, temp_schedules, cfg_scales, cfg_schedules, tokens_per_target, ) schedule = build_chained_generation_schedules( cond_domains=cond_domains, target_domains=target_domains, tokens_per_target=tokens_per_target, autoregression_schemes=autoregression_schemes, decoding_steps=decoding_steps, token_decoding_schedules=token_decoding_schedules, temps=temps, temp_schedules=temp_schedules, cfg_scales=cfg_scales, cfg_schedules=cfg_schedules, cfg_grow_conditioning=True, ) return schedule, target_domains, cond_domains, tokens_per_target def _schedule_segments(schedule): """Groups a flat generation schedule (list of dicts carrying 'target_domain') into ordered contiguous segments [(domain, start, end_exclusive), ...]. Chained schedules never interleave domains in practice, but a domain reappearing later would simply open a new segment rather than corrupting the grouping. """ segments = [] for i, step in enumerate(schedule): domain = step["target_domain"] if segments and segments[-1][0] == domain and segments[-1][2] == i: segments[-1] = (domain, segments[-1][1], i + 1) else: segments.append((domain, i, i + 1)) return segments N_RECON_FRAMES_VJEPA = 17 // 2 N_RECON_FRAMES_OTHERS = 17 def _build_modality_config(state): """Fresh per-call: PCA converters fit themselves on first use, so they must not be reused across generate() calls. """ pca_converters = {"vjepa": FeatureToPCAConverter(), "dinov2": FeatureToPCAConverter(), "siglip": FeatureToPCAConverter()} return { "tok_video_rgb@128": {"tokenizer": state["tokenizer_rgb"], "frames": 5, "kind": "decode", "name": "rgb"}, "tok_video_depth@128": {"tokenizer": state["tokenizer_depth"], "frames": 5, "kind": "decode", "name": "depth"}, "tok_video_opticalflow@128": { "tokenizer": state["tokenizer_opticalflow"], "frames": 5, "kind": "decode", "name": "opticalflow", "post_process": lambda x: torch.tensor(convert_raw_optical_flow(np.array(x), bound=20)), }, "tok_video_normal@128": {"tokenizer": state["tokenizer_normal"], "frames": 5, "kind": "decode", "name": "normal"}, "tok_video_vjepa@224": {"tokenizer": state["tokenizer_vjepa"], "frames": N_RECON_FRAMES_VJEPA, "kind": "pca", "pca": pca_converters["vjepa"], "modality": "vjepa", "name": "vjepa"}, "tok_video_dinov2@224": {"tokenizer": state["tokenizer_dinov2"], "frames": N_RECON_FRAMES_OTHERS, "kind": "pca", "pca": pca_converters["dinov2"], "modality": "dinov2", "name": "dinov2"}, "tok_video_siglipv2@224": {"tokenizer": state["tokenizer_siglipv2"], "frames": N_RECON_FRAMES_OTHERS, "kind": "pca", "pca": pca_converters["siglip"], "modality": "siglipv2", "name": "siglip"}, } def _decode_visual_domain(tokens, config, device): """Decodes either a generated out_dict[domain]['tensor'] (shape [1, N]) or raw GT tokens straight from an example (shape [N], no batch dim) into a uint8 numpy video [T, H, W, 3]. """ tokens = torch.as_tensor(tokens) tokens_batched = tokens.unsqueeze(0) if tokens.dim() == 1 else tokens if config["kind"] == "pca": predicted_tokens = load_and_decode_tokens( tokens_batched.squeeze(0), config["tokenizer"], device, config["frames"], modality=config["modality"] ) predicted_pca = config["pca"].convert_to_rgb(predicted_tokens) predicted_uint8 = (predicted_pca[0].cpu().numpy() * 255).astype(np.uint8) return resize_and_duplicate(predicted_uint8) reshaped_tensor = rearrange(tokens_batched, "1 (t h w) -> 1 t h w", t=config["frames"], h=16, w=16) with torch.no_grad(), torch.autocast(device_type="cuda"): latents = config["tokenizer"].regularization.indices_to_codes(reshaped_tensor.to(device)) disc_rec = config["tokenizer"].decoder(latents) reshaped_rec = rearrange(disc_rec.squeeze(0), "c t h w -> t c h w") video = torch.tensor(tensor_to_uint8(reshaped_rec)).permute(0, 2, 3, 1) if "post_process" in config: video = config["post_process"](video) return np.array(video) def _decode_all_outputs(out_dict, target_domains, state, device, modality_config): """Decodes every generated target domain into the results dict: numpy uint8 videos for visual modalities, strings for caption/transcription, a bbox-rendered video for det. """ results = {} for domain in target_domains: if domain not in out_dict or domain not in modality_config: continue config = modality_config[domain] results[config["name"]] = _decode_visual_domain(out_dict[domain]["tensor"], config, device) dec_dict = decode_dict_text_modalities(out_dict, state["text_tokenizer"], state["human_pose_tokenizer"]) if "caption" in target_domains: caption = dec_dict.get("caption") results["caption"] = caption[0] if isinstance(caption, list) else caption if "transcription" in target_domains: transcription = dec_dict.get("transcription") results["transcription"] = transcription[0] if isinstance(transcription, list) else transcription if "det" in target_domains and "det" in dec_dict: results["det"] = np.concatenate(visualize_temporal_bboxes(None, dec_dict["det"][0]), axis=0) return results def generate(video_path, chain=None, seed=0, top_p=0.8, top_k=0.0, overrides=None): """Runs the full RGB -> {chain} chained generation for one uploaded video. overrides: optional per-modality {"temp": float, "cfg": float, "decoding_steps": int} dict, see _build_schedule for details. Pass None to use CONFIGS defaults. Returns a dict: {modality_name: numpy uint8 video [T,H,W,3]} for visual modalities, plus {"caption": str, "transcription": str} for text modalities when they're in the chain. """ chain = chain or DEFAULT_CHAIN state = load_pipeline() device = state["device"] gt_rgb_tokens, original_rgb_detokenized, single_crop_video_rgb = tokenize_raw_video( video_path, state["tokenizer_rgb"], device=device ) schedule, target_domains, cond_domains, tokens_per_target = _build_schedule(chain, ["rgb"], overrides=overrides) batched_sample = {} for target_mod, ntoks in zip(target_domains, tokens_per_target): batched_sample = init_empty_target_modality( batched_sample, MODALITY_INFO, target_mod, 1, ntoks, device, False ) batched_sample["tok_video_rgb@128"] = image_mask_first_frame_conditional( gt_rgb_tokens, 1280, RGB_SEED_TOKENS, RGB_TARGET_TOKENS ) modality_config = _build_modality_config(state) with torch.no_grad(): out_dict = state["sampler"].generate( batched_sample, schedule, text_tokenizer=state["text_tokenizer"], verbose=True, seed=seed, top_p=top_p, top_k=top_k, prediction_window_size=13, perform_windowed_prediction=False, use_decomposed_inference=False, decomposed_inference_modalities=[], ) results = {"rgb_input": np.array(single_crop_video_rgb)} results.update(_decode_all_outputs(out_dict, target_domains, state, device, modality_config)) return results def _prepare_any_to_any(example_stem, input_modality, chain, overrides, raw_video_path, raw_caption_text): """Shared setup for generate_any_to_any / generate_any_to_any_stream: resolves the input (curated example, raw uploaded video, or typed caption), builds the chained generation schedule, initializes the empty target modalities and sets the full conditioning. """ chain = chain or [k for k in CONFIGS if k != input_modality] state = load_pipeline() device = state["device"] if raw_video_path is not None: assert input_modality == "rgb", "raw_video_path input is only supported for input_modality='rgb'" gt_rgb_tokens, _, _ = tokenize_raw_video(raw_video_path, state["tokenizer_rgb"], device=device) example = {"rgb": gt_rgb_tokens} elif raw_caption_text is not None: assert input_modality == "caption", "raw_caption_text input is only supported for input_modality='caption'" example = {"caption": raw_caption_text} else: example = _load_example_raw(example_stem) schedule, target_domains, cond_domains, tokens_per_target = _build_schedule(chain, [input_modality], overrides=overrides) batched_sample = {} for target_mod, ntoks in zip(target_domains, tokens_per_target): batched_sample = init_empty_target_modality( batched_sample, MODALITY_INFO, target_mod, 1, ntoks, device, False ) _set_full_conditioning(batched_sample, input_modality, example, state["text_tokenizer"], device) modality_config = _build_modality_config(state) return state, device, example, schedule, target_domains, batched_sample, modality_config def _decode_input_preview(input_modality, example, modality_config, device): """Decodes what was actually given as input into a displayable value (string for caption/transcription, uint8 numpy video otherwise), so the UI can show it alongside the predictions. Deterministic and RNG-free, so it is safe to run either before or after sampling. """ if input_modality == "caption": return example["caption"] if input_modality == "transcription": return " ".join(example["transcription"]) if input_modality == "det": sequence = _prepare_detections_text(example["detection_dict"], example["crop_settings"]) return np.concatenate(visualize_temporal_bboxes(None, sequence), axis=0) domain = CONFIGS[input_modality][0] return _decode_visual_domain(example[input_modality], modality_config[domain], device) def _decode_single_output(domain, out_dict, state, device, modality_config): """Single-domain counterpart of _decode_all_outputs: decodes one generated domain into (results_key, value), producing exactly the entry _decode_all_outputs would for that domain. decode_dict_text_modalities only reads the entries of the dict it is handed, so passing the one completed domain mid-chain is equivalent to decoding it at the end. """ if domain in modality_config: config = modality_config[domain] return config["name"], _decode_visual_domain(out_dict[domain]["tensor"], config, device) dec_dict = decode_dict_text_modalities({domain: out_dict[domain]}, state["text_tokenizer"], state["human_pose_tokenizer"]) if domain == "det": return "det", np.concatenate(visualize_temporal_bboxes(None, dec_dict["det"][0]), axis=0) decoded = dec_dict.get(domain) return domain, decoded[0] if isinstance(decoded, list) else decoded def generate_any_to_any( example_stem=None, input_modality="rgb", chain=None, seed=0, top_p=0.8, top_k=0.0, overrides=None, raw_video_path=None, raw_caption_text=None, ): """Runs input_modality -> {chain} chained generation for one input, which is either a curated example (example_stem), a user-uploaded raw video (raw_video_path, must be paired with input_modality="rgb"), or a user-typed caption (raw_caption_text, must be paired with input_modality="caption"). Exactly one of the three should be given. chain: which modalities to generate; defaults to every CONFIGS modality except input_modality (i.e. true any-to-all). overrides: see _build_schedule. Returns a dict: {modality_name: numpy uint8 video} for visual modalities (including f"input_{input_modality}", the decoded/displayed conditioning input), plus {"caption": str, "transcription": str} for text modalities when they're in the chain, plus f"input_{input_modality}" as a string when the input modality itself is caption/transcription. """ state, device, example, schedule, target_domains, batched_sample, modality_config = _prepare_any_to_any( example_stem, input_modality, chain, overrides, raw_video_path, raw_caption_text ) with torch.no_grad(): out_dict = state["sampler"].generate( batched_sample, schedule, text_tokenizer=state["text_tokenizer"], verbose=True, seed=seed, top_p=top_p, top_k=top_k, prediction_window_size=13, perform_windowed_prediction=False, use_decomposed_inference=False, decomposed_inference_modalities=[], ) results = _decode_all_outputs(out_dict, target_domains, state, device, modality_config) # Also surface what was actually given as input, so the UI can show it # alongside the predictions. results[f"input_{input_modality}"] = _decode_input_preview(input_modality, example, modality_config, device) return results def generate_any_to_any_stream( example_stem=None, input_modality="rgb", chain=None, seed=0, top_p=0.8, top_k=0.0, overrides=None, raw_video_path=None, raw_caption_text=None, progress_every=5, ): """Streaming counterpart of generate_any_to_any: same inputs, same schedule and per-step seeding (sampler.generate_iter runs the identical step functions as sampler.generate, so outputs match bit-for-bit for the same arguments), but yields (kind, key, payload) 3-tuples as generation progresses instead of returning one dict at the end: ("input", input_modality, preview) once, before sampling starts ("start", modality_key, total_steps) that modality's segment begins ("progress", modality_key, (done, total)) every `progress_every` steps; text modalities are a single autoregressive step and emit none ("result", modality_key, decoded_value) as soon as a modality finishes ("done", None, results_dict) same dict generate_any_to_any returns, incl. f"input_{...}" generate_iter is decorated @torch.no_grad() itself; deliberately NOT wrapped in an outer no_grad here, because a context manager held open across a yield would leak no-grad state into the caller while suspended. """ state, device, example, schedule, target_domains, batched_sample, modality_config = _prepare_any_to_any( example_stem, input_modality, chain, overrides, raw_video_path, raw_caption_text ) input_preview = _decode_input_preview(input_modality, example, modality_config, device) yield ("input", input_modality, input_preview) segments = _schedule_segments(schedule) results = {} seg_idx = 0 seg_domain, seg_start, seg_end = segments[0] yield ("start", _DOMAIN_TO_KEY[seg_domain], seg_end - seg_start) iterator = state["sampler"].generate_iter( batched_sample, schedule, text_tokenizer=state["text_tokenizer"], verbose=False, seed=seed, top_p=top_p, top_k=top_k, prediction_window_size=13, perform_windowed_prediction=False, use_decomposed_inference=False, decomposed_inference_modalities=[], ) for step_idx, mod_dict in enumerate(iterator): if step_idx + 1 == seg_end: # The segment's last decoding step just ran: this modality's # tokens are complete, decode only it. mod_dict is the sampler's # live dict (mutated in place between yields), so all decoding # must happen before the iterator is resumed -- which it does, # since this generator only advances when the caller pulls the # next event. results_key, value = _decode_single_output(seg_domain, mod_dict, state, device, modality_config) results[results_key] = value yield ("result", _DOMAIN_TO_KEY[seg_domain], value) seg_idx += 1 if seg_idx < len(segments): seg_domain, seg_start, seg_end = segments[seg_idx] yield ("start", _DOMAIN_TO_KEY[seg_domain], seg_end - seg_start) else: done = step_idx + 1 - seg_start if progress_every and done % progress_every == 0: yield ("progress", _DOMAIN_TO_KEY[seg_domain], (done, seg_end - seg_start)) results[f"input_{input_modality}"] = input_preview yield ("done", None, results) # --- Future frame prediction -------------------------------------------------- # Ported from cvpr_fvd_videos_no_poses_raw_rgb.py / future_pred.py: rather # than giving a modality's tokens in full (as in any-to-any above), only the # first `seed_tokens` tokens are given -- the model completes the rest of # that same modality's video (temporal extrapolation) alongside every other # modality in the chain, all steered by cfg_grow_conditioning=True. FUTURE_CHAIN = ["caption", "transcription", "siglip", "vjepa", "dinov2", "det", "depth", "normal", "opticalflow", "rgb"] # These four share VidTok's causal frame layout (1280 tokens = 5 causal # blocks -- block 0 is 1 frame, each subsequent block is 4 frames), which is # what makes the 256/512-seed-token "first frame" / "first 5 frames" framing # correct for them. siglip/dinov2/vjepa use a different per-frame tokenizer # layout and were never partial-conditioned this way in the research scripts, # so they're excluded. 'det' is also supported, via a completely different # mechanism (see _build_det_partial_conditioning below) since it's a # variable-length text sequence, not a fixed spatial token grid. FUTURE_SEED_MODALITIES = ["rgb", "depth", "normal", "opticalflow", "det"] # seed_tokens -> raw frames actually given: 256 tokens = 1 frame (block 0 # alone), 512 tokens = 5 frames (block 0 + one 4-frame block). Matches # cvpr_fvd_videos_no_poses_raw_rgb.py's RGB_SEED_TOKENS convention exactly. # Also reused as a frame count for 'det' (see SEED_TOKENS_TO_FRAMES). FUTURE_SEED_TOKEN_OPTIONS = {"First frame only": 256, "First 5 frames": 512} SEED_TOKENS_TO_FRAMES = {256: 1, 512: 5} # Optional extra full-conditioning input to steer the predicted trajectory # (e.g. a caption). Kept to caption/transcription only -- unlike picking an # existing example's detections as a *seed* (above), asking a user to author # new detection annotations from scratch for the optional extra-conditioning # slot isn't practical. FUTURE_EXTRA_COND_MODALITIES = ["caption", "transcription"] # --- det partial (future-prediction seed) conditioning ----------------------- # Ported from cvpr_fvd_videos_no_poses.py's partial-'det' branch. Unlike # image_mask_first_frame_conditional's fixed-token-count prefix (used for # rgb/depth/normal/opticalflow above), det's seed is a *frame count*: find # where the (n_frames+1)-th frame's sentinel token starts in the tokenized # sequence, and give everything before that as input context. # # [S_i] sentinel tokens are a separate, simpler text-level scheme from the # [FRAME_i] tokens used for full det conditioning elsewhere (_prepare_detections_text) -- # transform_tensor_with_markers below remaps them into the model's actual # frame-marker special-token id range (30004+, same range create_frame_ids # expects), so there's no real vocab mismatch, just a two-step encoding. DET_SENTINEL_TOKENS = [f"[S_{i}]" for i in range(1, 22)] DET_S1_VOCAB_ID = 5 # [S_1]'s fixed vocab id in the shared text tokenizer # frame count n -> vocab id of the sentinel marking the START of frame n+1 # (the boundary to slice at for "first n frames given"). DET_SENTINEL_VOCAB_MAPPING = {i: i + 5 for i in range(1, 18)} def _prepare_detections_text_sentinel(detection_dict, crop_settings): """Same detection preprocessing as _prepare_detections_text, but tagged with [S_i] sentinel tokens instead of [FRAME_i] -- only used for det partial (future-prediction seed) conditioning. """ detection_transforms = VideoDetectionTransform( det_threshold=0.2, det_max_instances=None, bbox_order="dist_to_orig", coord_bins=1000, min_visibility=0.0 ) starting_time, ending_time, i, j, h, w, h_flip = crop_settings[0] processed_detections = detection_transforms.image_augment( detection_dict, (starting_time, ending_time, i, j, h, w), False, None, (None, None), None, None ) gt_processed_detections, per_frame_instance_counts = detection_transforms.postprocess(processed_detections) return merge_detection_tokens_with_sentinel_tokens(gt_processed_detections, DET_SENTINEL_TOKENS, per_frame_instance_counts) def _build_det_partial_conditioning(example, n_frames, text_tokenizer, device): """Given only the first n_frames frames' worth of ground-truth detections, builds det's batched_sample entry so the model completes the rest (up to DETECTION_SIZE tokens) -- 1:1 port of cvpr_fvd_videos_no_poses.py's partial-'det' branch. """ sequence = _prepare_detections_text_sentinel(example["detection_dict"], example["crop_settings"]) seq_ids = text_tokenizer.encode(sequence).ids index_gt_sequence = seq_ids.index(DET_SENTINEL_VOCAB_MAPPING[n_frames]) partial_gt_token_tensor = torch.tensor(seq_ids[0:index_gt_sequence]).to(device).unsqueeze(0) max_length = (DETECTION_SIZE + 1) * 2 tensor = torch.zeros(max_length, dtype=torch.int) target_mask = torch.ones(max_length, dtype=torch.bool) input_mask = torch.ones(max_length, dtype=torch.bool) decoder_attention_mask = torch.ones(max_length, dtype=torch.int) input_seq_ids = transform_tensor_with_markers(partial_gt_token_tensor).squeeze(0) tensor[:len(input_seq_ids)] = input_seq_ids.to(dtype=torch.int) input_mask[:len(input_seq_ids)] = 0 tensor[len(input_seq_ids)] = DET_S1_VOCAB_ID target_mask[len(input_seq_ids):len(input_seq_ids) + DETECTION_SIZE] = 0 # Last target position must be the "stop" sentinel so the model knows # where to end generation. tensor[len(input_seq_ids) + DETECTION_SIZE - 1] = 22 - n_frames decoder_attention_mask[len(input_seq_ids):len(input_seq_ids) + DETECTION_SIZE] = 0 return { "tensor": tensor.unsqueeze(0).to(device), "input_mask": input_mask.unsqueeze(0).to(device), "target_mask": target_mask.unsqueeze(0).to(device), "decoder_attention_mask": decoder_attention_mask.unsqueeze(0).to(device), } def _build_future_schedule(chain, seed_modality, seed_tokens, extra_cond_modality=None, overrides=None): """chain: the fixed FUTURE_CHAIN, filtered/ordered by the UI's checkboxes -- must already include seed_modality (and extra_cond_modality, if any) at their normal fixed positions; this function does not reorder it (unlike _build_schedule's any-to-any prepending, which doesn't apply here since rgb's fixed *last* position -- conditioning on everything else -- is what the original future-prediction chain design relies on). overrides: see _build_schedule. """ overrides = overrides or {} assert seed_modality in chain, "seed_modality must be present in chain (its checkbox is always forced on)" if extra_cond_modality: assert extra_cond_modality in chain, "extra_cond_modality must be present in chain (its checkbox is always forced on)" cond_domains_user = [seed_modality] + ([extra_cond_modality] if extra_cond_modality else []) cond_domains = [CONFIGS[k][0] for k in cond_domains_user] target_domains = [CONFIGS[k][0] for k in chain] autoregression_schemes = [CONFIGS[k][1] for k in chain] decoding_steps = [ overrides.get(k, {}).get("decoding_steps", CONFIGS[k][2]) if CONFIGS[k][2] is not None else None for k in chain ] token_decoding_schedules = [CONFIGS[k][3] for k in chain] tokens_per_target = [CONFIGS[k][4] for k in chain] temps = [float(overrides.get(k, {}).get("temp", CONFIGS[k][5])) for k in chain] cfg_scales = [float(overrides.get(k, {}).get("cfg", CONFIGS[k][6])) for k in chain] temp_schedules = ["constant"] * len(target_domains) cfg_schedules = ["constant"] * len(target_domains) # seed_modality is partial (partial_tokens != None) -> pop_conditioning_domain # keeps it in target_domains, to be completed. extra_cond_modality (if # any) is fully given (partial_tokens=None) -> popped, nothing to complete. partial_conditioning_tokens = [seed_tokens] + ([None] if extra_cond_modality else []) (target_domains, autoregression_schemes, decoding_steps, token_decoding_schedules, temps, temp_schedules, cfg_scales, cfg_schedules, tokens_per_target, _) = pop_conditioning_domain( cond_domains, target_domains, partial_conditioning_tokens, True, autoregression_schemes, decoding_steps, token_decoding_schedules, temps, temp_schedules, cfg_scales, cfg_schedules, tokens_per_target, ) # seed_modality's target token budget is the completion budget (GT minus # the seed), not its full per-modality budget -- except 'det', whose # CONFIGS budget (DETECTION_SIZE) is already exactly the completion span # regardless of how many frames were given (the variable-length prefix # sits before it, not carved out of a shared token pool). if seed_modality != "det": seed_domain = CONFIGS[seed_modality][0] seed_gt_tokens = CONFIGS[seed_modality][4] tokens_per_target[target_domains.index(seed_domain)] = seed_gt_tokens - seed_tokens schedule = build_chained_generation_schedules( cond_domains=cond_domains, target_domains=target_domains, tokens_per_target=tokens_per_target, autoregression_schemes=autoregression_schemes, decoding_steps=decoding_steps, token_decoding_schedules=token_decoding_schedules, temps=temps, temp_schedules=temp_schedules, cfg_scales=cfg_scales, cfg_schedules=cfg_schedules, cfg_grow_conditioning=True, ) return schedule, target_domains, cond_domains, tokens_per_target def generate_future_prediction( example_stem, seed_modality, seed_tokens, chain, extra_cond_modality=None, seed=0, top_p=0.8, top_k=0.0, overrides=None, override_caption_text=None, override_transcription_text=None, ): """Given only the first seed_tokens tokens (256 = first frame, 512 = first 5 frames; for 'det' this maps to a frame count via SEED_TOKENS_TO_FRAMES instead) of seed_modality's ground truth, predicts the rest of that same modality (temporal extrapolation) plus every other modality checked in chain, optionally also fully conditioned on one extra modality (caption/transcription) to steer the predicted trajectory. seed_modality must be one of FUTURE_SEED_MODALITIES. chain must already include seed_modality and extra_cond_modality (see _build_future_schedule). overrides: see _build_schedule. override_caption_text/override_transcription_text: see _set_full_conditioning -- lets the Future Prediction tab's editable extra-conditioning box steer generation instead of the example's own text. Returns a dict: {modality_name: video/str} for every domain in chain -- seed_modality's own entry already contains the given seed frames plus the predicted future frames combined, nothing separate to surface for it -- plus f"input_{extra_cond_modality}" (video or str) when extra_cond_modality is given, same convention as generate_any_to_any's f"input_{input_modality}". """ assert seed_modality in FUTURE_SEED_MODALITIES state = load_pipeline() device = state["device"] example = _load_example_raw(example_stem) schedule, target_domains, cond_domains, tokens_per_target = _build_future_schedule( chain, seed_modality, seed_tokens, extra_cond_modality, overrides=overrides ) batched_sample = {} for target_mod, ntoks in zip(target_domains, tokens_per_target): batched_sample = init_empty_target_modality( batched_sample, MODALITY_INFO, target_mod, 1, ntoks, device, False ) if seed_modality == "det": batched_sample["det"] = _build_det_partial_conditioning( example, SEED_TOKENS_TO_FRAMES[seed_tokens], state["text_tokenizer"], device ) else: seed_domain = CONFIGS[seed_modality][0] seed_gt_tokens = CONFIGS[seed_modality][4] batched_sample[seed_domain] = image_mask_first_frame_conditional( example[seed_modality], seed_gt_tokens, seed_tokens, seed_gt_tokens - seed_tokens ) if extra_cond_modality: _set_full_conditioning( batched_sample, extra_cond_modality, example, state["text_tokenizer"], device, override_caption_text=override_caption_text, override_transcription_text=override_transcription_text, ) modality_config = _build_modality_config(state) with torch.no_grad(): out_dict = state["sampler"].generate( batched_sample, schedule, text_tokenizer=state["text_tokenizer"], verbose=True, seed=seed, top_p=top_p, top_k=top_k, prediction_window_size=13, perform_windowed_prediction=False, use_decomposed_inference=False, decomposed_inference_modalities=[], ) results = _decode_all_outputs(out_dict, target_domains, state, device, modality_config) if extra_cond_modality: input_key = f"input_{extra_cond_modality}" if extra_cond_modality == "caption": # Surface whatever was actually fed to the model -- the user's # edited text when given (see _set_full_conditioning above), # not always the example's own ground-truth caption. results[input_key] = override_caption_text if override_caption_text is not None else example["caption"] elif extra_cond_modality == "transcription": results[input_key] = ( override_transcription_text if override_transcription_text is not None else " ".join(example["transcription"]) ) elif extra_cond_modality == "det": sequence = _prepare_detections_text(example["detection_dict"], example["crop_settings"]) results[input_key] = np.concatenate(visualize_temporal_bboxes(None, sequence), axis=0) else: domain = CONFIGS[extra_cond_modality][0] results[input_key] = _decode_visual_domain(example[extra_cond_modality], modality_config[domain], device) return results