Video-4M / precompute_cache.py
Muhammad Uzair Khattak
Claude Sonnet 5
Add Chained/Direct presets, Future Prediction UI cleanup, and gallery tweaks
31c6b21
Raw History Blame Contribute Delete
22 kB
"""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