"""Precomputed-results cache: chain policy, request canonicalization, lookup. GPU time on the HF Space is rationed (ZeroGPU gives free-tier visitors a few minutes a day, and load_pipeline's ~37.5 GB of checkpoints is loaded from inside the @spaces.GPU allocation). So the default configuration of each tab is generated ahead of time on the cluster and served from disk: the page shows a real model output the moment it opens, and clicking Generate without changing anything costs nothing. Two halves, and they have to agree exactly: - scripts/precompute_demo_outputs.py calls resolve_a2a / resolve_fp to build a *canonical spec* dict, runs the model, and writes both to the manifest. - app.py calls the same resolvers on the live UI arguments and looks up the spec by plain dict equality. Everything that can change a generated result is in the spec; nothing that can't is (see _effective_overrides). That is what makes a fresh page load compare equal to a run recorded weeks earlier -- and, just as importantly, what makes a nudged slider miss instead of silently serving the wrong video. This module deliberately does NOT import gradio or app: the precompute script has to reach the chain-policy helpers without building ~1000 UI components. """ import hashlib import json import os import traceback from collections import namedtuple from inference import ( CONFIGS, HYPERPARAM_PRESETS, FUTURE_CHAIN, get_precomputed_file, load_precomputed_manifest, ) # Bump when the manifest layout changes incompatibly; a mismatched manifest is # ignored outright rather than half-understood. SCHEMA_VERSION = 1 # --- modality registry (shared with the UI) ----------------------------------- # Fixed order for the hyperparameter accordion only (unrelated to the # dynamic, chain-order-following output slots). ALL_MODALITIES = ["rgb", "depth", "normal", "opticalflow", "dinov2", "vjepa", "siglip", "det", "caption", "transcription"] TEXT_MODALITIES = {"caption", "transcription"} # Which tuned HYPERPARAM_PRESETS set fits a given input modality -- dense # visual/feature-map inputs carry a lot of information already, sparse # text-like inputs need the more exploratory sparse-to-dense tuning. DENSE_INPUT_MODALITIES = {"rgb", "depth", "normal", "opticalflow", "dinov2", "vjepa", "siglip"} # Two recommended coarse-to-fine orders -- one for dense (pixel/feature-map) # inputs, one for sparse (text-like) inputs -- each modality drops its own # key from whichever list it belongs to. Previously each input modality had # its own hand-tuned order; simplified to these two shared orders. Edit # freely; the only rule is that a list must not contain its own key # (enforced below, not hand-maintained per entry anymore). DENSE_CHAIN = ["siglip", "dinov2", "vjepa", "caption", "transcription", "det", "depth", "normal", "opticalflow", "rgb"] SPARSE_CHAIN = ["caption", "transcription", "siglip", "dinov2", "vjepa", "det", "depth", "normal", "opticalflow", "rgb"] SPARSE_INPUT_MODALITIES = ("caption", "transcription", "det") # Both source lists contain all 10 modalities (each is a target list for the # other group's inputs too), so the dict keys must be restricted to each # group's actual INPUT modalities -- iterating the full source lists here # would let the second comprehension silently overwrite the first for every # shared key. COARSE_TO_FINE_BY_INPUT = { **{mod: [m for m in DENSE_CHAIN if m != mod] for mod in DENSE_INPUT_MODALITIES}, **{mod: [m for m in SPARSE_CHAIN if m != mod] for mod in SPARSE_INPUT_MODALITIES}, } DEFAULT_INPUT_MODALITY = "rgb" DEFAULT_ABS_MODALITY = "depth" DEFAULT_CHAIN_STATE = list(COARSE_TO_FINE_BY_INPUT[DEFAULT_INPUT_MODALITY]) # Future Prediction tab defaults. Shared by the UI (fp_targets_state), the # page-load prefill and the precompute recipe: the prefill re-derives its # lookup from the tab's actual default state, so if these three ever # disagreed the cache would simply never be found. DEFAULT_FP_OBS_MODALITY = DEFAULT_INPUT_MODALITY DEFAULT_FP_SEED_TOKENS = 512 # "First 5 frames" # Predict every modality we can, so the showcase shows the model's full # any-modality reach rather than a two-card sample. The observed modality is # excluded (its given frames are completed in place, not predicted). DEFAULT_FP_TARGETS = [m for m in FUTURE_CHAIN if m != DEFAULT_FP_OBS_MODALITY] # A hand-picked caption prompt shown as the first quick-pick chip in the # Caption input mode -- unlike the rest of CAPTION_SAMPLES (real GT captions # of curated gallery clips), this one has no backing example_stem, so it must # be listed here explicitly for scripts/precompute_demo_outputs.py to also # precompute it. Single source of truth: app.py imports this rather than # hardcoding its own copy, so the two can't drift apart. EXTRA_CAPTION_SAMPLE = ( "A young woman with light skin and reddish-blonde hair in pigtails is facing the camera. " "She wears a green and blue tie-dye top and is in a room with posters on the wall. " "She appears to be in the middle of a burp, exclaiming 'Oh'." ) def hyperparam_preset_key_for_modality(input_modality): return "rgb_to_others" if input_modality in DENSE_INPUT_MODALITIES else "text_to_rgb" # Flat temp/cfg overrides for each tab's "Direct" preset -- deliberately a # single fixed pair per modality family rather than a tuned HYPERPARAM_PRESETS # set, since skipping the coarse-to-fine chain calls for simpler generation # hyperparameters too. DIRECT_HYPERPARAMS_A2A = { **{m: {"temp": 2.0, "cfg": 4.0} for m in DENSE_INPUT_MODALITIES}, **{m: {"temp": 1.0, "cfg": 1.0} for m in TEXT_MODALITIES | {"det"}}, } DIRECT_HYPERPARAMS_FP = {m: {"temp": 1.0, "cfg": 1.0} for m in ALL_MODALITIES} # --- chain policy -------------------------------------------------------------- def clean_chain(input_modality, chain): """Drops the input modality from chain if it's stuck there from before a modality switch (rather than wiping the whole chain), and dedupes. """ seen = set() cleaned = [] for k in chain: if k != input_modality and k not in seen: cleaned.append(k) seen.add(k) return cleaned def preset_chain(preset, input_modality): """The chain a preset stands for, given the current input modality. Returns None for "custom" (leave the chain exactly as the user built it). """ if preset == "coarse": fallback = ["caption", "dinov2", "depth", "rgb"] return clean_chain(input_modality, list(COARSE_TO_FINE_BY_INPUT.get(input_modality, fallback))) if preset == "direct": return ["depth"] if input_modality == "rgb" else ["rgb"] return None def chain_for_input_change(preset, new_input_modality, current_chain): """What the chain becomes when the input modality changes: coarse rebuilds its per-input ladder, direct keeps its single chosen target while it stays valid (falling back to the direct default if the new input swallowed it), custom just drops the new input from the chain. """ if preset == "coarse": return preset_chain("coarse", new_input_modality) cleaned = clean_chain(new_input_modality, current_chain) if preset == "direct": return cleaned[:1] if cleaned else preset_chain("direct", new_input_modality) return cleaned def fp_extra(cond_mode, cond_modality): """Effective conditioning modality: None when unconditional.""" return None if cond_mode == "none" else cond_modality def fp_chain(seed_modality, extra, targets): """The chain generate_future_prediction runs: always the fixed, tuned FUTURE_CHAIN order filtered down to what's needed -- the user picks WHICH modalities to predict, never the order. The seed modality (its given frames are completed in place) and the conditioning modality (fully given; the schedule builder pops it) are structurally required in-chain. """ keep = set(targets) | {seed_modality} | ({extra} if extra else set()) return [m for m in FUTURE_CHAIN if m in keep] def fp_preset_targets(preset, obs_modality): """(targets, complete_seed) for a Future Prediction preset button -- the tab 2 analogue of preset_chain. 'chained' is today's only behavior: predict every other modality, chained through the fixed FUTURE_CHAIN order. 'direct' mirrors the A2A tab's single-hop "Direct" preset: land on rgb without chaining through any intermediate abstract modality, or, when rgb is itself what's being observed, just continue rgb's own remaining frames (there's nothing more "direct" than that). """ if preset == "direct": return ([], True) if obs_modality == "rgb" else (["rgb"], False) return ([m for m in FUTURE_CHAIN if m != obs_modality], True) # --- canonicalization ---------------------------------------------------------- # Mirrors the flat spec _build_hyperparam_controls emits, in the same order. HYPERPARAM_SPEC = [(key, param) for key in ALL_MODALITIES for param in ("temp", "cfg")] _CONFIG_IDX = {"temp": 5, "cfg": 6} def default_slider_values(preset_key): """Exactly the values a freshly-loaded hyperparameter accordion holds, in HYPERPARAM_SPEC order. preset_key=None sources straight from CONFIGS (what the Future Prediction tab does). _build_hyperparam_controls builds its sliders from this, so UI defaults and the canonicalizer cannot drift. """ hp_defaults = HYPERPARAM_PRESETS[preset_key] if preset_key else {} return [ hp_defaults.get(key, {}).get(param, CONFIGS[key][_CONFIG_IDX[param]]) for key, param in HYPERPARAM_SPEC ] def _num(value): """Fixed-width float text. A step=0.01 range input can hand us 0.7000000000000001; comparing those as floats would miss forever.""" return "%.4f" % float(value) def _effective_overrides(chain, overrides): """Only the hyperparameters that can actually change the output. _build_schedule reads overrides[k] for k in chain and nothing else, and pop_conditioning_domain then drops the input modality's entry entirely -- so the sliders for off-chain modalities, and the input modality's own sliders, are inert. Including them would make the cache miss whenever a visitor idly dragged an irrelevant slider. """ overrides = overrides or {} effective = {} for key in chain: given = overrides.get(key) or {} entry = { "temp": _num(given.get("temp", CONFIGS[key][5])), "cfg": _num(given.get("cfg", CONFIGS[key][6])), } # decoding_steps is ignored for the autoregressive modalities # (caption/transcription/det), so it must not enter their key. if CONFIGS[key][2] is not None: entry["decoding_steps"] = int(given.get("decoding_steps", CONFIGS[key][2])) effective[key] = entry return effective def canonical_a2a_spec(input_modality, source, chain, seed, top_p, top_k, overrides): """source: {"kind": "example", "stem": ...} or {"kind": "caption", "text": ...}.""" return { "v": SCHEMA_VERSION, "task": "a2a", "input_modality": input_modality, "source": source, "chain": list(chain), # ordered: the chain order changes the result "seed": int(seed), "top_p": _num(top_p), "top_k": int(round(float(top_k))), "overrides": _effective_overrides(chain, overrides), } def canonical_fp_spec(stem, obs_mod, seed_tokens, extra, steer_text, chain, seed, top_p, top_k, overrides, complete_seed=True): # When not completing the seed modality, its hyperparameter sliders are # inert (same reasoning as _effective_overrides' docstring for off-chain # modalities) -- excluded here so idly dragging them can't miss the cache. overrides_chain = chain if complete_seed else [k for k in chain if k != obs_mod] spec = { "v": SCHEMA_VERSION, "task": "fp", "stem": stem, "obs": obs_mod, "seed_tokens": int(seed_tokens), "extra": extra, "chain": list(chain), # already canonical: fp_chain filters FUTURE_CHAIN "complete_seed": bool(complete_seed), "seed": int(seed), "top_p": _num(top_p), "top_k": int(round(float(top_k))), "overrides": _effective_overrides(overrides_chain, overrides), } if extra: # steer_text only reaches the model when a conditioning modality is on. spec["steer_text"] = steer_text or "" return spec def code_fingerprint(): """Fingerprint of everything a spec is computed from. A mismatch means the cache will simply miss -- worth a startup warning, never a hard failure.""" blob = json.dumps( {"configs": CONFIGS, "presets": HYPERPARAM_PRESETS, "future_chain": FUTURE_CHAIN, "coarse": COARSE_TO_FINE_BY_INPUT}, sort_keys=True, separators=(",", ":"), default=list, ) return hashlib.sha256(blob.encode("utf-8")).hexdigest()[:16] # --- UI arguments -> generation call (one implementation, two callers) --------- A2ARequest = namedtuple("A2ARequest", "input_modality chain gen_kwargs overrides seed top_p top_k spec") FPRequest = namedtuple("FPRequest", "stem obs_mod seed_tokens extra chain steer_text overrides seed top_p top_k spec complete_seed raw_video_path") def resolve_a2a(src_mode, example_stem, upload_path, caption_text, abs_mod, chain, seed, top_p, top_k, hyperparam_values, hyperparam_spec=HYPERPARAM_SPEC): """Validates and normalizes the any-to-any tab's arguments. Raises ValueError for anything the user has to fix; app.py re-raises those as gr.Error *before* requesting a GPU, so a mis-click costs no quota. spec is None for uploads -- they have no stable identity, so they are never cacheable and always run live. """ # Pin down what's really being generated from the source mode rather than # trusting any UI state that could lag. if src_mode == "caption": input_modality = "caption" elif src_mode == "abstract": input_modality = abs_mod else: input_modality = "rgb" chain = clean_chain(input_modality, chain) if not chain: raise ValueError("Please add at least one modality to the chain first.") gen_kwargs = {} source = None if src_mode == "video" and upload_path: gen_kwargs["raw_video_path"] = upload_path elif src_mode == "video": if not example_stem: raise ValueError("Please pick an example clip or upload a video first.") gen_kwargs["example_stem"] = example_stem source = {"kind": "example", "stem": example_stem} elif src_mode == "caption": text = (caption_text or "").strip() if not text: raise ValueError("Please type a caption first.") gen_kwargs["raw_caption_text"] = text source = {"kind": "caption", "text": text} else: if not example_stem: raise ValueError("Please pick an example clip first.") gen_kwargs["example_stem"] = example_stem source = {"kind": "example", "stem": example_stem} # Decoding steps aren't exposed as sliders -- seed overrides from the tuned # preset matching the input modality, then layer the user-adjustable # temp/cfg slider values on top. preset_key = hyperparam_preset_key_for_modality(input_modality) overrides = {key: dict(vals) for key, vals in HYPERPARAM_PRESETS[preset_key].items()} for (key, param), value in zip(hyperparam_spec, hyperparam_values): overrides.setdefault(key, {})[param] = value seed, top_p, top_k = int(seed), float(top_p), float(top_k) spec = None if source is not None: spec = canonical_a2a_spec(input_modality, source, chain, seed, top_p, top_k, overrides) return A2ARequest(input_modality, chain, gen_kwargs, overrides, seed, top_p, top_k, spec) def resolve_fp(example_stem, obs_mod, seed_tokens, cond_mode, cond_mod, steer_text, targets, seed, top_p, top_k, hyperparam_values, hyperparam_spec=HYPERPARAM_SPEC, complete_seed_modality=True, raw_video_path=None): """Validates and normalizes the future-prediction tab's arguments. complete_seed_modality: mirrors run_generation.py's --should_complete_partial_conditioned_modalities -- see inference.py's _build_future_schedule for what this actually changes. raw_video_path: an uploaded clip instead of a curated example -- same convention as resolve_a2a's raw_video_path (only valid for RGB observation, has no stable identity so it's never cacheable: spec is None below, same as an any-to-any upload). """ if raw_video_path: if obs_mod != "rgb": raise ValueError("An uploaded video can only be used for RGB observation.") elif not example_stem: raise ValueError("Please pick an example clip or upload a video first.") extra = fp_extra(cond_mode, cond_mod) if raw_video_path and extra: # No ground-truth caption/transcription exists for an uploaded clip # (unlike a curated example) -- the steering text has to be typed in. if not (steer_text or "").strip(): raise ValueError( f"An uploaded video has no ground-truth {extra} -- please type one to condition on." ) targets = [t for t in targets if t != obs_mod and t != extra] # Completing the seed modality's own remaining frames now counts as # "predicting something" too (see the Future Prediction tab's chip for # it), so this is only a hard requirement when that's also off. if not targets and not complete_seed_modality: raise ValueError("Pick at least one modality to predict the future in.") chain = fp_chain(obs_mod, extra, targets) # No HYPERPARAM_PRESETS base here -- Future Prediction uses plain CONFIGS # hyperparameters directly (sliders already default to those values; only # overridden if changed). overrides = {} for (key, param), value in zip(hyperparam_spec, hyperparam_values): overrides.setdefault(key, {})[param] = value # Stripped so a stray trailing newline in the textarea can't miss the cache. steer_text = (steer_text or "").strip() or None seed, seed_tokens = int(seed), int(seed_tokens) top_p, top_k = float(top_p), float(top_k) complete_seed_modality = bool(complete_seed_modality) spec = None if raw_video_path else canonical_fp_spec( example_stem, obs_mod, seed_tokens, extra, steer_text, chain, seed, top_p, top_k, overrides, complete_seed=complete_seed_modality, ) return FPRequest(example_stem, obs_mod, seed_tokens, extra, chain, steer_text, overrides, seed, top_p, top_k, spec, complete_seed_modality, raw_video_path) # --- the cache itself ---------------------------------------------------------- class CacheStore: """Read-only view over the precomputed manifest. Every accessor is total: a missing, truncated or half-uploaded cache degrades to "no cache", never to an exception -- the demo has to boot offline (same contract as _list_examples_safe in app.py).""" def __init__(self, data=None): self.data = data or {} self.entries = [e for e in self.data.get("entries", []) if isinstance(e, dict)] @property def enabled(self): return bool(self.entries) def lookup(self, spec): """The entry generated with exactly this configuration, or None.""" if not spec: return None return next((e for e in self.entries if e.get("spec") == spec), None) def prefill_entry(self, task): return next((e for e in self.entries if e.get("task") == task and e.get("prefill")), None) def prefill_stem(self, task): """Which example clip the page-load prefill for this tab is about. Read back out of the entry's own spec so it cannot drift from it.""" entry = self.prefill_entry(task) spec = (entry or {}).get("spec") or {} if task == "fp": return spec.get("stem") return (spec.get("source") or {}).get("stem") def resolve_results(self, entry): """{result_key: local mp4 path | text}, shaped exactly like generate_any_to_any's return so the normal render path consumes it unchanged. Videos that failed to download are dropped rather than passed on as None -- a bad path in a build-time gr.Video(value=...) raises at import and the Space never comes up.""" results = {} for key, item in ((entry or {}).get("results") or {}).items(): if not isinstance(item, dict): continue if item.get("type") == "text": results[key] = item.get("value") continue path = get_precomputed_file(item.get("file") or "") if path and os.path.exists(path): results[key] = path return results def load_store(): """Fetches and validates the manifest. Never raises.""" try: data = load_precomputed_manifest() except Exception: traceback.print_exc() return CacheStore() if not data: return CacheStore() if data.get("v") != SCHEMA_VERSION: print(f"[precompute] manifest schema v{data.get('v')} != v{SCHEMA_VERSION}; cache disabled") return CacheStore() if data.get("code_fingerprint") != code_fingerprint(): print("[precompute] WARNING: CONFIGS / HYPERPARAM_PRESETS / chain presets have changed " "since these results were generated — lookups will miss and fall back to live GPU. " "Re-run scripts/precompute_demo_outputs.py.") store = CacheStore(data) print(f"[precompute] loaded {len(store.entries)} precomputed entr" f"{'y' if len(store.entries) == 1 else 'ies'}") return store