Spaces:
Running on Zero
Running on Zero
Muhammad Uzair Khattak
Claude Sonnet 5
Add Chained/Direct presets, Future Prediction UI cleanup, and gallery tweaks
31c6b21 Download precompute_cache.py from EPFL-VILAB/Video-4M: direct link, hf CLI and curl.
- Browser
- Download file 22 kB
-
https://huggingface.co/spaces/EPFL-VILAB/Video-4M/resolve/main/precompute_cache.py
- Command line
-
hf download hf://spaces/EPFL-VILAB/Video-4M/precompute_cache.py
-
curl -L -o precompute_cache.py https://huggingface.co/spaces/EPFL-VILAB/Video-4M/resolve/main/precompute_cache.py
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)] | |
| 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 | |