Compare commits

...
46 Commits
Author SHA1 Message Date
shaoxiongduanandClaude Opus 4.7 c2e89f22d3 [feat] eval: async VideoPool + GPU-side common metrics + safe optical-flow chunks
Pipelines path → tensor decode behind GPU metric compute via a new
VideoPool, so multi-sample eval runs no longer serialize disk I/O and
metric work. The Evaluator owns one pool per evaluate(samples=...)
call; each EvalWorker is a single-GPU consumer that grabs decoded
samples from the shared queue (work-stealing across replicas when
num_gpus > 1).

Worker pre-uploads video/reference to its device once per sample so
every metric in the loop consumes the same GPU-resident tensor (no
per-metric .to(device) traffic).

SSIM and PSNR move to the GPU — at 1080p × 121 frames the CPU path
was both slow (5–10 s/pair) and contended with the loader thread for
DDR bandwidth. LPIPS gains a chunk_size knob (default 8) that caps
peak from ~60 GB to ~5 GB with bit-identical output. Optical-flow
metrics drop chunk_size to 1 because DPFlow's cost volume is ~4 GB
per frame pair at 1080p (matches mhuo/ptlflow upstream).

physics_iq and a handful of vbench metrics drop their list-batch
shims — the per-sample contract is now uniform across the suite.

Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
2026-05-11 04:48:45 +00:00
shaoxiongduan e47a3c5aad [chore] eval: clear pre-commit --all-files lint backlog
CI runs `pre-commit run --all-files` which surfaces every latent issue
across the eval suite, not just the ones in changed files. Before this
commit, that surfaced 9 ruff errors and 47 mypy errors that had built
up over earlier PRs (last touched in `[style] eval: ruff/yapf pass` and
related). Cleaning them in one sweep so subsequent eval PRs land green.

Ruff
- yapf reformat: 3 files (entrypoints/cli/eval.py, optical_flow/_shared.py,
  physics_iq/utils.py).
- B024: drop ABC inheritance from PromptDataset (it has no abstract
  methods; subclasses just populate self._rows).
- B027: BaseMetric.setup is intentionally an optional no-op override
  (metrics with no eager state inherit it). Document and noqa instead
  of forcing every native metric to declare an empty override.
- SIM105 / SIM115: contextlib.suppress for ipc_collect, mkstemp for
  the temp video file in vbench.scene.
- UP038: isinstance(x, (A, B)) -> isinstance(x, A | B).

Mypy
- Optional-narrowing pattern across 12 metric files: each metric class
  initialised self._model (and sometimes _processor / _tokenizer /
  _head) to None and reassigned in setup(), but mypy narrows the type
  to None and flags every later attribute access. Annotate the slots
  as Any. Same fix already applied to physics_iq + human_action +
  videoscore2 in earlier commits; this extends it to the rest of the
  suite.
- Untyped helpers: annotate _safe_amt_forward (motion_smoothness),
  _patch_detectron2_registries (_grit_helper), _Loader (vbench/__init__).
- _grit_helper: function-attribute writes (`func._patched_idempotent`)
  type: ignore'd at the assignment site.
- imaging_quality: rename `all_scores` rebinding to `chunks` / `per_frame`
  so mypy doesn't carry the list[Any] type into the cat'd tensor.
- worker._resolve_video_input: add Any -> Any annotation.

No runtime behaviour changes. 33/33 eval pytest still pass.
2026-05-08 01:42:25 +00:00
shaoxiongduan d40fbfc534 [docs] eval: tone down doc voice, fix two stale agent-workflow paths
Voice pass over the three eval-related docs to bring them in line with
the project's existing voice (see docs/contributing/pull_requests.md,
docs/contributing/testing.md, docs/getting_started/installation/gpu.md):

* fastvideo/eval/README.md
* docs/contributing/eval-metrics.md
* .agents/workflows/evaluation-development.md

Concretely: cut em-dashes from ~30 across the three files to 1 (left
in a table cell where it reads naturally), removed "no X, no Y"
negation patterns, dropped buzzy phrasing ("first-class", "drop a
file", "out of the box"), and demoted bold imperatives ("Do not X")
to plain prose where the surrounding context already conveys the
instruction.

Two stale path references in evaluation-development.md fixed while
the file was open:

* "Update .agents/skills/evaluate-video-quality.md" pointed at a flat
  file; skills in this repo are directories. Corrected to
  .agents/skills/evaluate-video-quality/SKILL.md.
* "Check the evaluation_registry.md" referenced a bare filename that
  does not exist; aligned to .agents/memory/evaluation-registry/README.md
  to match the other references in the same doc.

No code changes; tests still 33/4 on a 1-GPU borrow.
2026-05-08 01:17:07 +00:00
shaoxiongduan 126a52ce32 [bugfix] eval: vbench group skips missing-dep metrics + correct install hint + checkpoint/transformers fixes
create_evaluator(metrics='vbench') now scores the 11 vbench sub-metrics
that don't need detectron2 instead of crashing on construction. Explicit
metric names (e.g. metrics=['vbench.color']) still raise ImportError
with the full two-step install command.

Group selectors filter missing deps
- evaluator._resolve_metric_names: when expanding a group prefix or
  'all', drop metrics whose declared dependencies aren't importable.
  One warning per skipped metric. Explicit names pass through unchanged
  so the missing dep surfaces as ImportError -- the friendly contract
  for "user asked for this specific metric."
- registry.missing_dependencies(name): new helper exposing per-metric
  dep status without instantiation.

Install hint actually satisfies the dep
- registry._install_hint(metric, dep) replaces _extra_for() in the
  ImportError formatter. detectron2 needs the base extra PLUS a git+
  install with build-isolation off; the old hint sent users in a
  circle. All other deps still resolve via the existing extra map.
- README install table aligned to use uv pip install verbatim
  (matches docs/getting_started/installation/gpu.md).

Other VBench/VLM fixes uncovered by running the group end-to-end
- vbench.human_action: filename was 'l16_25m.pth' which 404s on the
  OpenGVLab/VBench_Used_Models HF repo. The actual UMT-L Kinetics-400
  checkpoint there is 'l16_ptk710_ftk710_ftk400_f16_res224.pth'
  (matches the metric's expected vit_large_patch16_224 / num_classes=400
  / all_frames=16 shape exactly).
- videoscore2: switch to AutoModelForImageTextToText with a
  AutoModelForVision2Seq fallback (the legacy alias is being phased
  out in transformers 4.45+); pass dtype=torch.bfloat16 explicitly
  (transformers 4.57 deprecated torch_dtype= in favor of dtype=).

Mypy hygiene scoped to the two metric files I touched
- self._model / _processor / _tokenizer typed as Any to silence the
  pre-existing None-narrowing errors that surface only when these files
  are staged. Same pattern used across the eval suite; fixing it
  branch-wide is out of scope for this change.

Tested
- pytest fastvideo/tests/eval/ -> 33 passed (single-GPU subset).
- create_evaluator(metrics=['vbench.color']) -> raises ImportError with
  the full uv pip install ... && uv pip install --no-build-isolation
  'git+https://github.com/facebookresearch/detectron2.git' command.
- _resolve_metric_names('vbench') -> 11 of 16 vbench metrics, with
  4 detectron2-deps + 1 qwen_omni_utils-dep correctly filtered.
- vbench.human_action loads against the correct checkpoint name
  (verified end-to-end on a noise tensor; metric runs forward pass).
2026-05-08 00:57:04 +00:00
shaoxiongduan 419e1c68f1 [feat] eval physics_iq: self-contained dataset + auto-fetch from public bucket
`get_dataset("physics_iq")` now works with no kwargs. The manifest CSV is
vendored under fastvideo/eval/metrics/physics_iq/_vendored/; per-scenario
videos, masks, and switch-frames auto-fetch on first use from the public
DeepMind bucket (https://storage.googleapis.com/physics-iq-benchmark) into
${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq/, sibling to the existing
models/torch/clip/ subdirs.

Why
- Old default dataset_root was /root/physics-IQ-benchmark, which doesn't
  exist on shared hosts and isn't documented anywhere.
- Examples crashed with FileNotFoundError before any user-facing message
  pointed at where to get the data.
- The official upstream download script needs gcloud SDK; the bucket is
  also reachable over plain HTTPS (verified) so no SDK dependency is
  needed.

Behavior
- dataset_root= is now an opt-in override (mirroring vbench's
  full_info_path=); defaults to get_cache_dir() / "datasets" /
  "physics_iq".
- auto_download=True by default; flip to False for air-gapped runs.
- FASTVIDEO_PHYSICS_IQ_BUCKET_URL env var redirects to internal mirrors.
- Atomic .part -> rename for safe concurrent SLURM-rank fetches.

Bench script fix
- bench_physics_iq.py: --limit was applied as a post-construction
  list slice, after the dataset module had already eagerly resolved
  every scenario's on-disk paths. With auto-download that means we'd
  pull all 198 scenarios on every smoke run. Pass limit=args.limit
  through to get_dataset() so partial-data smoke runs only fetch what
  they need. --dataset-root is now optional (defaults to the cache
  path).

Convention for future vendored files
- _vendored/ subdirs hold upstream-provenance content. They're
  auto-skipped by metric discovery (the leading _) and by codespell
  (single */_vendored/* glob in [tool.codespell].skip), so future
  vendored files require no further config.
- docs/contributing/eval-metrics.md updated to point at this convention
  for new metrics that ship paired-reference datasets.

Tested
- pytest fastvideo/tests/eval/ -> 33 passed (1-GPU subset); 4 multi-GPU
  tests skipped on 1-GPU borrow as expected.
- Smoke: get_dataset("physics_iq", limit=1) auto-fetches the 5 expected
  assets; physics_iq.{mse,spatial_iou} on take-1 self-pair returns
  0.0 / 1.0.
2026-05-08 00:35:46 +00:00
shaoxiongduan 938bc3c972 [refactor] eval: per-benchmark extras + drop audio metrics from this PR
**Per-benchmark eval extras (Option B layout)**

Replace the single ``[eval]`` rollup with per-benchmark groups so users
can install only what they need:

* ``[eval-vbench]`` — ``openai-clip``, ``pyiqa``, ``easydict`` (covers
  12 of the 16 vbench sub-metrics; the four GRiT ones still need
  ``detectron2`` installed manually).
* ``[eval-physics-iq]`` — empty group; documents intent (all
  physics_iq metrics already use base fastvideo deps).
* ``[eval]`` — sensible default rollup: common.lpips, optical_flow.*,
  videoscore2, plus eval-vbench + eval-physics-iq. The 80% case.
* ``[eval-full]`` — adds ``qwen-omni-utils`` for the AVoCaDO-based
  ``vbench.scene`` metric. Detectron2 still manual.

Most deps are *not* in any of these — they're already in base
fastvideo's pinned dependencies (transformers, timm, einops, scipy,
omegaconf, opencv-python, imageio). Per-metric runtime patching keeps
versions consistent across all metrics.

Registry's missing-dep ImportError now points at the right extra
(e.g. ``pip install 'fastvideo[eval-vbench]'`` for vbench metrics)
instead of the generic ``[eval]`` it used to say.

**Drop audio metrics from this PR**

Five ``audio.*`` metrics (clap_score, frechet_distance, kl_divergence,
wer, audiobox_aesthetics; 452 lines) came in with the original wm-eval
port and have not been iterated, tested, or exercised end-to-end since.
None of the test suite or example scripts touched them. They would
have forced 7 untested deps into a public extra.

Pull them out of this PR. Code is preserved on the side branch
``shao/eval-audio`` (pushed to origin) for a follow-up audio-eval PR
that lands them with proper tests + a real end-to-end run.

**Also: actually remove the ``_assets/`` gitignore rule**

The previous commit (``381a7aae``) claimed to revert the
``fastvideo/eval/_assets/`` gitignore carve-out but the change was
unstaged at commit time, so the rule remained. This commit removes it
for real. Untracked content under that path is now visible to ``git
status`` again, which is what we wanted — that scratch dir is not the
kind of carve-out the project root ``.gitignore`` should carry.

Net registry: 32 metrics → 27; tests: 33 pass / 4 multi-GPU skipped
(unchanged); ``shao/eval-audio`` branch pushed for the follow-up.
2026-05-07 22:19:14 +00:00
shaoxiongduan 381a7aae66 [cleanup] eval: drop duplicate scripts under scripts/eval/, revert _assets/ gitignore
* Delete ``scripts/eval/score_folder.py`` and ``scripts/eval/run_vbench_e2e.py``.
  Both pre-dated the simpler ``examples/inference/eval/score_folder.py`` /
  ``bench_vbench.py`` versions and now overlap (one even name-collides).
  Users follow the ``examples/`` path going forward.
* Drop the ``fastvideo/eval/_assets/`` ignore rule. It was specific to a
  branch-local synthetic-flow scratch dir that no longer ships in this
  PR; not the kind of carve-out the project root ``.gitignore`` should
  carry.
2026-05-07 21:55:03 +00:00
shaoxiongduan 7fa4fb50bb [cleanup] eval: PR-readiness audit — empty inits, fix stale docs, drop unused deps
Three loose ends from earlier hygiene passes that the final PR review
caught:

* Empty out 27 sub-metric ``__init__.py`` files that still re-exported
  the metric class. The earlier ``54f26bea`` commit only cleaned 2 of
  them; the rest still carried ``from .metric import FooMetric  # noqa:
  F401`` lines copied from the original wm-eval port. Auto-discovery
  doesn't need them, and ``get_metric("group.name")`` is the canonical
  access pattern. Now uniformly empty across all sub-metric inits.

* Fix stale layout references in ``fastvideo/eval/README.md`` and
  ``.agents/workflows/evaluation-development.md`` that still described
  the abandoned ``<bench>/external/upstream/`` per-metric layout. The
  actual layout is the flat ``fastvideo/third_party/eval/<bench>/`` we
  switched to during the port. Also drops the stale ``_third_party``
  example from ``fastvideo/eval/metrics/__init__.py``'s docstring (the
  underscore-skip rule is unchanged; the example was just wrong).

* Drop ``open_clip_torch`` and ``torchmetrics`` from the
  ``[eval]`` extra. Neither is imported anywhere under
  ``fastvideo/eval/`` (or in the vbench submodule we vendor); they were
  carried over from the wm-eval port's heavier dep graph.

Verified: 32 metrics still register, 33/4 tests pass/skip on 1 GPU.
2026-05-07 21:47:05 +00:00
shaoxiongduan de3cd6aab2 [feat] eval: accept video paths at the worker boundary; bound batch memory
``Evaluator.evaluate`` and the one-shot ``evaluate()`` helper now accept
``video`` / ``reference`` as either a pre-loaded ``(T, C, H, W)`` tensor
or a path-like (``str`` / ``Path``). Paths are decoded inside the
worker thread that picks the sample up — see
``EvalWorker._resolve_video_input``.

Why this matters for batch eval: the dispatcher in
``Evaluator.evaluate(samples=[...])`` submits every sample to the
``ThreadPoolExecutor`` upfront. With pre-loaded tensors that meant
every video in the batch was resident in CPU RAM at once (≈ 3 GB per
video at 1088×1920×121); a full benchmark of hundreds of clips would
OOM before scoring started.

With paths, the queued futures hold cheap strings; only ``num_gpus``
videos are decoded concurrently. Peak resident memory becomes
``O(num_gpus)`` instead of ``O(len(samples))``. No dispatcher rewrite
needed — the change is ~10 lines at the worker boundary.

* ``EvalWorker.evaluate`` now normalizes ``sample["video"]`` and
  ``sample["reference"]`` through ``_resolve_video_input`` (path →
  ``load_video`` decode; ``(1, T, C, H, W)`` → squeeze; tensor →
  passthrough). Metrics keep the existing contract: by the time
  ``compute()`` runs, ``sample["video"]`` is always a 4-D tensor.
* ``score_folder.py`` and ``bench_vbench.py`` simplified to pass paths
  directly. ``bench_physics_iq.py`` already did.
* ``fastvideo.eval.api.evaluate`` signature widened to ``Tensor | str |
  Path``; the worker handles the decode either way.

Tests: new ``test_evaluator_paths.py`` covers the kwargs form, the
samples-list form, mixing paths and tensors in the same batch, parity
between path-form and tensor-form scores, and the missing-path
exception path. End-to-end ``score_folder.py`` smoke run confirmed.
2026-05-07 20:29:04 +00:00
shaoxiongduan a64e5e62ab [bugfix] eval CLI: clean up _expand_paths and _cmd_run lint flags
Two small fixes in ``fastvideo/entrypoints/cli/eval.py``:

* ``_expand_paths`` deduplicated through ``[x for x in out if not (x in
  seen or seen.add(x))]`` — a known idiom that ruff flags because
  ``set.add`` returns ``None`` (B023/func-returns-value). Replace with
  an explicit loop so the side effect doesn't ride on the boolean
  expression.
* ``_cmd_run`` constructed ``metrics_arg`` through an
  ``if/else``-block where SIM108 (project policy) prefers a ternary.
  Fold to a single conditional expression.

Pure mechanical cleanups; ``ruff check`` is now clean on the file and
the eval CLI smoke test (``score_video.py`` against a self-paired
mp4 → ``common.psnr=100``) still passes.
2026-05-07 09:50:53 +00:00
shaoxiongduan f4200cc3f3 [style] eval: ruff/yapf pass across the suite
Pure formatting pass — no behavior changes. Most edits are one of:

* yapf splitting / re-flowing kwarg-heavy callsites (argparse setup,
  ``MetricResult(...)`` construction).
* ruff auto-fixes around imports (collapsing ``Iterable`` to
  ``collections.abc``, removing the ``typing.Tuple`` shim where
  ``tuple[...]`` works, dropping unused imports).
* ``zip(..., strict=False)`` added to every two-iterable zip call.
* ``getattr(np, "trapz")`` → ``np.trapz`` where the fallback can read
  the attribute directly (the ``trapezoid`` lookup still uses
  ``getattr`` because the name itself is the moving target).

Verified the test suite still passes after the pass.
2026-05-07 09:47:35 +00:00
shaoxiongduan 363230b372 [examples] eval: expose num-frames/height/width on bench_* scripts
The two end-to-end runners called ``VideoGenerator.generate_video``
without specifying generation dimensions, leaving the model to fall
back on whatever default sampling shape it carries internally — which
isn't necessarily what the user wants for a benchmark run.

Mirror the basic ``examples/inference/basic/basic_ltx2.py`` script:
default to ``121 x 1088 x 1920`` (LTX2's intended sampling shape) and
expose ``--num-frames`` / ``--height`` / ``--width`` so users can
downscale for smoke runs without editing the script.

Verified end-to-end on a borrowed H200: ``bench_vbench.py --limit 1
--num-frames 49 --height 480 --width 768`` generates an mp4 in ~97s
and scores it with ``vbench.aesthetic_quality`` (auto-downloads
``Davids048/LTX2-Base-Diffusers`` weights through ``from_pretrained``).
2026-05-07 09:06:44 +00:00
shaoxiongduan 1dcdacc35a [examples] eval: add 4 simple end-user scripts
Add four small example scripts under ``examples/inference/eval/``,
each driving the public eval API in a different real-world shape:

* ``score_video.py`` — score one mp4 on one GPU. Smallest possible
  use of ``create_evaluator`` + ``Evaluator.evaluate``. Optional
  ``--reference``, ``--text-prompt``, ``--fps`` for paired / prompt-
  aware / fps-aware metrics.
* ``score_folder.py`` — score every mp4 in a directory across ``--num-gpus``
  replicas via ``Evaluator.evaluate(samples=[...])``. Pair each
  generated video with a same-name reference under ``--reference-dir``
  if you want paired metrics.
* ``bench_vbench.py`` — full VBench end-to-end: ``get_dataset("vbench",
  dimensions=...)`` → ``VideoGenerator`` (LTX2 by default) → score
  with the matching ``vbench.*`` sub-metrics → print per-metric
  averages. ``--skip-generation`` re-uses existing mp4s under
  ``--videos-dir`` so you can iterate on metric selection without
  re-paying generation cost.
* ``bench_physics_iq.py`` — full Physics-IQ end-to-end:
  ``get_dataset("physics_iq", dataset_root=...)`` → ``VideoGenerator``
  → score with the composite ``physics_iq`` metric → aggregate via the
  upstream's ``aggregate_components`` recipe.

The pre-existing ``basic_ltx2_eval.py`` and ``eval_ltx2_vbench.py`` are
left in place — they target the single-prompt generate-and-score case
and complement (rather than overlap with) the new full-dataset runners.

Verified end-to-end on a borrowed GPU node:
* ``score_video.py``: PSNR=100, SSIM=1.0 on a self-paired mp4 (sanity).
* ``score_folder.py``: produces well-formed scores.json.
2026-05-07 08:20:20 +00:00
shaoxiongduan 0a7a4a21f6 [test] eval: cover registry, single-replica, multi-GPU, dataset paths
Add four end-to-end test modules under ``fastvideo/tests/eval/``,
all driving the public API only — no reaches into ``EvalWorker``,
``BaseMetric``, or the ``_REGISTRY`` dict. Each test mirrors a real
caller flow.

* ``test_registry.py`` — ``list_metrics`` invariants, group-prefix
  resolution (verified against the no-model ``physics_iq`` group so
  the test stays cheap), unknown-metric error path.
* ``test_evaluator_single.py`` — one-shot ``evaluate(...)``,
  long-lived ``Evaluator`` with both kwargs and ``samples=[...]``
  shapes, deterministic scoring, the legacy ``(1, T, C, H, W)``
  back-compat unwrap. Runs on CPU using ``common.psnr`` /
  ``common.ssim`` so it's free in CI.
* ``test_evaluator_multi_gpu.py`` — auto-skips when fewer than 2 CUDA
  devices visible. Pins the round-robin contract by computing a
  single-GPU baseline and asserting bit-equivalent scores under
  multi-GPU dispatch with the same input list, plus the kwargs-form
  → worker-0 invariant and ``release_cuda_memory`` no-crash check.
* ``test_evaluator_with_dataset.py`` — full ``get_dataset("vbench")
  → Evaluator.evaluate(**row)`` flow. Synthesizes random video tensors
  per row (no diffusion model needed) and verifies that extra dataset
  keys (``prompt``, ``n_samples``, ``dimensions``, ``auxiliary_info``)
  flow through unused metrics without breaking them.

29 tests total: 25 pass on a single GPU, all 29 pass on 2 GPUs.
2026-05-07 08:01:09 +00:00
shaoxiongduan c30d731992 [bugfix] eval CLI: coerce numpy / torch / Path leaves before json.dumps
``fastvideo eval run --output scores.json`` would crash with
``TypeError: Object of type float32 is not JSON serializable`` whenever
the chosen metric set landed numpy or torch values in
``MetricResult.details``. The optical-flow metrics in particular
populate ``per_frame_metrics`` with ``np.float64`` scalars, and pretty
much any metric is one ``np.percentile`` call away from the same crash.

Pass a ``default=`` callback to ``json.dumps`` that walks the unknown
leaves and coerces:

* ``np.integer`` / ``np.floating`` / ``np.bool_`` → native Python scalars
* ``np.ndarray`` → ``.tolist()``
* ``torch.Tensor`` → detached CPU ``.tolist()``
* ``pathlib.Path`` → ``str``

Anything else still raises ``TypeError`` — the goal is to handle the
known metric outputs cleanly, not to silently coerce arbitrary objects.
2026-05-07 07:49:50 +00:00
shaoxiongduan f8eaff4292 [refactor] eval: move physics_iq dataset traversal under eval/datasets
The Physics-IQ metric package was carrying a ~200-line dataset loader
(``PhysicsIQDataLoader`` + ``PhysicsIQScenario`` + an FPS-conversion
helper + manifest-walking constants) inside its metric directory, with
``__init__.py`` re-exporting the loader as a public name. Two distinct
concerns were mixed: dataset traversal (a user-of-the-metric concern)
and metric-pipeline configuration (an internal concern).

Split them:

* New ``fastvideo/eval/datasets/physics_iq.py`` registers a
  ``PhysicsIQPromptDataset(PromptDataset)`` under
  ``@register_dataset("physics_iq")``. It walks ``descriptions.csv``,
  resolves per-take video and real-mask paths, FPS-converts the source
  release on cache miss, and yields one sample dict per take-1
  scenario shaped to drop straight into ``Evaluator.evaluate(**row)``.
  The previously-public ``PhysicsIQScenario`` dataclass moves with it.
* Metric defaults (``DEFAULT_TARGET_FPS=30``,
  ``DEFAULT_DURATION_SECONDS=5``) move into
  ``metrics/physics_iq/utils.py``, where they were already used as
  default kwargs.
* ``metrics/physics_iq/models.py`` is deleted; ``__init__.py`` is
  emptied to match the rest of the metric directories.

Users now do ``get_dataset("physics_iq", dataset_root=...)`` to walk
the corpus instead of reaching for ``PhysicsIQDataLoader`` directly.
The old import path is dropped without a back-compat shim — the eval
suite hasn't shipped yet, so there's no API contract to honor.

The eval-suite README (``fastvideo/eval/README.md``) gets a short
``Prompt datasets`` section showing the ``get_dataset`` workflow. The
project root README is unchanged.
2026-05-07 07:45:45 +00:00
shaoxiongduan 76e7048b3d [cleanup] eval: drop class re-exports from sub-metric __init__.py files
Two sub-metric ``__init__.py`` files re-exported their metric class
(``VBClapScoreMetric``, ``AestheticQualityMetric``) while every other
sub-metric leaves ``__init__.py`` empty. Auto-discovery imports each
``metric.py`` module by full path, so the re-export was redundant —
and the inconsistency obscured the contract for new contributors
(adding a metric should not require touching ``__init__.py``).

Standardize on empty sub-metric ``__init__.py``; users instantiate
metrics through ``get_metric("group.name")`` rather than reaching for
the class object directly.

The ``physics_iq`` group's ``__init__.py`` also re-exports a dataset
loader and scenario dataclass; that one's left alone in this commit
pending a separate move of those helpers into ``fastvideo/eval/datasets/``.
2026-05-07 07:39:52 +00:00
shaoxiongduan b1386a7a78 [refactor] eval: tighten compute(sample) contract to a single MetricResult
The ``EvalWorker`` always invokes metrics on a single video and only
ever reads ``result[0]`` from the returned list, so every metric had a
dead ``for b in range(B)`` loop. Tighten the contract:

* ``BaseMetric.compute(sample) -> MetricResult`` — return one result,
  not a one-element list.
* ``BaseMetric._skip(sample, reason) -> MetricResult`` likewise.
* Sample-side: ``video`` and ``reference`` are ``(T, C, H, W)`` —
  no leading batch dim. The worker still unwraps a ``(1, T, C, H, W)``
  caller for back-compat, so existing user code keeps working.
* ``Evaluator.metric_names`` now reads through a public
  ``EvalWorker.metric_names`` property instead of poking the private
  ``_metrics`` dict.

All 28 metrics — common.{ssim,psnr,lpips}, optical_flow.*, audio.*,
vbench.* (16), physics_iq.* (5 incl. composite), videoscore2 — drop
their ``for b in range(B)`` loop and return a single ``MetricResult``.
List-shaped optional inputs (``text_prompt``, ``audio``,
``auxiliary_info``, ``actions``) are still accepted: each metric
unwraps a single-element list before use, so callers can keep the
existing ``[prompt]`` convention or pass a scalar — both work.
2026-05-07 06:45:14 +00:00
shaoxiongduan 664d6b3c23 [docs] eval: reflect new optical_flow group in README and contributor guide
Update the layout snippets in both ``fastvideo/eval/README.md`` and
``docs/contributing/eval-metrics.md`` to show the new
``optical_flow/`` group sibling to ``common/``, with the two
sub-metrics (``gt_optical_flow``, ``synthetic_optical_flow``) listed
explicitly.
2026-05-07 06:16:14 +00:00
shaoxiongduan 0688c5131a [refactor] eval: split optical_flow into its own group with gt + synthetic sub-metrics
Move ``common.optical_flow`` to a dedicated ``optical_flow`` group so
flow-based comparisons can register additional reference modes without
piling into ``common``. Two metrics live under the new group:

* ``optical_flow.gt_optical_flow`` — the existing video-vs-video
  comparison, ported verbatim minus a thin shared-helper extraction.
* ``optical_flow.synthetic_optical_flow`` — new metric that takes a
  per-frame action stream + a ``ThirdPersonCalibration`` JSON and
  predicts the reference flow analytically (Longuet-Higgins + off-pivot
  translation, no depth) instead of extracting it from a GT video.

Both metrics share ``optical_flow/_shared.py`` for ptlflow loading,
per-frame metric computation, and temporal aggregation, so scores are
directly comparable across the two reference modes.

The third-person predictor is vendored at
``optical_flow/synthetic_optical_flow/_thirdperson.py`` to keep the
metric self-contained; the calibration *fitter* (``calibrate.py`` etc.)
is intentionally not part of the eval suite — it's an offline fitting
tool.
2026-05-07 06:15:33 +00:00
shaoxiongduan 0b605d0a41 [cleanup] eval: drop legacy EvalRunner/EvalResult/WM_EVAL_CACHE refs
Several scaffolds that pre-date the runner refactor still pointed at
removed APIs:

* ``EvalResult`` (and its ``Evaluator.evaluate_dataset`` docstring)
  exposed a class that is never returned anywhere; drop it from the
  public API.
* ``Evaluator``'s module docstring still referenced the removed
  ``EvalRunner`` layer; trim to the current Evaluator → EvalWorker
  shape.
* ``WM_EVAL_CACHE`` survived as a fallback env var from the wm-eval
  port; remove and keep only ``FASTVIDEO_EVAL_CACHE``.
* The contributor docs documented ``batch_unit`` and
  ``trial_forward``/``Evaluator.calibrate()`` which were dropped from
  ``BaseMetric`` in earlier refactors; update to the current contract.
2026-05-07 06:15:10 +00:00
shaoxiongduan 8f5ebe2aeb [bugfix] eval: human_action — load kinetics labels from upstream submodule
The metric resolved the Kinetics-400 label file via a wm-eval-era
``_third_party/umt/kinetics_400_categories.txt`` path that no longer
exists after the vbench-as-submodule port. ``os.path.exists`` was
silently False, leaving the label dict empty; every prediction then
mapped to the empty string and every score returned 0.0.

Resolve the path through ``vbench.third_party.umt.__file__`` instead,
which always points at the pinned upstream submodule.
2026-05-07 06:14:55 +00:00
shaoxiongduan 84803076a0 [fix] eval: motion_smoothness OOM recovery + free-memory autoscale
Two changes that together let motion_smoothness score 1088×1920×121
videos on shared GPUs:

1. _get_scale() now queries torch.cuda.mem_get_info() free memory on
   every call instead of caching total_memory at setup() time.
   Adapts to whatever's actually available — other metric replicas
   already loaded, residual generator allocations, another process
   sharing the GPU. Upstream cached total_memory once which on a
   shared/loaded GPU lets AMT attempt 30+ GB correlation reshapes.

2. _safe_amt_forward() wraps the model call with two-axis OOM retry:
   - Halve batch until size 1 (per-pair memory dominates).
   - Halve scale_factor until 1/16 (AMT's internal feature-map
     resolution and therefore correlation volume size).
   - Bottom out at scale=1/16, batch=1; if still OOM, the resolution
     is genuinely impossible at this headroom and we re-raise.

The second change is needed because upstream's autoscale formula in
_get_scale mis-extrapolates at high resolution: it scales linearly
in pixel count, but AMT's correlation volume grows quadratically.
Rather than rewriting upstream's formula, the retry path makes the
metric robust to whatever the formula picks.

Verified on fs-mbz-gpu-085 (shared with another job holding 45 GB):

Before: motion_smoothness OOM at 31.75 GB allocation on 1088×1920×121.
After:  motion_smoothness=0.9897 across all 8 metrics, no OOM, all
        other scores (aesthetic_quality=0.5430, subject_consistency=
        0.8629, etc.) unchanged.

Parity against runs/vbench_smoke/results_v3.json byte-identical.
2026-05-06 01:56:41 +00:00
shaoxiongduan 5cc337fdb6 [feat] eval: add basic_ltx2_eval.py + score_folder.py scripts
Two new entrypoints showcasing the post-refactor eval surface:

- examples/inference/eval/basic_ltx2_eval.py — generate one LTX2
  video using the same parameters as basic_ltx2.py (same prompt,
  model, 1088×1920×121), then score it with the prompt-aware VBench
  subset. Builds Evaluator directly, calls evaluate(**kwargs) per
  the single-sample API; no runner, no argparse.

- scripts/eval/score_folder.py — bulk-score every video in a folder
  with prompt-free VBench metrics. Optional --prompts-json maps
  filenames to prompts to enable prompt-aware metrics. Uses
  Evaluator.evaluate(samples=[...]) for multi-GPU fan-out.

Tested on fs-mbz-gpu-085:

- score_folder.py: 3 duplicate mp4s × 4 metrics → 3 identical score
  sets, summary aggregated, scores.json written. Clean run.

- basic_ltx2_eval.py: generation succeeded; scoring path verified
  for 7 of 8 metrics on the LTX2 output (aesthetic_quality,
  subject_consistency, background_consistency, imaging_quality,
  temporal_flickering, dynamic_degree, overall_consistency).
  vbench.motion_smoothness OOMs at 1088×1920 on a shared GPU because
  VBench's AMT memory autoscale reads total_memory rather than
  mem_get_info() free memory, so it underestimates the required
  resolution scale-down when another process holds 45 GB. Documented
  in the script docstring; not refactor-related (same behavior at
  HEAD pre-refactor). Drop motion_smoothness from METRICS or run on
  a dedicated GPU to score it.

Added a torch.cuda.empty_cache() between generator.shutdown() and
evaluator construction to free residual generation memory.
2026-05-06 01:46:29 +00:00
shaoxiongduan 2b8e5a56a8 [refactor] eval: drop batch_unit/trial_forward from remaining metrics
Completes the d1bc1413 cleanup pass on the five files that were
skipped because dataloader had pending modifications. With those
modifications now landed (commits 7cdbcc94 + c2a22560), the dead
code in optical_flow + the four aux-using vbench metrics
(color/multiple_objects/object_class/spatial_relationship) can go.

Mechanical removal of batch_unit class attrs and trial_forward
method overrides — neither is read anywhere since calibrate() was
deleted in 502f061d.

Parity verified on fs-mbz-gpu-085: byte-identical to v3 baseline.
2026-05-06 00:51:00 +00:00
shaoxiongduan e117167eb9 [feat] eval: full per-frame optical-flow metric set
Replace the single mean-EPE port with mhuo's complete validation set
(see fastvideo/training/ptlflow_validation.py in mhuo's tree):

  Per-frame: mf_epe, mf_angle_err, mf_cosine, mf_mag_ratio,
             pixel_epe_mean/max, px_angle_rmse, grid_epe_mean/max,
             fl_all, foe_dist, flow_kl_2d
  Aggregated over time: <name>_mean / _std / _max / _auc, plus
                        divergence_onset_frame / divergence_threshold

Added supporting math: least-squares Focus-of-Expansion estimation,
2D KL divergence over (angle, log-magnitude) histogram, trapezoid
integral with numpy 2.x compat shim.

Headline ``score`` is pixel_epe_mean_mean (lower is better);
everything else lives in MetricResult.details so downstream
consumers can pick whichever scalar they care about.

Pairs gen and ref videos across all B inputs and batches the model
through them together (chunk_size=16) for GPU efficiency.

Authored by dataloader.
2026-05-06 00:48:51 +00:00
shaoxiongduan 4f9af56b78 [bugfix] eval: per-row skip in aux-using vbench metrics
The four aux-using vbench metrics (color, multiple_objects,
object_class, spatial_relationship) used to bail at the top of
compute() if auxiliary_info was None, then unconditionally indexed
aux[b]["<key>"] for every row. This crashed with KeyError on rows
that had aux but lacked the metric's specific key — common once
the dataset yields heterogeneous rows under "vbench" / multi-dim
runs (each row carries flat aux populated only for its own
dimensions).

Switch to per-row skip: each row checks for its own required key
and emits MetricResult(score=None, details={"skipped": "..."}) if
absent. Lets a single evaluator.evaluate(samples=[...]) call score
heterogeneous-aux rows in one pass.

Also handle two structural sub-cases:
- multiple_objects requires "<a> and <b>" in aux["object"]; rows
  without the separator (single-object data leaking in) skip.
- object_class is the inverse: rows whose "object" contains " and "
  belong to multiple_objects' territory and skip here.
- spatial_relationship's nested {object_a/object_b/relationship}
  sub-dict is read defensively; missing keys skip with a reason.

Pairs with the flat-aux-at-load behavior in VBenchPromptDataset
(commit 06735161). Without these per-row skips, dimensions="all"
runs crash on the first row that lacks the active metric's key.

Authored by dataloader.
2026-05-06 00:48:36 +00:00
shaoxiongduan 33efe7a673 [feat] eval: add EvalResult aggregate type; ignore _assets/
Two small additions paired with the in-progress eval orchestration:

- types.py: EvalResult dataclass with summary/per_video fields plus
  from_raw / save / print helpers. Used by scripts and any future
  callers that aggregate per-sample MetricResults into a corpus-level
  summary (e.g. scripts/eval/run_vbench_e2e.py).

- .gitignore: exclude fastvideo/eval/_assets/ which holds calibration
  data, downloaded videos, and run dumps used by the synthetic-flow
  evaluator that we don't want to track.

Authored by dataloader.
2026-05-06 00:48:19 +00:00
shaoxiongduan b07cc3f9d4 [docs] eval/vbench: document the aux-flatten invariant per dimension
The flatten loop in VBenchPromptDataset strips exactly one level of
upstream's {dim_name: ...} wrapper. For most VBench dims this leaves
a flat scalar dict ({"color": "red"}, {"object": "person"}). For
spatial_relationship, upstream double-wraps, so one level of unwrap
leaves {"spatial_relationship": {object_a, object_b, relationship}} —
which is exactly what SpatialRelationshipMetric reads.

That last case looks accidental; it isn't. Documenting all four shapes
inline so a future reader doesn't "simplify" the wrapping away and
silently break spatial_relationship scoring.

Empirically verified output for all four aux dims matches the
documented shapes.
2026-05-06 00:44:40 +00:00
shaoxiongduan 29eb4109cc [refactor] eval: drop dead batch_unit and trial_forward across metrics
The Evaluator's calibrate() is gone, so batch_unit class attrs and
trial_forward() method overrides have no callers — pure cruft from the
old auto-calibration design. Mechanical removal across 15 metric files.

Internal time-dim chunking (the part of batching that actually does
work) lives on as metric-owned _chunk_size constants set in __init__,
unaffected by this change.

Five files (color/multiple_objects/object_class/spatial_relationship/
optical_flow metrics) skipped — they have dataloader's in-flight
modifications uncommitted; left for that branch's cleanup pass.

Parity verified on fs-mbz-gpu-085: results byte-identical to v3
baseline.
2026-05-06 00:29:58 +00:00
shaoxiongduan d07a7691fe [refactor] eval: drop EvalRunner; flat script + helpers in eval.io
Match FastVideo's existing depth: VideoGenerator is the top-level
inference object with no Runner above it; loops live in scripts.
Eval should mirror that — Evaluator is the top-level scoring object,
EvalWorker × N is the layer below, and end-to-end pipelines (prompts
→ generate → score) are scripts, not classes.

Removed:
- fastvideo/eval/runner.py (EvalRunner + classmethod constructors).
- EvalRunner export from fastvideo.eval.

Added:
- fastvideo/eval/io/paths.py with sanitize_prompt, default_filename,
  glob_videos, build_eval_kwargs as free functions. The reusable bits
  of the runner survive; the class wrapping them does not.

Rewritten:
- scripts/eval/run_vbench_e2e.py is now top-to-bottom procedural:
  parse_args → dataset → optional generate() loop → optional score()
  loop. Reads in one pass; no classmethod indirection to chase.

Parity verified on fs-mbz-gpu-085 against runs/vbench_smoke/videos/
(2 videos × 3 metrics, dimensions=subject_consistency):
- 1-GPU run (results_v4.json) is byte-identical to v3 baseline.
- 2-GPU run (results_v4_2gpu.json) is byte-identical to v3 baseline
  (order-preserving fan-out via Evaluator.evaluate(samples=[...])).
2026-05-05 23:45:29 +00:00
shaoxiongduan 960485519f [refactor] eval: dict-based PromptDataset + EvalRunner orchestrator
Continues the eval refactor. Datasets now yield plain dicts (matches
ValidationDataset / VideoGenerator / Evaluator's kwargs-flowing-through
pattern). End-to-end orchestration moves into EvalRunner so neither the
dataset nor the Evaluator owns videos_dir / fps / filename / generator
state.

Layering:
  EvalRunner          dataset / generator / file conventions / manifest
    └── Evaluator     thin dispatcher (single evaluate(), 1 or N samples)
         └── EvalWorker × N    single-GPU metric replicas

Dataset (fastvideo/eval/datasets/):
- PromptDataset is just an Iterable[dict]. _rows holds dicts with
  prompt / n_samples / dimensions / auxiliary_info / ... — unused
  fields stay absent rather than living as Optional[None] on a
  dataclass. BasePromptDataset alias kept for back-compat.
- BenchmarkSample dataclass deleted. Sample TypedDict documents the
  recognized keys without forcing the schema.
- VBench flattens its nested {dim: {key: val}} aux into {key: val} at
  load time. Aux-using metrics already read flat keys; no metric edits.

Runner (fastvideo/eval/runner.py, new):
- Three named constructors: from_dataset, from_videos, from_samples.
- generate() drives the generator over the dataset (writes manifest);
  score() loads videos and dispatches via Evaluator.evaluate(list).
  run() = generate then score.
- Multi-GPU lives in the underlying Evaluator; the runner just hands
  num_gpus through. No double-orchestration.
- File naming convention (filename_fn) and eval-kwargs assembly
  (eval_kwargs_fn) are runner-side hooks, not dataset overrides.

Script (scripts/eval/run_vbench_e2e.py):
- Rewritten against EvalRunner; ~95 LOC, was 149.

Includes prior scaffolding from dataloader (datasets/registry.py,
the test file, the script) carried over so the refactor lands as one
coherent state.
2026-05-05 23:14:01 +00:00
shaoxiongduan 412c95b1a0 [refactor] eval: split Evaluator into EvalWorker + thin dispatcher
Mirrors FastVideo's VideoGenerator → Worker layering, in-process:

  Evaluator (user-facing dispatcher)
    └── EvalWorker × N  (single-GPU, owns metric replicas)

EvalWorker (new) holds metric replicas on one device and scores one
sample at a time. Evaluator builds num_gpus workers eagerly in __init__
(every metric loaded on every replica) and exposes a single evaluate()
method that handles both shapes:

  ev.evaluate(video=..., ...)        → one sample, runs on worker 0
  ev.evaluate([s1, s2, ...])         → fan out across workers

Removed (dead code or moved to runner in a follow-up):
- Evaluator.calibrate, _compute_chunked, _evaluate_multi_gpu,
  _gpu_metrics, evaluate_dataset, B>1 input path
- BaseMetric.batch_unit and trial_forward (kept _chunk_size as a plain
  default class attr; metrics use it for internal time-dim chunking)
- memory.is_batch_too_large, slice_sample (clear_cache stays)

Per-metric batch_unit/trial_forward overrides remain as harmless
class-attr cruft pending a mechanical removal pass.

Net: evaluator.py 447 → 135 LOC; metrics/base.py 92 → 66; memory.py
40 → 14; +78 LOC for worker.py.
2026-05-05 22:57:58 +00:00
shaoxiongduan f6275f8005 [feat] Port wm-eval into fastvideo as fastvideo.eval
Adds an in-process evaluation suite covering native (SSIM/PSNR/LPIPS/
optical_flow), audio, physics_iq, vbench (16 sub-metrics), and
videoscore2. Public API: create_evaluator/evaluate, BaseMetric +
@register, ensure_checkpoint, get_cache_dir. CLI: fastvideo eval
list/run.

VBench upstream pinned as a git submodule at
fastvideo/third_party/eval/vbench (Vchitect/VBench@45e79ec); modern-dep
compat is achieved via runtime shims in vbench/__init__.py rather than
on-disk patches. CLIP/torch.hub caches are routed through
${FASTVIDEO_EVAL_CACHE}; HF cache stays at the system default.
ensure_checkpoint delegates to fastvideo.utils.get_lock + huggingface_hub
primitives. Evaluator supports release_cuda_memory(), unload(),
reload() for training-time eval that frees GPU between calls.

Verified parity against upstream vbench (5/8 metrics bit-exact, others
within 1% drift driven by transformers/torch version skew) and against
upstream VideoScore2 (regex anchored on the actual model's output
format, soft-score formula matches upstream's argmax*max_prob/total).

Includes docs/contributing/eval-metrics.md as the porting guide.
Out of scope (deferred follow-ups): MIND, VBench-2.0, FVD as a
registered metric, training-time EvalCallback.
2026-05-01 23:17:19 +00:00
Mook 7b872cc41e [Perf] Skip bool-mask round-trip in block-sparse VSA attention (#1243) 2026-04-26 15:14:37 -07:00
alexzms 37418946c8 [docs]: clarify real_score_guidance_scale CFG parameterization (#1256) 2026-04-26 16:38:00 +08:00
William Lin 95fd29e0cb [feat] Streaming WebSocket server skeleton (single generator + fMP4) (#1251) 2026-04-26 00:33:49 -07:00
Junda Suandmergify[bot] e17cd2633c [bugfix]: normalize uint8 pil_image in I2V VAE encoding (#1249)
Co-authored-by: mergify[bot] <37929162+mergify[bot]@users.noreply.github.com>
2026-04-24 09:16:01 +00:00
William Lin e0dc5f2b0c [feat] Add typed LTX-2 continuation state and streaming session store (#1250) 2026-04-24 01:28:07 -07:00
William Lin 70ee5d230c [feat] [6/n] Improve API: LTX-2 public preset + asset wiring + gpu_pool translation (#1239) 2026-04-23 11:36:45 -07:00
William Lin 24ced500f5 [test] add LTX-2 distilled T2V SSIM regression test (#1240) 2026-04-21 12:03:38 -07:00
William Lin 4ddcdf541f [feat] [5.5/n] Improve API: streaming server config surface + serve dispatch (#1238) 2026-04-17 15:36:21 -07:00
William Lin 0e3529869c [feat] [5/n] Improve API: wire ServeConfig.default_request into OpenAI serving (#1237) 2026-04-17 13:26:18 -07:00
William Lin e1e0d91c00 [misc] small cleanup for API handling (#1235) 2026-04-16 16:21:21 -07:00
William Lin 145a3f166b [feat] [4/n] Improve API: refactor sampling param and merge with presets (#1234) 2026-04-16 14:10:02 -07:00
William Lin 88a5a933ab [feat] [3/n] Improve API: extend support to cli (#1226) 2026-04-14 15:20:47 -07:00
287 changed files with 19479 additions and 2421 deletions
+96
View File
@@ -0,0 +1,96 @@
#!/usr/bin/env bash
# Sync .agents/skills/ into .claude/skills/ via per-skill symlinks.
#
# Why: Claude Code only scans .claude/skills/ and ~/.claude/skills/ for
# user-invocable skills (no skillsPath config exists — see
# https://code.claude.com/docs/en/skills.md). This repo's skills live
# in .agents/skills/ so they travel with the repo and stay under git.
# Run this once after cloning (or after adding/removing a skill) to
# expose them to Claude Code without maintaining a parallel tree.
#
# Usage:
# .agents/scripts/sync-skills.sh
#
# Idempotent and safe to re-run. Prunes stale symlinks whose source
# has been removed from .agents/skills/. Leaves hand-written
# .claude/skills/<name>/ directories untouched (only symlinks are
# managed).
set -euo pipefail
REPO_ROOT="$(git -C "$(dirname "$0")" rev-parse --show-toplevel)"
SRC_DIR="$REPO_ROOT/.agents/skills"
DST_DIR="$REPO_ROOT/.claude/skills"
if [[ ! -d "$SRC_DIR" ]]; then
echo "Error: $SRC_DIR does not exist." >&2
exit 1
fi
mkdir -p "$DST_DIR"
linked=0
unchanged=0
skipped=0
pruned=0
link_skill() {
local name="$1"
local src="$SRC_DIR/$name"
local dst="$DST_DIR/$name"
# Relative target keeps symlinks portable across clones.
local rel="../../.agents/skills/$name"
if [[ -L "$dst" ]]; then
if [[ "$(readlink "$dst")" == "$rel" ]]; then
unchanged=$((unchanged + 1))
return
fi
rm "$dst"
elif [[ -e "$dst" ]]; then
echo "Skipped (not a symlink): .claude/skills/$name" >&2
skipped=$((skipped + 1))
return
fi
ln -s "$rel" "$dst"
echo "Linked: .claude/skills/$name -> $rel"
linked=$((linked + 1))
}
prune_stale() {
local link="$1"
local target
target="$(readlink "$link")"
case "$target" in
../../.agents/skills/*) ;;
*) return ;;
esac
local name="${target##*/}"
if [[ ! -d "$SRC_DIR/$name" ]]; then
rm "$link"
echo "Pruned stale: .claude/skills/$(basename "$link")"
pruned=$((pruned + 1))
fi
}
for src in "$SRC_DIR"/*/; do
[[ -d "$src" ]] || continue
name="$(basename "$src")"
# Only treat directories that actually contain a SKILL.md as skills.
[[ -f "$src/SKILL.md" ]] || continue
link_skill "$name"
done
shopt -s nullglob
for link in "$DST_DIR"/*; do
[[ -L "$link" ]] || continue
prune_stale "$link"
done
shopt -u nullglob
printf "\nSummary: %d linked, %d unchanged, %d pruned" "$linked" "$unchanged" "$pruned"
if [[ "$skipped" -gt 0 ]]; then
printf ", %d skipped (non-symlink collision)" "$skipped"
fi
printf "\n"
+1
View File
@@ -5,3 +5,4 @@
{"name": "evaluate-video-quality", "description": "Evaluate generated video quality using available metrics (SSIM, loss trajectory, caption consistency)", "path": "evaluate-video-quality/SKILL.md", "status": "draft", "trust": "low"}
{"name": "index-related-work", "description": "Ingest a paper or repository into the related work index", "path": "index-related-work/SKILL.md", "status": "draft", "trust": "low"}
{"name": "search-related-work", "description": "Query the related work index for relevant papers, repos, or comparisons", "path": "search-related-work/SKILL.md", "status": "draft", "trust": "low"}
{"name": "seed-ssim-references", "description": "Run a new or updated fastvideo/tests/ssim/ test on Modal, pull generated videos, and upload them to FastVideo/ssim-reference-videos so the test has a regression baseline", "path": "seed-ssim-references/SKILL.md", "status": "draft", "trust": "low"}
@@ -0,0 +1,250 @@
---
name: seed-ssim-references
description: Seed HF reference videos for a single newly-added SSIM test. Runs the test on Modal L40S, downloads the generated mp4s via `modal volume get`, pauses for the user to eyeball quality, then uploads only that test's files to `FastVideo/ssim-reference-videos`. Use when a new `fastvideo/tests/ssim/test_*_similarity.py` has just been added and has no references on HF yet.
---
# Seed SSIM Reference Videos
## Purpose
A brand-new SSIM test in `fastvideo/tests/ssim/` fails forever until its
reference videos exist on the HF dataset (`FastVideo/ssim-reference-videos`).
This skill:
1. Runs the test on Modal's L40S pool to generate the videos.
2. Downloads them to the local repo via `modal volume get`.
3. Pauses so the user can eyeball the mp4s and confirm quality.
4. Uploads only the new test's files to HF, with a guard that refuses to
overwrite anything already present.
The skill is run **manually**, once per new test. Before invoking it, the user
has already sanity-tested the new test locally — it launches `VideoGenerator`
and writes an mp4 without crashing. The skill does not re-test locally; it
goes straight to Modal L40S (which is what CI uses).
## When to use
- A new `test_*_similarity.py` file has been added in `fastvideo/tests/ssim/`
and the HF dataset has no `reference_videos/default/L40S_reference_videos/<model_id>/`
subtree for it yet.
## When not to use
- Regular CI runs — once refs exist, `pytest fastvideo/tests/ssim/` downloads
them automatically.
- Re-seeding an existing test. That requires `--force` on the upload step, and
is out of scope here; treat as a separate, deliberate operation.
## Inputs
The skill has **one required input**: the path to the new SSIM test file.
Prompt the user for it if they didn't supply it.
| Parameter | Required | Description |
|-----------|----------|-------------|
| `test_file` | Yes | e.g. `fastvideo/tests/ssim/test_ltx2_similarity.py`. The skill's first action is to ask for this if missing. |
Everything else is fixed:
- Modal runner GPU: **L40S** (hardcoded in `fastvideo/tests/modal/ssim_test.py`).
- Device folder: `L40S_reference_videos`.
- Quality tier: `default` (the tier CI runs). The `full_quality` tier is not
seeded by this skill.
- HF repo: `FastVideo/ssim-reference-videos` (dataset).
- Multi-model test files: all model ids in `*_MODEL_TO_PARAMS` are seeded
together; the Modal run produces one mp4 per (model, prompt, backend) and
the upload scopes by `--model-id`, looping if there is more than one.
## Prerequisites
The user has confirmed:
- `modal` CLI authenticated.
- `HF_API_KEY` (or `HUGGINGFACE_HUB_TOKEN` / `HF_TOKEN`) exported with write
access to `FastVideo/ssim-reference-videos`.
- The test file runs locally end-to-end (generates an mp4; SSIM assertion
failure due to missing reference is expected and fine).
Fail fast if the token env var is missing.
## Steps
### 1. Ask for the test file
If the user didn't name one, ask: *"Which SSIM test file do you want to seed
references for? (e.g. `fastvideo/tests/ssim/test_ltx2_similarity.py`)"*.
Validate:
- Path exists and matches `fastvideo/tests/ssim/test_*_similarity.py`.
- File defines a `*_MODEL_TO_PARAMS` dict — grep it to extract the set of
model ids. Those ids drive step 5.
If either check fails, stop and tell the user what's wrong.
### 2. Run the test on Modal L40S
Pick a subdir name so repeated runs don't collide:
```bash
SHORT_COMMIT=$(git rev-parse --short=12 HEAD)
TIMESTAMP=$(date -u +%Y%m%d_%H%M%S)
SUBDIR="${TIMESTAMP}_${SHORT_COMMIT}"
```
Then launch the Modal run:
```bash
modal run fastvideo/tests/modal/ssim_test.py \
--git-repo="$(git config --get remote.origin.url)" \
--git-commit="$(git rev-parse HEAD)" \
--hf-api-key="$HF_API_KEY" \
--test-files="<test_file>" \
--sync-generated-to-volume \
--generated-volume-subdir="$SUBDIR" \
--skip-reference-download \
--no-fail-fast
```
Flag rationale:
- `--skip-reference-download`: no refs exist yet, so conftest must not try to
pull them.
- `--no-fail-fast`: lets the test finish generation before `_assert_similarity`
raises `FileNotFoundError: Reference video folder does not exist`. The
expected failure is what we want — the mp4 has already been written.
- `--sync-generated-to-volume` + `--generated-volume-subdir`: copies the
generated mp4s to the `hf-model-weights` Modal volume under
`ssim_generated_videos/default/<SUBDIR>/generated_videos/` so we can pull
them locally.
The Modal run will end with a nonzero exit (expected) and print a
`modal volume get hf-model-weights ssim_generated_videos/default/<SUBDIR>/generated_videos ./generated_videos_modal/default`
command. Capture that `<SUBDIR>` — you need it for step 3.
### 3. Download generated videos locally
```bash
modal volume get --force hf-model-weights \
ssim_generated_videos/default/"$SUBDIR"/generated_videos \
./generated_videos_modal/default
```
`--force` is required when the parent `./generated_videos_modal/default`
already exists; without it, `modal volume get` errors with `[Errno 21] Is a
directory`. Safe to pass on the first run too.
After this, the mp4s live at
`./generated_videos_modal/default/generated_videos/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
The extra `generated_videos/` level comes from the volume layout in
`_sync_generated_videos_to_volume` (`ssim_test.py`) — the command copies
`<repo>/fastvideo/tests/ssim/generated_videos/<tier>` to
`ssim_generated_videos/<tier>/<SUBDIR>/generated_videos/`, and `modal volume
get` preserves that trailing `generated_videos/` segment.
### 4. PAUSE — user reviews quality
Print the list of downloaded mp4s and their paths, then stop. Tell the user:
> "Generated videos downloaded to `./generated_videos_modal/default/generated_videos/L40S_reference_videos/`. Please open them and confirm the quality looks correct. Reply **`upload`** to continue, or anything else to abort."
Do not proceed until the user explicitly says `upload`. If they abort, leave
everything on disk so they can inspect further — no cleanup.
### 5. Copy into the local reference layout
Scoped copy — only the new test's mp4s. Loop over each `<model_id>` extracted
in step 1:
```bash
python fastvideo/tests/ssim/reference_videos_cli.py copy-local \
--quality-tier default \
--device-folder L40S_reference_videos \
--generated-dir ./generated_videos_modal/default/generated_videos/L40S_reference_videos
```
(The `--generated-dir` points at the device-folder root inside the
downloaded tree; `copy-local` walks all `<model>/<backend>/*.mp4`
underneath it. Since the Modal run was scoped to a single test file via
`--test-files`, only that test's model(s) are present — so the copy is
implicitly per-test.)
Result: `fastvideo/tests/ssim/reference_videos/default/L40S_reference_videos/<model_id>/<backend>/<prompt>.mp4`.
### 6. Upload to HF — scoped per model_id, with overwrite guard
For each `<model_id>`:
```bash
python fastvideo/tests/ssim/reference_videos_cli.py upload \
--quality-tier default \
--device-folder L40S_reference_videos \
--model-id "<model_id>"
```
The upload command:
- Uploads **only** `reference_videos/default/L40S_reference_videos/<model_id>/`.
- **Refuses** if any file already exists at that path on HF (this is the
guard — seeding a new test should never clobber existing refs). To override,
the user must re-run with `--force`. If the guard fires, stop and report
exactly which files exist; do not silently `--force`.
Reads the HF token from `HF_API_KEY` / `HUGGINGFACE_HUB_TOKEN` / `HF_TOKEN`.
### 7. Report success
List what was uploaded (paths in repo) and remind the user to push any
related code changes. Do **not** auto-verify by re-running Modal — the user
can run `pytest fastvideo/tests/ssim/<test_file>` later to confirm end-to-end;
it will auto-download the refs they just uploaded.
## Failure modes and how to handle them
- **`HF_API_KEY` unset.** Stop before step 2. The Modal run needs it (passed
via `--hf-api-key`), and step 6 needs it for upload.
- **Modal run fails before generation.** No mp4s on the volume — nothing to
download. Fix the test locally (`pytest fastvideo/tests/ssim/<test_file>`)
and retry from step 2.
- **`./generated_videos_modal/default/L40S_reference_videos/` missing after
`modal volume get`.** The run didn't produce videos (most likely the test
crashed before writing, or `REQUIRED_GPUS` exceeded the partition capacity
— see Modal logs).
- **Upload guard fires (files already exist).** The test name / model id
collides with something already on HF. Verify the user actually wants to
replace existing refs; if so, re-run the upload with `--force`. If not,
rename the model id in `*_MODEL_TO_PARAMS` and re-seed.
- **Quality looks wrong in step 4.** Abort. The mp4s stay on disk for
inspection. The fix is usually in the test's params (resolution, steps,
seed) — edit the test, then re-run the skill.
## Design notes (for future skill maintainers)
- The skill deliberately runs on Modal, **not** locally, because the CI
runner is L40S. Seeding from a different GPU SKU produces refs that CI's
L40S runs can't match (SSIM drifts across SKUs).
- The skill is default-tier only. `full_quality` refs are seeded by a
separate, deliberate operation — they double runtime and aren't what CI
gates on.
- The overwrite guard in `reference_videos_cli.py upload` is default-on
specifically because this skill exists. Re-seeding is a distinct operation
that requires explicit `--force`.
## References
- `fastvideo/tests/modal/ssim_test.py` — Modal orchestrator; see
`--sync-generated-to-volume`, `--generated-volume-subdir`,
`--skip-reference-download`, `--no-fail-fast`.
- `fastvideo/tests/ssim/reference_videos_cli.py` — `copy-local`, `upload`
(with `--model-id`, `--force`), `download`, `ensure` subcommands.
- `fastvideo/tests/ssim/README.md` — reference layout, HF repo conventions.
- `fastvideo/tests/ssim/inference_similarity_utils.py` —
`run_text_to_video_similarity_test` + `_build_init_kwargs`: what each test
config passes to `VideoGenerator.from_pretrained`.
## Changelog
| Date | Change |
|------|--------|
| 2026-04-17 | Initial version (Modal sync-to-volume flow). |
| 2026-04-21 | Rewrite: single-test scope, explicit user-review pause, per-`model_id` upload, HF overwrite guard. Dropped `scripts/seed_ssim.sh`. |
| 2026-04-21 | Post-first-run fixes: `modal volume get` needs `--force` when parent exists; download tree has an extra `generated_videos/` level so `--generated-dir` must reflect it. |
+51 -12
View File
@@ -4,21 +4,24 @@ description: How to develop, validate, and register a new evaluation metric
# Evaluation Development SOP
Standard procedure for adding new video quality evaluation metrics to the
FastVideo agent toolkit.
Standard procedure for adding new video quality evaluation metrics to
the FastVideo agent toolkit.
## When to Use
## When to use
- You need a metric that doesn't exist in `.agents/memory/evaluation-registry/README.md`.
- You need a metric that does not exist in
`.agents/memory/evaluation-registry/README.md`.
- An existing metric needs significant changes to its methodology.
- You're exploring a new evaluation approach.
- You are exploring a new evaluation approach.
## Steps
### 1. Research
- Search `.agents/memory/related-work/` for existing evaluation approaches.
- Check the `evaluation_registry.md` for current metrics and their limitations.
- Search `.agents/memory/related-work/` for existing evaluation
approaches.
- Check `.agents/memory/evaluation-registry/README.md` for current
metrics and their limitations.
- Review literature: FVD, CLIP-Score, human preference, etc.
### 2. Prototype
@@ -29,21 +32,25 @@ FastVideo agent toolkit.
### 3. Validate
- **Known-good test**: Metric should score high on reference-quality videos.
- **Known-bad test**: Metric should score low on degraded/unrelated videos.
- **Sensitivity test**: Small quality differences should produce meaningful
score differences.
- **Known-good test**: metric should score high on reference-quality
videos.
- **Known-bad test**: metric should score low on degraded or unrelated
videos.
- **Sensitivity test**: small quality differences should produce
meaningful score differences.
- Document thresholds and their justification.
### 4. Register
Update `.agents/memory/evaluation-registry/README.md`:
- Add the metric with status `Active`.
- Document location, thresholds, and trust level.
### 5. Integrate
Update `.agents/skills/evaluate-video-quality.md`:
Update `.agents/skills/evaluate-video-quality/SKILL.md`:
- Add the new metric as a section.
- Include code examples and interpretation guide.
@@ -52,3 +59,35 @@ Update `.agents/skills/evaluate-video-quality.md`:
- Move the exploration log content into the skill.
- Clean up the exploration file or mark it as `promoted`.
- If anything went wrong during development, create a lesson.
## Where the metrics live
The eval suite is `fastvideo/eval/`. New metrics register themselves
via `@register("<group>.<name>")` and are auto-discovered when
`fastvideo.eval.metrics` is imported.
- **Native metrics** (SSIM, PSNR, LPIPS, optical flow, VLM): add a
file under the appropriate group dir
(`fastvideo/eval/metrics/common/`, `optical_flow/`, `videoscore2/`,
`physics_iq/`).
- **Metrics that wrap upstream research code**: follow the vbench
pattern in `fastvideo/eval/metrics/vbench/`. The contract is:
- Upstream lives as a git submodule under
`fastvideo/third_party/eval/<bench>/`, pinned to a SHA in repo-root
`.gitmodules`.
- The metric package's `__init__.py` inserts the submodule path on
`sys.path` and installs runtime compat shims (attribute-level
monkey-patches) for any modern-dep drift. Do not modify upstream
files on disk, and do not ship a `setup.sh`.
- See `fastvideo/eval/README.md` for the worked vbench example.
- Full porting guide:
[`docs/contributing/eval-metrics.md`](../../docs/contributing/eval-metrics.md).
## Out of scope of the initial eval port
The following land in follow-up PRs:
- **MIND** metrics (depends on a separate `vipe` submodule).
- **VBench-2.0** sibling package.
- Native conversion of **FVD** under `fastvideo/eval/metrics/fvd/`.
- The training-time `EvalCallback`.
+1 -1
View File
@@ -105,7 +105,7 @@ pull_request_rules:
- files~=^fastvideo/pipelines/samplers/
- files~=^fastvideo/entrypoints/
- files~=^fastvideo/worker/
- files~=^fastvideo/configs/sample/
- files~=^fastvideo/api/sampling_param
- files~=^fastvideo/configs/pipelines/
- files~=^examples/inference/
- -closed
+3
View File
@@ -4,3 +4,6 @@
[submodule "fastvideo-kernel/include/cutlass"]
path = fastvideo-kernel/include/cutlass
url = https://github.com/NVIDIA/cutlass.git
[submodule "fastvideo/third_party/eval/vbench"]
path = fastvideo/third_party/eval/vbench
url = https://github.com/Vchitect/VBench.git
+2 -2
View File
@@ -62,9 +62,9 @@ This page contains the complete API reference for the FastVideo library.
show_root_toc_entry: true
heading_level: 4
#### fastvideo.configs.sample
#### fastvideo.api.sampling_param
::: fastvideo.configs.sample
::: fastvideo.api.sampling_param
options:
show_source: true
show_root_heading: true
+1 -1
View File
@@ -173,7 +173,7 @@ Applied by Mergify based on which paths you modified. Multiple scope labels can
| Label | File paths that trigger it |
|-------|---------------------------|
| `scope: training` | `fastvideo/train/`, `fastvideo/training/`, `fastvideo/distillation/`, `examples/train/`, `examples/training/`, `examples/distill/` |
| `scope: inference` | `fastvideo/pipelines/basic/`, `fastvideo/pipelines/stages/`, `fastvideo/pipelines/samplers/`, `fastvideo/entrypoints/`, `fastvideo/worker/`, `fastvideo/configs/sample/`, `fastvideo/configs/pipelines/`, `examples/inference/` |
| `scope: inference` | `fastvideo/pipelines/basic/`, `fastvideo/pipelines/stages/`, `fastvideo/pipelines/samplers/`, `fastvideo/entrypoints/`, `fastvideo/worker/`, `fastvideo/api/sampling_param.py`, `fastvideo/configs/pipelines/`, `examples/inference/` |
| `scope: attention` | `fastvideo/attention/` |
| `scope: kernel` | `fastvideo-kernel/`, `csrc/` |
| `scope: data` | `fastvideo/dataset/`, `fastvideo/pipelines/preprocess/`, `examples/preprocessing/` |
+5 -4
View File
@@ -44,7 +44,7 @@ FastVideo maps a Diffusers-style repo into a pipeline like:
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
weight name translation.
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
- `fastvideo/api/sampling_param.py`: runtime sampling parameters.
- `fastvideo/pipelines/basic/*`: end-to-end pipeline logic built from stages.
- `model_index.json`: the HF repo entrypoint that maps component names to
classes and weight files.
@@ -55,7 +55,7 @@ Minimal usage example (based on `examples/inference/basic/basic.py`):
```python
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
@@ -319,7 +319,8 @@ Purpose:
- `fastvideo/configs/pipelines/` describes pipeline wiring and model module
names.
- `fastvideo/configs/sample/` defines default runtime parameters.
- `fastvideo/api/sampling_param.py` defines runtime sampling parameters.
Defaults come from profiles in `fastvideo/pipelines/basic/<family>/profiles.py`.
Action:
@@ -474,7 +475,7 @@ FastVideo integration.
3. Pipeline wiring.
- Pipeline: `fastvideo/pipelines/basic/wan/wan_pipeline.py`
- Pipeline config: `fastvideo/configs/pipelines/wan.py`
- Sampling defaults: `fastvideo/configs/sample/wan.py`
- Sampling defaults: `fastvideo/pipelines/basic/wan/profiles.py`
4. Minimal example.
- Script: `examples/inference/basic/basic.py`
+559
View File
@@ -0,0 +1,559 @@
# Porting Eval Metrics into `fastvideo.eval`
This guide is for contributors adding new evaluation metrics to
FastVideo's eval suite. To run the existing metrics, see
[`fastvideo/eval/README.md`](../../fastvideo/eval/README.md).
## When to use this guide
Use this guide when you are:
- Adding a new metric (native or wrapping a third-party library).
- Porting a benchmark (e.g. VBench, MIND, EvalCrafter) whose Python
code needs to be importable from a pinned upstream.
- Adding a new metric group (audio, vlm, etc.).
## TL;DR
Metrics are auto-discovered from
`fastvideo/eval/metrics/<group>/<name>/metric.py`. Each declares itself
with `@register("<group>.<name>")` and subclasses `BaseMetric`. Three
recipes:
1. **Native metric** (pure-PyTorch, no submodule). Drop a file,
declare deps, implement `compute(sample)`.
2. **Library-wrapped metric** (CLIP, torch.hub, transformers, pyiqa).
Same as above, plus route the library's cache through
`get_cache_dir()` if it has a `download_root=` / `cache_dir=`
kwarg.
3. **Upstream-submodule-wrapped metric** (vbench-style). Pin upstream
as a git submodule under `fastvideo/third_party/eval/<bench>/`. The
adapter `__init__.py` does the `sys.path` insert and any runtime
compat shims for modern dep versions. Patches live as Python in
that file rather than as on-disk patches to the submodule.
The full recipes are below.
---
## 0) Layout and auto-discovery
```
fastvideo/eval/metrics/
├── base.py # BaseMetric + lifecycle contract
├── common/ # group: SSIM, PSNR, LPIPS
├── optical_flow/ # group: gt_optical_flow, synthetic_optical_flow
├── vlm/ # group: VideoScore-2
├── physics_iq/ # group + sub-metrics
└── vbench/ # group: 16 sub-metrics
├── __init__.py # sys.path bootstrap + runtime compat shims
├── _grit_helper.py # shared upstream-touching helpers
└── <sub_metric>/metric.py
```
Auto-discovery (`fastvideo/eval/metrics/__init__.py`) walks each group
dir and imports every `metric.py` it finds, which fires the
`@register` decorators. Names starting with `_` are skipped. Use that
prefix for shared helpers or vendored code that should not register
itself.
---
## 1) The `BaseMetric` contract
Every metric subclasses `fastvideo.eval.metrics.base.BaseMetric` and
declares:
```python
class YourMetric(BaseMetric):
name: str = "common.your_metric" # must match @register
requires_reference: bool = True # needs sample["reference"]
higher_is_better: bool = True # for ranking / aggregates
dependencies: list[str] = [] # importable module names;
# registry surfaces a clean
# ImportError if missing
needs_gpu: bool = False
backbone: str | None = None # e.g. "clip_vit_l14"
```
You must implement:
```python
def compute(self, sample: dict) -> list[MetricResult]:
"""sample['video'] is (1, T, C, H, W). Return a one-element list.
The leading 1 is preserved for forward-compat with batched eval;
today :class:`EvalWorker` always invokes metrics with B=1.
"""
```
You may override:
- `setup(self) -> None`. Eager model loading. Called once by
`create_evaluator`. Idempotent (re-entrant). Use the `if self._model
is not None: return` pattern.
- `to(self, device)`. Move the metric and its submodels to `device`.
If a required input is missing (e.g. an fps-aware metric called
without `fps`), return `self._skip(sample, reason)` instead of
raising.
---
## 2) Recipe A: native metric (no external deps)
Smallest case. Pixel math, simple closed-form.
```python
# fastvideo/eval/metrics/common/your_metric/metric.py
from __future__ import annotations
import torch
from fastvideo.eval.metrics.base import BaseMetric
from fastvideo.eval.registry import register
from fastvideo.eval.types import MetricResult
@register("common.your_metric")
class YourMetric(BaseMetric):
name = "common.your_metric"
requires_reference = True
higher_is_better = True
needs_gpu = False
dependencies: list[str] = [] # nothing extra
def compute(self, sample: dict) -> list[MetricResult]:
gen, ref = sample["video"], sample["reference"] # (B,T,C,H,W) each
per_video = ((gen - ref) ** 2).mean(dim=(1, 2, 3, 4)).sqrt()
return [
MetricResult(name=self.name, score=float(s), details={})
for s in per_video
]
```
That is the whole recipe. Drop the file and the registry picks it up.
---
## 3) Recipe B: library-wrapped metric (CLIP, torch.hub, transformers, pyiqa)
If your metric loads a backbone from a Python package, route the
library at the eval cache so users get one knob (`FASTVIDEO_EVAL_CACHE`)
to redirect everything.
### Cache routing rules
| Library | How to route | Location after redirect |
|---|---|---|
| `clip.load("ViT-X")` | pass `download_root=str(get_cache_dir() / "clip")` | `${FASTVIDEO_EVAL_CACHE}/clip/` |
| `torch.hub.load(...)` | nothing; `TORCH_HOME` is redirected at `fastvideo.eval` import time | `${FASTVIDEO_EVAL_CACHE}/torch/hub/` |
| `transformers.from_pretrained(...)` | nothing; leave HF's default cache (`~/.cache/huggingface/hub/`) so users dedupe with other ML projects | `~/.cache/huggingface/hub/` |
| `huggingface_hub.snapshot_download` / `hf_hub_download` | use `ensure_checkpoint(...)` (it wraps these with filelock) | same as above |
| `pyiqa.create_metric(...)` | no env var or kwarg honored; document in metric docstring | pyiqa-internal |
| `lpips`, `ptlflow` | torch.hub-based, auto-redirected | `${FASTVIDEO_EVAL_CACHE}/torch/hub/` |
| Raw URL (no HF Hub) | use `ensure_checkpoint(name, source="https://...")` | `${FASTVIDEO_EVAL_CACHE}/models/<name>` |
| Dataset asset (raw video/mask/image) auto-fetched from a public bucket | download into `get_cache_dir() / "datasets" / "<bench>"`, mirroring upstream's relative layout. Vendor any small manifest (CSV/JSON ≤1 MB) under the metric folder so the dataset can be used without external setup. | `${FASTVIDEO_EVAL_CACHE}/datasets/<bench>/` |
### Dataset assets: vendor the manifest, auto-fetch the rest
If your metric ships with its own paired-reference dataset (Physics-IQ
is the canonical example), follow this layout:
- **Manifest** (CSV/JSON ≤1 MB): vendor it under
`fastvideo/eval/metrics/<bench>/_vendored/<manifest>.<ext>`, with a
sibling `_vendored/LICENSE` recording attribution and provenance.
The `_vendored/` subdir is the project-wide convention for
upstream-provenance files: it is auto-skipped by metric discovery
(the `_` prefix) and by codespell (one `*/_vendored/*` glob in
`[tool.codespell].skip`), so dropping in a new vendored file
requires no further config. Read the manifest from the dataset
module via a `Path(__file__)`-relative resolver. Mirror
`_VENDORED_DESCRIPTIONS_CSV` in
`fastvideo/eval/datasets/physics_iq.py`.
- **Heavy assets** (videos, masks, images): do not vendor. Auto-fetch
on first miss into `get_cache_dir() / "datasets" / "<bench>"`,
mirroring upstream's relative directory layout one-for-one so a
pre-downloaded mirror at any path works as a drop-in
`dataset_root=`. Use atomic `.part` then final-rename to be safe
under concurrent SLURM ranks.
- **Bucket override**: expose `FASTVIDEO_<BENCH>_BUCKET_URL` so users
with internal mirrors can redirect.
- **Opt-out**: accept `auto_download: bool = True` in the dataset
constructor; on `False`, raise `FileNotFoundError` instead of
fetching. This covers air-gapped runs and CI.
The end-state is `get_dataset("<bench>")` with no kwargs.
### Example: CLIP backbone + LAION head
```python
# fastvideo/eval/metrics/your_group/your_metric/metric.py
from __future__ import annotations
import torch
import torch.nn as nn
from fastvideo.eval.metrics.base import BaseMetric
from fastvideo.eval.registry import register
from fastvideo.eval.types import MetricResult
@register("your_group.your_metric")
class YourMetric(BaseMetric):
name = "your_group.your_metric"
requires_reference = False
needs_gpu = True
dependencies = ["clip"] # "openai-clip" PyPI; importable as `clip`
def __init__(self) -> None:
super().__init__()
self._clip = None
self._head = None
def setup(self) -> None:
if self._clip is not None:
return
import clip
from fastvideo.eval.models import ensure_checkpoint, get_cache_dir
# Backbone: route CLIP's cache through our root.
self._clip, _ = clip.load(
"ViT-L/14",
device=self.device,
download_root=str(get_cache_dir() / "clip"),
)
self._clip.eval()
# URL-fetched head: ensure_checkpoint downloads to
# ${FASTVIDEO_EVAL_CACHE}/models/ with filelock + atomic rename.
ckpt = ensure_checkpoint(
"your_head.pth",
source="https://example.com/path/to/your_head.pth",
)
self._head = nn.Linear(768, 1)
self._head.load_state_dict(
torch.load(ckpt, map_location="cpu", weights_only=True)
)
self._head.to(self.device).eval()
def to(self, device):
super().to(device)
if self._clip is not None:
self._clip = self._clip.to(self.device)
if self._head is not None:
self._head = self._head.to(self.device)
return self
def compute(self, sample: dict) -> list[MetricResult]:
...
```
### Do not redirect other `~/.cache/...` dirs
If a third-party library hard-codes `~/.cache/<lib>/` and offers no
override, document the exception in the metric's docstring. Forcing
redirection by setting `os.environ` or patching `os.path.expanduser`
is fragile and breaks user expectations of where the library's cache
lives.
---
## 4) Recipe C: upstream-submodule-wrapped metric (vbench pattern)
Use this when the upstream benchmark ships Python code (`vbench/`,
`MIND/`, etc.) that is not pip-installable cleanly. See
`fastvideo/eval/metrics/vbench/__init__.py` for the worked example.
### 4.1 Pin the upstream as a submodule
```bash
git submodule add <upstream-url> fastvideo/third_party/eval/<bench>
cd fastvideo/third_party/eval/<bench>
git checkout <pinned-sha>
cd -
git add .gitmodules fastvideo/third_party/eval/<bench>
```
The submodule pulls under the standard `git submodule update --init
--recursive` flow that users already run for kernel deps.
### 4.2 Bootstrap on `sys.path`
```python
# fastvideo/eval/metrics/<bench>/__init__.py
from __future__ import annotations
import sys
from pathlib import Path
# fastvideo/eval/metrics/<bench>/__init__.py → ../../../../third_party/eval/<bench>
_UPSTREAM = Path(__file__).resolve().parents[3] / "third_party" / "eval" / "<bench>"
if _UPSTREAM.is_dir() and str(_UPSTREAM) not in sys.path:
sys.path.insert(0, str(_UPSTREAM))
```
We do not `pip install` the upstream because its egg-link/.pth would
just re-do this `sys.path.insert`, and skipping the install also skips
the upstream's `setup.py` (which often gates on a specific CUDA
version).
### 4.3 Modern-dep compat: runtime shims
Upstream code pinned to e.g. `transformers==4.33.2`, `numpy<2`
typically breaks against modern versions in 3-4 known places (API
renames). Fix those at import time, in the same `__init__.py`:
```python
def _install_compat_shims() -> None:
# Example: transformers.modeling_utils API moved.
try:
import transformers.modeling_utils as _mu
import transformers.pytorch_utils as _pu
for _n in ("apply_chunking_to_forward",
"find_pruneable_heads_and_indices",
"prune_linear_layer"):
if not hasattr(_mu, _n) and hasattr(_pu, _n):
setattr(_mu, _n, getattr(_pu, _n))
except ImportError:
pass
# Example: numpy.lib.function_base.disp removed in numpy>=2.
try:
import types, numpy.lib as _nl
if not hasattr(_nl, "function_base"):
_stub = types.ModuleType("numpy.lib.function_base")
_stub.disp = lambda *a, **k: None
sys.modules["numpy.lib.function_base"] = _stub
_nl.function_base = _stub
except ImportError:
pass
_install_compat_shims()
```
For function-level patches that cannot be expressed as attribute
writes (e.g. wrapping a model factory function), use a
`sys.meta_path` finder that wraps the loader. See
`_install_modeling_finetune_hook()` in
`fastvideo/eval/metrics/vbench/__init__.py` for the pattern (about 30
lines).
Why shims rather than `git apply` patches: patches go stale when the
upstream SHA changes; shims are versioned Python code in our repo,
they are grep-able, and they only run if the targeted module is
imported.
### 4.4 Per-sub-metric files
Each sub-metric is a normal `BaseMetric` subclass that imports from
the upstream:
```python
# fastvideo/eval/metrics/<bench>/<sub>/metric.py
from fastvideo.eval.metrics.base import BaseMetric
from fastvideo.eval.registry import register
from fastvideo.eval.types import MetricResult
@register("<bench>.<sub>")
class YourSubMetric(BaseMetric):
...
def setup(self) -> None:
if self._model is not None:
return
# The sys.path bootstrap fired when fastvideo.eval.metrics.<bench>
# was imported (which auto-discovery does before importing this
# sub-package). Upstream imports just work:
from <bench>.something import SomeModel
...
```
### 4.5 Conditional registration when the upstream is missing
If a user installed `fastvideo[eval]` but did not run `git submodule
update --init`, `<bench>.*` metrics should not register. The
auto-discovery walker imports each sub-package's `metric` module; have
that import bail out cleanly:
```python
# fastvideo/eval/metrics/<bench>/__init__.py — at the bottom
_AVAILABLE = (_UPSTREAM / "<bench>" / "__init__.py").is_file()
```
```python
# fastvideo/eval/metrics/<bench>/<sub>/__init__.py
from fastvideo.eval.metrics.<bench> import _AVAILABLE
if _AVAILABLE:
from .metric import YourSubMetric # noqa
```
`fastvideo eval list` then reflects what the user actually has rather
than what they could have.
### 4.6 Leave the upstream alone unless it blocks the metric
The upstream is pinned. If a metric works against the pinned SHA,
leave the upstream files untouched. If it actively breaks against
modern deps (the import-drift cases above), shim it. Avoid
fastvideo-side forks of upstream code; they make patches go stale and
parity drift.
---
## 5) Model checkpoints: `ensure_checkpoint`
Use `ensure_checkpoint(name, source, filename=None)` for any
non-package weights. It resolves a local path, downloading on miss,
with filelock safety across processes and SLURM ranks.
| `source` form | What happens |
|---|---|
| `"/abs/path/to/file.pth"` | passthrough, returned unchanged |
| `"https://..."` | downloaded to `${FASTVIDEO_EVAL_CACHE}/models/<name>` via `huggingface_hub.http_get`, atomic rename, filelock |
| `"org/repo"` (no `filename`) | `snapshot_download(repo_id)` → `~/.cache/huggingface/hub/` |
| `"org/repo"` (with `filename`) | `hf_hub_download(repo_id, filename)` → `~/.cache/huggingface/hub/` |
`name` is only used as the local filename for URL sources. HF sources
ignore it (HF manages its own cache key by content hash).
```python
from fastvideo.eval.models import ensure_checkpoint
# URL: name matters
ckpt = ensure_checkpoint(
"amt-s.pth",
source="https://huggingface.co/lalala125/AMT/resolve/main/amt-s.pth",
)
# HF single file: name is decorative
ckpt = ensure_checkpoint(
"raft-things.pth", # ignored; HF cache uses repo+sha
source="OpenGVLab/VBench_Used_Models",
filename="raft-things.pth",
)
```
---
## 6) Declaring `dependencies`
Set `dependencies = ["pkg1", "pkg2"]` on your metric class with
importable module names (not PyPI distribution names). The registry
checks each via `importlib.util.find_spec` at instantiation time and
raises a clean `ImportError` pointing the user at the right install
extra:
```python
class YourMetric(BaseMetric):
dependencies = ["clip", "timm"] # importable as `import clip`, `import timm`
```
If a dep is in `[project.optional-dependencies.eval-<group>]`, you do
not need to do anything more. If it is a new dep, add it to that
group in `pyproject.toml`.
---
## 7) Common gotchas
- The standard `git submodule update --init --recursive` is enough
for a benchmark; do not write a `setup.sh`. Modern-dep compat goes
into your `__init__.py` as runtime shims.
- Do not modify upstream files on disk. The submodule should always
match its pinned SHA. Compat lives in our `__init__.py`.
- Do not pip-install the upstream. The egg-link is a glorified
`sys.path.insert`, which we do directly in `__init__.py`.
- Do not call `torch.hub.set_dir(...)` from your metric. It is done
globally in `fastvideo/eval/__init__.py`.
- Do not put cache-redirection env vars in your metric's `setup()`.
By the time `setup()` runs, the library has likely already cached
the default-location decision. Set env vars at package-init time.
- Skip rather than raise when an input is missing. Use
`self._skip(sample, reason)` for any expected-missing input. It
returns a list of `MetricResult(score=None)` so other metrics in
the same evaluator continue.
- Watch for upstream re-registration conflicts. If the upstream uses
a global registry (detectron2's `META_ARCH_REGISTRY`, MMCV, etc.),
loading the same model twice in the same process will throw. The
evaluator already loads each metric once; if you write a custom
setup-then-call pattern, mirror that single-load discipline.
---
## 8) Training-time eval: keep evaluators hot, free caches between calls
When wiring eval into a training loop, the working pattern is:
1. Construct the `Evaluator` once and attach it to the pipeline
(`self._eval = create_evaluator(...)`). Do not recreate it per
validation round; that re-pays the model load cost.
2. Save validation videos to disk (the diffusion path already does
this). Pass paths to `evaluator.evaluate`, not in-memory tensors
that share GPU memory with the training model.
3. Run validation only on rank 0 of each sequence-parallel group.
Gather paths from other ranks and let rank 0 score everything.
4. After every `evaluate(...)` call, call
`evaluator.release_cuda_memory()` in a `finally` block. That runs
`gc.collect()` + `torch.cuda.empty_cache()` +
`torch.cuda.ipc_collect()`. The eval model stays loaded; only
transient activation buffers from the just-finished call get
freed:
```python
for video_path in batch:
try:
scores = self._eval.evaluate(video=load_video(video_path))
finally:
self._eval.release_cuda_memory()
```
5. If memory pressure spikes (rare on H200), call
`evaluator.unload()` to drop every metric reference and let the
GPU memory be GC'd. `unload` is reversible:
`evaluator.reload()` rebuilds the same metrics with the original
config (re-paying the model load cost). Calling `evaluate`
between `unload` and `reload` raises a clear `RuntimeError`.
For most metrics (sub-1 GB backbones, e.g. CLIP/DINO/RAFT/AMT) the
eval model can stay co-resident with the training model in
`transformer.eval()` mode without any swap. For larger ones
(VideoScore2 at 14 GB), measure first; if it fits on the rank-0 GPU
during validation (training model in eval mode means no
grads/optimizer updates), keep it hot. If not, `unload` between
rounds.
## 9) Local verification
Native and library-wrapped metrics: a single-GPU smoke is enough.
```python
import torch
from fastvideo.eval import create_evaluator
ev = create_evaluator(metrics=["<group>.<your_metric>"], device="cuda")
video = torch.randn(1, 49, 3, 256, 256, device="cuda").clamp(0, 1)
print(ev.evaluate(video=video))
```
Submodule-wrapped metrics: also do a parity check against the
upstream once. Clone upstream into a separate venv, run the same
video through both, and expect an exact match on bit-deterministic
metrics and ≤1% drift on backbone-heavy ones (driven by
transformers/torch version differences).
For quick parity in CI: pin a tiny test video, record expected
scores ± tolerance, and add a calibration test under
`fastvideo/tests/eval/`.
---
## 10) When not to add a metric
- **Set-vs-set distribution metrics** (FVD, FID-style) do not fit
`BaseMetric.compute(sample)` cleanly; they need a population.
Adding them requires a stateful accumulator interface that does
not exist yet. Open an issue first.
- **Metrics requiring a single-GPU model larger than available
memory.** Eval is not the place for tensor-parallel sharding;
metrics are expected to fit on one GPU.
- **Metrics that need `mmcv` with a conflicting CUDA ABI.** Document
the affected sub-metrics as unsupported and skip them. Building
isolation infrastructure (subprocess engine, per-metric venv) is
out of scope.
@@ -1,7 +1,7 @@
status_definitions:
kept: "Public field remains on a public adapter surface with the same meaning."
moved: "Public field remains supported but normalizes into a different nested path."
profile_owned: "Public field remains supported only through a model/profile-specific surface."
preset_owned: "Public field remains supported only through a model/preset-specific surface."
compatibility_only: "Legacy public field remains adapter-only during migration and is not part of the canonical typed schema."
private_only: "Field should only be handled by private adapters and is not a public FastVideo compatibility promise."
internal_only: "Field is runtime/config plumbing and should not be part of the new public typed inference API."
@@ -29,7 +29,7 @@ surfaces:
vae_cpu_offload: generator.engine.offload.vae
pin_cpu_memory: generator.engine.offload.pin_cpu_memory
enable_torch_compile: generator.engine.compile.enabled
torch_compile_kwargs: generator.engine.compile.kwargs
torch_compile_kwargs: generator.engine.compile.backend,fullgraph,mode,dynamic,extras
disable_autocast: generator.engine.disable_autocast
enable_stage_verification: generator.engine.enable_stage_verification
prompt_txt: request.inputs.prompt_path
@@ -40,12 +40,12 @@ surfaces:
init_weights_from_safetensors_2: generator.pipeline.components.transformer_2_weights
override_pipeline_cls_name: generator.pipeline.components.override_pipeline_cls_name
boundary_ratio: request.sampling.boundary_ratio
profile_owned:
ltx2_vae_tiling: generator.pipeline.profile_overrides.ltx2.vae_tiling
ltx2_vae_spatial_tile_size_in_pixels: generator.pipeline.profile_overrides.ltx2.vae.spatial_tile_size_in_pixels
ltx2_vae_spatial_tile_overlap_in_pixels: generator.pipeline.profile_overrides.ltx2.vae.spatial_tile_overlap_in_pixels
ltx2_vae_temporal_tile_size_in_frames: generator.pipeline.profile_overrides.ltx2.vae.temporal_tile_size_in_frames
ltx2_vae_temporal_tile_overlap_in_frames: generator.pipeline.profile_overrides.ltx2.vae.temporal_tile_overlap_in_frames
ltx2_vae_tiling: generator.pipeline.vae_tiling
preset_owned:
ltx2_vae_spatial_tile_size_in_pixels: generator.pipeline.preset_overrides.ltx2.vae.spatial_tile_size_in_pixels
ltx2_vae_spatial_tile_overlap_in_pixels: generator.pipeline.preset_overrides.ltx2.vae.spatial_tile_overlap_in_pixels
ltx2_vae_temporal_tile_size_in_frames: generator.pipeline.preset_overrides.ltx2.vae.temporal_tile_size_in_frames
ltx2_vae_temporal_tile_overlap_in_frames: generator.pipeline.preset_overrides.ltx2.vae.temporal_tile_overlap_in_frames
ltx2_initial_latent_path: request.extensions.ltx2.initial_latent_path
compatibility_only:
mode: "Legacy multi-mode FastVideoArgs switch; typed inference config should not expose execution mode."
@@ -69,16 +69,16 @@ surfaces:
pipeline_config_base:
moved:
pipeline_config_path: generator.pipeline.components.pipeline_config_path
profile_owned:
embedded_cfg_scale: generator.pipeline.profile_overrides.embedded_cfg_scale
flow_shift: generator.pipeline.profile_overrides.flow_shift
flow_shift_sr: generator.pipeline.profile_overrides.flow_shift_sr
is_causal: generator.pipeline.profile_overrides.is_causal
vae_tiling: generator.pipeline.profile_overrides.vae_tiling
vae_sp: generator.pipeline.profile_overrides.vae_sp
dmd_denoising_steps: generator.pipeline.profile_overrides.dmd_denoising_steps
ti2v_task: generator.pipeline.profile_overrides.ti2v_task
boundary_ratio: generator.pipeline.profile_overrides.boundary_ratio
preset_owned:
embedded_cfg_scale: generator.pipeline.preset_overrides.embedded_cfg_scale
flow_shift: generator.pipeline.preset_overrides.flow_shift
flow_shift_sr: generator.pipeline.preset_overrides.flow_shift_sr
is_causal: generator.pipeline.preset_overrides.is_causal
vae_tiling: generator.pipeline.preset_overrides.vae_tiling
vae_sp: generator.pipeline.preset_overrides.vae_sp
dmd_denoising_steps: generator.pipeline.preset_overrides.dmd_denoising_steps
ti2v_task: generator.pipeline.preset_overrides.ti2v_task
boundary_ratio: generator.pipeline.preset_overrides.boundary_ratio
compatibility_only:
model_path: "Redundant with generator.model_path."
disable_autocast: "Duplicated by generator.engine.disable_autocast during migration."
@@ -97,7 +97,7 @@ surfaces:
postprocess_text_funcs: "Internal text postprocessing hooks."
pipeline_config_extensions:
profile_owned:
preset_owned:
conditioning_strategy:
sources:
- fastvideo.configs.pipelines.cosmos.CosmosConfig
@@ -309,8 +309,8 @@ surfaces:
compatibility_only:
batch_size: "Gen3C inference-only tuning field pending typed batching design."
gradient_checkpointing: "Gen3C inference-only compatibility field pending typed batching design."
guidance_scale: "Gen3C pipeline-level default pending profile/default-request cleanup."
num_inference_steps: "Gen3C pipeline-level default pending profile/default-request cleanup."
guidance_scale: "Gen3C pipeline-level default pending preset/default-request cleanup."
num_inference_steps: "Gen3C pipeline-level default pending preset/default-request cleanup."
internal_only:
audio_decoder_config: "Legacy internal component config object."
audio_decoder_precision: "Precision override pending dedicated component precision design."
@@ -345,6 +345,7 @@ surfaces:
num_inference_steps: request.sampling.num_inference_steps
num_inference_steps_sr: request.sampling.num_inference_steps_sr
guidance_scale: request.sampling.guidance_scale
guidance_scale_2: request.sampling.guidance_scale_2
guidance_rescale: request.sampling.guidance_rescale
boundary_ratio: request.sampling.boundary_ratio
sigmas: request.sampling.sigmas
@@ -353,96 +354,36 @@ surfaces:
return_frames: request.output.return_frames
return_trajectory_latents: request.runtime.return_trajectory_latents
return_trajectory_decoded: request.runtime.return_trajectory_decoded
profile_owned:
continuation_state: request.state
return_continuation_state: request.output.return_state
preset_owned:
t_thresh: request.stage_overrides.refine.t_thresh
spatial_refine_only: request.stage_overrides.refine.spatial_refine_only
num_cond_frames: request.stage_overrides.refine.num_cond_frames
trajectory_type: request.extensions.gen3c.trajectory_type
movement_distance: request.extensions.gen3c.movement_distance
camera_rotation: request.extensions.gen3c.camera_rotation
prompt_attention_mask: request.extensions.hyworld.prompt_attention_mask
negative_attention_mask: request.extensions.hyworld.negative_attention_mask
camera_states: request.extensions.hunyuangamecraft.camera_states
camera_trajectory: request.extensions.hunyuangamecraft.camera_trajectory
action_list: request.extensions.hunyuangamecraft.action_list
action_speed_list: request.extensions.hunyuangamecraft.action_speed_list
gt_latents: request.extensions.hunyuangamecraft.gt_latents
conditioning_mask: request.extensions.hunyuangamecraft.conditioning_mask
ltx2_cfg_scale_video: request.extensions.ltx2.cfg_scale_video
ltx2_cfg_scale_audio: request.extensions.ltx2.cfg_scale_audio
ltx2_modality_scale_video: request.extensions.ltx2.modality_scale_video
ltx2_modality_scale_audio: request.extensions.ltx2.modality_scale_audio
ltx2_rescale_scale: request.extensions.ltx2.rescale_scale
ltx2_stg_scale_video: request.extensions.ltx2.stg_scale_video
ltx2_stg_scale_audio: request.extensions.ltx2.stg_scale_audio
ltx2_stg_blocks_video: request.extensions.ltx2.stg_blocks_video
ltx2_stg_blocks_audio: request.extensions.ltx2.stg_blocks_audio
internal_only:
data_type: "Derived from the request shape and not a public input."
sampling_param_extensions:
moved:
guidance_scale_2:
target: request.sampling.guidance_scale_2
sources:
- fastvideo.configs.sample.lingbotworld.LingBotWorld_SamplingParam
- fastvideo.configs.sample.lingbotworld.Wan2_2_I2V_A14B_SamplingParam
- fastvideo.configs.sample.wan.SelfForcingWan2_2_T2V_A14B_480P_SamplingParam
- fastvideo.configs.sample.wan.Wan2_2_I2V_A14B_SamplingParam
- fastvideo.configs.sample.wan.Wan2_2_T2V_A14B_SamplingParam
profile_owned:
action_list:
target: request.extensions.hunyuangamecraft.action_list
sources:
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
action_speed_list:
target: request.extensions.hunyuangamecraft.action_speed_list
sources:
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
camera_states:
target: request.extensions.hunyuangamecraft.camera_states
sources:
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
camera_trajectory:
target: request.extensions.hunyuangamecraft.camera_trajectory
sources:
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
conditioning_mask:
target: request.extensions.hunyuangamecraft.conditioning_mask
sources:
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
gt_latents:
target: request.extensions.hunyuangamecraft.gt_latents
sources:
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraftSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft65FrameSamplingParam
- fastvideo.configs.sample.hunyuangamecraft.HunyuanGameCraft129FrameSamplingParam
prompt_attention_mask:
target: request.extensions.hyworld.prompt_attention_mask
sources: [fastvideo.configs.sample.hyworld.HYWorld_SamplingParam]
negative_attention_mask:
target: request.extensions.hyworld.negative_attention_mask
sources: [fastvideo.configs.sample.hyworld.HYWorld_SamplingParam]
ltx2_cfg_scale_audio:
target: request.extensions.ltx2.cfg_scale_audio
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
ltx2_cfg_scale_video:
target: request.extensions.ltx2.cfg_scale_video
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
ltx2_modality_scale_audio:
target: request.extensions.ltx2.modality_scale_audio
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
ltx2_modality_scale_video:
target: request.extensions.ltx2.modality_scale_video
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
ltx2_rescale_scale:
target: request.extensions.ltx2.rescale_scale
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
ltx2_stg_blocks_audio:
target: request.extensions.ltx2.stg_blocks_audio
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
ltx2_stg_blocks_video:
target: request.extensions.ltx2.stg_blocks_video
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
ltx2_stg_scale_audio:
target: request.extensions.ltx2.stg_scale_audio
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
ltx2_stg_scale_video:
target: request.extensions.ltx2.stg_scale_video
sources: [fastvideo.configs.sample.ltx2.LTX2BaseSamplingParam]
sampling_param_extensions: {}
openai_image_request:
kept:
@@ -495,230 +436,14 @@ cli:
notes:
- "CLI parity is checked against the actual generate/serve parser dest sets."
- "The inventory tracks parser dest names, excluding argparse's implicit help action."
- "The refactored inference CLI is config-only: subcommands expose only --config, and any additional CLI input must use dotted override paths."
generate:
explicit_local_fields:
- config
expected_dests:
- VSA_sparsity
- boundary_ratio
- bsa_cdf_threshold
- bsa_chunk_k
- bsa_chunk_q
- bsa_sparsity
- config
- disable_autocast
- dist_timeout
- distributed_executor_backend
- dit_config.prefix
- dit_config.quant_config
- dit_cpu_offload
- dit_layerwise_offload
- dit_precision
- dmd_denoising_steps
- embedded_cfg_scale
- enable_bsa
- enable_stage_verification
- enable_torch_compile
- flow_shift
- fps
- guidance_rescale
- guidance_scale
- height
- hsdp_replicate_dim
- hsdp_shard_dim
- image_encoder_cpu_offload
- image_encoder_precision
- image_path
- inference_mode
- init_weights_from_safetensors
- init_weights_from_safetensors_2
- lora_nickname
- lora_path
- lora_target_modules
- ltx2_initial_latent_path
- ltx2_vae_spatial_tile_overlap_in_pixels
- ltx2_vae_spatial_tile_size_in_pixels
- ltx2_vae_temporal_tile_overlap_in_frames
- ltx2_vae_temporal_tile_size_in_frames
- ltx2_vae_tiling
- master_port
- moba_config_path
- mode
- model_path
- negative_prompt
- num_cond_frames
- num_frames
- num_gpus
- num_inference_steps
- num_videos_per_prompt
- output_path
- output_type
- output_video_name
- override_pipeline_cls_name
- override_text_encoder_quant
- override_text_encoder_safetensors
- override_transformer_cls_name
- pin_cpu_memory
- pipeline_config_path
- preprocess.dataloader_num_workers
- preprocess.dataset_output_dir
- preprocess.dataset_path
- preprocess.dataset_type
- preprocess.do_temporal_sample
- preprocess.drop_short_ratio
- preprocess.flush_frequency
- preprocess.max_height
- preprocess.max_width
- preprocess.model_path
- preprocess.num_frames
- preprocess.preprocess_video_batch_size
- preprocess.samples_per_file
- preprocess.seed
- preprocess.speed_factor
- preprocess.train_fps
- preprocess.training_cfg_rate
- preprocess.video_length_tolerance_range
- preprocess.video_loader_type
- preprocess.with_audio
- prompt
- prompt_path
- prompt_txt
- refine_from
- return_frames
- return_trajectory_decoded
- return_trajectory_latents
- revision
- save_video
- seed
- sp_size
- spatial_refine_only
- t_thresh
- text_encoder_configs
- text_encoder_cpu_offload
- text_encoder_precisions
- torch_compile_kwargs
- tp_size
- trust_remote_code
- use_fsdp_inference
- vae_config.blend_num_frames
- vae_config.load_decoder
- vae_config.load_encoder
- vae_config.tile_sample_min_height
- vae_config.tile_sample_min_num_frames
- vae_config.tile_sample_min_width
- vae_config.tile_sample_stride_height
- vae_config.tile_sample_stride_num_frames
- vae_config.tile_sample_stride_width
- vae_config.use_parallel_tiling
- vae_config.use_temporal_tiling
- vae_config.use_tiling
- vae_cpu_offload
- vae_precision
- vae_sp
- vae_tiling
- video_path
- width
- workload_type
serve:
explicit_local_fields:
- config
- host
- output_dir
- port
expected_dests:
- VSA_sparsity
- bsa_cdf_threshold
- bsa_chunk_k
- bsa_chunk_q
- bsa_sparsity
- config
- disable_autocast
- dist_timeout
- distributed_executor_backend
- dit_config.prefix
- dit_config.quant_config
- dit_cpu_offload
- dit_layerwise_offload
- dit_precision
- dmd_denoising_steps
- embedded_cfg_scale
- enable_bsa
- enable_stage_verification
- enable_torch_compile
- flow_shift
- host
- hsdp_replicate_dim
- hsdp_shard_dim
- image_encoder_cpu_offload
- image_encoder_precision
- inference_mode
- init_weights_from_safetensors
- init_weights_from_safetensors_2
- lora_nickname
- lora_path
- lora_target_modules
- ltx2_initial_latent_path
- ltx2_vae_spatial_tile_overlap_in_pixels
- ltx2_vae_spatial_tile_size_in_pixels
- ltx2_vae_temporal_tile_overlap_in_frames
- ltx2_vae_temporal_tile_size_in_frames
- ltx2_vae_tiling
- master_port
- mode
- model_path
- num_gpus
- output_dir
- output_type
- override_pipeline_cls_name
- override_text_encoder_quant
- override_text_encoder_safetensors
- override_transformer_cls_name
- pin_cpu_memory
- pipeline_config_path
- port
- preprocess.dataloader_num_workers
- preprocess.dataset_output_dir
- preprocess.dataset_path
- preprocess.dataset_type
- preprocess.do_temporal_sample
- preprocess.drop_short_ratio
- preprocess.flush_frequency
- preprocess.max_height
- preprocess.max_width
- preprocess.model_path
- preprocess.num_frames
- preprocess.preprocess_video_batch_size
- preprocess.samples_per_file
- preprocess.seed
- preprocess.speed_factor
- preprocess.train_fps
- preprocess.training_cfg_rate
- preprocess.video_length_tolerance_range
- preprocess.video_loader_type
- preprocess.with_audio
- prompt_txt
- revision
- sp_size
- text_encoder_cpu_offload
- text_encoder_precisions
- torch_compile_kwargs
- tp_size
- trust_remote_code
- use_fsdp_inference
- vae_config.blend_num_frames
- vae_config.load_decoder
- vae_config.load_encoder
- vae_config.tile_sample_min_height
- vae_config.tile_sample_min_num_frames
- vae_config.tile_sample_min_width
- vae_config.tile_sample_stride_height
- vae_config.tile_sample_stride_num_frames
- vae_config.tile_sample_stride_width
- vae_config.use_parallel_tiling
- vae_config.use_temporal_tiling
- vae_config.use_tiling
- vae_cpu_offload
- vae_precision
- vae_sp
- vae_tiling
- workload_type
+6 -5
View File
@@ -12,7 +12,7 @@ FastVideo maps a Diffusers-style repo into a pipeline like this:
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
weight name translation.
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
- `fastvideo/api/sampling_param.py`: runtime sampling parameters.
- `fastvideo/pipelines/basic/*`: end-to-end pipelines.
- `fastvideo/pipelines/stages/*`: reusable pipeline stages.
- `fastvideo/models/loader/*`: component loaders for Diffusers-style repos.
@@ -26,7 +26,7 @@ Minimal usage (from `examples/inference/basic/basic.py`):
```python
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
@@ -49,8 +49,9 @@ runtime parameters consistent:
- `fastvideo/configs/models/`: architecture definitions, layer shapes, and
`param_names_mapping` rules for key renaming.
- `fastvideo/configs/pipelines/`: pipeline wiring and required components.
- `fastvideo/configs/sample/`: default sampling parameters (steps, frames,
guidance scale, resolution, fps).
- `fastvideo/api/sampling_param.py`: sampling parameters (steps, frames,
guidance scale, resolution, fps). Defaults come from profiles in
`fastvideo/pipelines/basic/<family>/profiles.py`.
- `fastvideo/registry.py`: unified registry for pipeline config + sampling
defaults and model metadata resolution, defined via explicit
`register_configs(...)` blocks (no separate dict registries).
@@ -142,7 +143,7 @@ How this maps to FastVideo:
- `T5TokenizerFast` -> loaded via HF in `fastvideo/models/loader/`
- `UniPCMultistepScheduler` -> loaded via Diffusers scheduler utilities
- Pipeline defaults -> `fastvideo/configs/pipelines/wan.py`
- Sampling defaults -> `fastvideo/configs/sample/wan.py`
- Sampling defaults -> `fastvideo/pipelines/basic/wan/profiles.py`
## Pipeline system
+177
View File
@@ -0,0 +1,177 @@
# Streaming WebSocket Server Contract
The streaming server (`fastvideo/entrypoints/streaming/server.py`) speaks
a JSON-over-WebSocket protocol with binary fMP4 chunks for media. This
document is the authoritative spec for the message catalogue and the
session state machine. Any change to either must update this document
in the same PR that touches `protocol.py` or `session.py`.
## Endpoint
| Path | Protocol | Purpose |
|---|---|---|
| `WS /v1/stream` | WebSocket (JSON + binary) | Per-session realtime streaming |
| `GET /health` | HTTP | Liveness probe (`status`, `stream_mode`, active `sessions`) |
The server is launched by `fastvideo serve --config <serve.yaml>` when
the config carries a `streaming:` block. Without that block the same CLI
launches the OpenAI stateless HTTP server instead.
## Connection lifecycle
Every WebSocket connection holds exactly one `Session`. Sessions move
through the states in `SessionState` (`fastvideo/entrypoints/streaming/session.py`).
```
┌──────────────┐
│ INITIALIZING │ ← WebSocket accepted, before init frame
└──────┬───────┘
│ session_init_v2 received
┌──────────────┼──────────────┐
▼ ▼ ▼
QUEUED GPU_BINDING REJECTED
│ │ ↑
│ slot ready │ │ max-sessions hit
▼ ▼ │ or invalid init
┌────────┐ │
│ ACTIVE │ ────────┘
└────┬───┘
segment loop │
│
┌───────────┼───────────┐
▼ ▼ ▼
COMPLETE ERROR TIMEOUT
(clean leave) (any failure) (idle / segment_cap reached)
```
Terminal states (`COMPLETE`, `ERROR`, `TIMEOUT`, `REJECTED`) are sinks —
no transitions out. The transition matrix is enforced in
`session.py::_VALID_TRANSITIONS`; bad transitions raise.
`SessionManager` enforces the per-process budgets pulled from
`StreamingConfig`:
- `session_timeout_seconds` — idle reaper drops sessions that haven't
advanced; non-terminal sessions transition to `TIMEOUT`.
- `generation_segment_cap` — a session that hits the cap transitions to
`COMPLETE` after the last segment ships.
## Message catalogue
Every JSON frame carries `{"type": <str>, ...}`. Pydantic models in
`protocol.py` are the source of truth; this table is the human-readable
view.
### Client → server
| `type` | Required fields | Purpose |
|---|---|---|
| `session_init_v2` | — | Opening frame. Carries preset, curated prompts, optional initial image, feature toggles, optional `continuation_state` to resume from a snapshot. |
| `segment_prompt_source` | `prompt` | Request the next segment using the supplied prompt; optional sampling overrides (`seed`, `num_inference_steps`, `guidance_scale`, `negative_prompt`). |
| `seed_prompts_updated` | `seed_prompts` | Replace the session's seed-prompt list; takes effect on the next segment. |
| `enhancement_updated` | `enabled` | Toggle prompt enhancement for subsequent segments. |
| `auto_extension_updated` | `enabled` | Toggle automatic per-segment prompt extension. |
| `loop_generation_updated` | `enabled` | Toggle loop-generation mode. |
| `generation_paused_updated` | `paused` | Pause/resume segment generation; queued requests defer. |
| `snapshot_state` | — | Request the current `ContinuationState` for export; server replies with `continuation_state_snapshot`. |
The opening frame must be `session_init_v2`. Any other first frame is
rejected with an `error` (code `invalid_message`) and the WebSocket is
closed.
### Server → client
| `type` | Carries | When emitted |
|---|---|---|
| `queue_status` | `position`, `queue_depth` | After `session_init_v2` accepted, before GPU binding. |
| `gpu_assigned` | GPU id, model id | Once a generator slot is bound. |
| `ltx2_stream_start` | session-level metadata | Once the session enters `ACTIVE`. |
| `ltx2_segment_start` | `segment_idx`, `prompt`, prompt source | When a `segment_prompt_source` request begins generation. |
| `step_complete` | `segment_idx`, denoise timings | After the segment's denoising loop finishes (before media emission). |
| `media_init` | `segment_idx`, mime, stream id | First frame of fMP4 output for the segment. |
| binary frame | fMP4 fragment bytes | Subsequent media chunks; the protocol enforces that `media_init` precedes any binary frames. |
| `media_segment_complete` | `segment_idx`, chunk count, byte count | Last media chunk for the segment. |
| `ltx2_segment_complete` | `segment_idx`, segment summary | Segment fully shipped; ready for the next `segment_prompt_source`. |
| `ltx2_stream_complete` | session summary | Session reached `generation_segment_cap` or client requested clean shutdown. |
| `session_timeout` | reason | Session hit `session_timeout_seconds`; immediately followed by close. |
| `continuation_state_snapshot` | `kind`, `payload` | Reply to `snapshot_state`. The payload is the same shape produced by `LTX2ContinuationState.to_continuation_state(...)`. |
| `error` | `code`, `message` | Any validation/runtime error. Non-fatal errors keep the connection open; fatal errors precede a `close`. |
## Continuation state
The session optionally accepts a `continuation_state` dict inside the
opening `session_init_v2` frame. When present, the server hydrates it
into a `ContinuationState(kind, payload)` envelope and feeds it as the
`request.state` on the first segment's `GenerationRequest` — letting a
client resume after a disconnect, migrate sessions across processes,
or replay a prior session.
After every segment, if the runtime returns a fresh state, the server
persists it to the `SessionStore` so a `snapshot_state` request can
export it. The store and serialization contracts live with the model
family (e.g. `fastvideo/pipelines/basic/ltx2/continuation.py` for LTX-2).
## Example flow
```
client server
────── ──────
WS /v1/stream ─────── connect ─────────────────────────►
◄────── (accept)
{"type": "session_init_v2",
"preset": "ltx2_two_stage",
"curated_prompts": ["a fox in snow", "the fox jumps"],
"initial_image": {...},
"stream_mode": "av_fmp4"} ─────────────────────────────►
(validate, queue, bind)
◄──── {"type": "queue_status",
"position": 0, "queue_depth": 0}
◄──── {"type": "gpu_assigned",
"gpu_id": 0, "model_id": "..."}
◄──── {"type": "ltx2_stream_start", ...}
{"type": "segment_prompt_source",
"prompt": "a fox in snow",
"source": "curated"} ───────────────────────────────────►
(run pipeline)
◄──── {"type": "ltx2_segment_start",
"segment_idx": 1, ...}
◄──── {"type": "step_complete",
"segment_idx": 1, "timings": {...}}
◄──── {"type": "media_init",
"segment_idx": 1,
"mime": "video/mp4", ...}
◄──── <binary fMP4 init segment>
◄──── <binary fMP4 fragment>
◄──── <binary fMP4 fragment>
◄──── {"type": "media_segment_complete",
"segment_idx": 1, "chunks": 12}
◄──── {"type": "ltx2_segment_complete",
"segment_idx": 1, ...}
{"type": "segment_prompt_source",
"prompt": "the fox jumps"} ─────────────────────────────►
(segment 2 …)
{"type": "snapshot_state"} ──────────────────────────────►
◄──── {"type": "continuation_state_snapshot",
"kind": "ltx2.v1",
"payload": {"schema_version": 1, ...}}
(close) ──────────────────────────────────────────────────►
(session → COMPLETE)
```
## Backward / forward compatibility
- Adding a new client message: append a Pydantic model to `protocol.py`
with a unique `type`; add the discriminator entry to `ClientMessage`;
add a row to the table above. Old clients that don't send the new
message remain compatible.
- Adding a new server message: emit only when a new feature flag is
enabled (or always emit, since clients ignore unknown types).
- Changing an existing message: bump the `type` (e.g. `session_init_v2`
→ `session_init_v3`) and accept both for one release cycle. Never
silently change field semantics under the same `type`.
+24 -1
View File
@@ -16,7 +16,8 @@ Both models are trained on **61×448×832** resolution but support generating vi
First install [VSA](../attention/vsa/index.md). Set `MODEL_BASE` to your own model path and run:
```bash
bash scripts/inference/v1_inference_wan_dmd.sh
FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN \
fastvideo generate --config scripts/inference/inference_wan_VSA_DMD_1_3B.yaml
```
## 🗂️ Dataset
@@ -85,3 +86,25 @@ sbatch examples/distill/Wan2.2-TI2V-5B-Diffusers/Data-free/distill_dmd_t2v_5B.sh
- Learning rate: 2e-5
- Training steps: 3000 (~12 hours)
- HSDP shard dim: 1
## 🧭 Note on `real_score_guidance_scale`
The teacher CFG used inside the DMD loss follows the DMD2 reference
implementation and uses the parameterization
```
x = x_cond + w * (x_cond - x_uncond)
```
rather than the Ho & Salimans form `x_uncond + w * (x_cond - x_uncond)`. The
two are mathematically equivalent up to a constant offset:
| `real_score_guidance_scale` (`w`) | Equivalent standard CFG (`w + 1`) | Output |
|-----------------------------------|-----------------------------------|-----------------------|
| `-1` | `0` | unconditional |
| `0` | `1` | conditional |
| `3.5` (default) | `4.5` | strong guidance |
So `real_score_guidance_scale` should be read as the **extra** guidance
strength added on top of the conditional prediction. When porting values
from a paper that uses the Ho & Salimans form, subtract 1.
+1 -1
View File
@@ -33,7 +33,7 @@ The following two classes `PipelineConfig` and `SamplingParam` are used to confi
### SamplingParam
::: fastvideo.configs.sample.base.SamplingParam
::: fastvideo.api.sampling_param.SamplingParam
options:
show_root_heading: true
show_source: false
+10 -15
View File
@@ -128,19 +128,14 @@ Concrete hierarchy: `DiTConfig` → `DiTArchConfig`, `VAEConfig` →
- `dump_to_json()` / `load_from_json()` — JSON persistence. Callable
fields and `arch_config` are excluded from dumps.
### SamplingParam (`fastvideo/configs/sample/`)
### SamplingParam (`fastvideo/api/sampling_param.py`)
Generation parameters separate from pipeline config. Each model family
provides defaults:
provides defaults via a profile (see `fastvideo/pipelines/basic/<family>/profiles.py`):
```python
@dataclass
class WanT2V_1_3B_SamplingParam(SamplingParam):
height: int = 480
width: int = 832
num_frames: int = 81
guidance_scale: float = 3.0
num_inference_steps: int = 50
sp = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# sp.height == 480, sp.width == 832, sp.num_frames == 81, etc.
```
## Component Loading
@@ -430,9 +425,9 @@ User: generator.generate_video(prompt, ...)
`fastvideo/configs/pipelines/<model>.py`. Set DiT/VAE/encoder configs,
flow_shift, precision defaults.
2. **Sampling param** — Create a `SamplingParam` subclass in
`fastvideo/configs/sample/<model>.py`. Set default height, width,
num_frames, guidance_scale, num_inference_steps.
2. **Sampling param profile** — Create a profile in
`fastvideo/pipelines/basic/<model>/profiles.py` with default height,
width, num_frames, guidance_scale, num_inference_steps.
3. **Register configs** — In `fastvideo/registry.py`, add a
`register_configs()` call inside `_register_configs()` with
@@ -455,6 +450,6 @@ User: generator.generate_video(prompt, ...)
`fastvideo/pipelines/stages/`, implement `forward()`, optionally
implement `verify_input()`/`verify_output()`.
7. **Verify** — Run `fastvideo generate --model-path <path> --prompt
"test" --num-inference-steps 2` to confirm the pipeline loads and
generates output.
7. **Verify** — Run `fastvideo generate --config <config.yaml>` with a
minimal nested config to confirm the pipeline loads and generates
output.
+42 -81
View File
@@ -1,71 +1,29 @@
# FastVideo CLI Inference
The FastVideo CLI exposes the same core inference controls as the Python API.
The FastVideo CLI is config-first. Inference runs are driven by a nested JSON or
YAML config, with optional dotted-path overrides on the command line. The
contract matches training: use an explicit subcommand plus `--config`, then add
any dotted overrides you need.
## Basic Usage
Use either:
1. `--model-path` + `--prompt`
2. `--model-path` + `--prompt-txt` (batch prompts, one line per prompt)
3. `--config` (JSON/YAML)
```bash
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--prompt "A cat playing with a ball of yarn"
fastvideo generate --config config.yaml
fastvideo serve --config serve.yaml
```
```bash
fastvideo generate --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--prompt-txt prompts.txt
```
You cannot provide both `--prompt` and `--prompt-txt` in the same run.
## View All Arguments
```bash
fastvideo generate --help
```
Arguments come from:
The subcommands intentionally expose only `--config`. Any per-run CLI changes
must use dotted override paths such as:
- FastVideo runtime args (`FastVideoArgs`)
- Sampling args (`SamplingParam`)
- Pipeline config args (`PipelineConfig`)
## Common Arguments
### Parallelism
- `--num-gpus`
- `--sp-size`
- `--tp-size`
### Sampling
- `--num-frames`
- `--height` / `--width`
- `--num-inference-steps`
- `--guidance-scale`
- `--seed`
- `--negative-prompt`
### Output
- `--output-path`
- `--save-video` / `--no-save-video`
- `--return-frames`
### Offloading and Performance
- `--dit-layerwise-offload`
- `--use-fsdp-inference`
- `--text-encoder-cpu-offload`
- `--image-encoder-cpu-offload`
- `--vae-cpu-offload`
- `--enable-torch-compile`
- `--torch-compile-kwargs`
- `--generator.engine.num_gpus 2`
- `--request.sampling.seed 42`
- `--server.port 9000`
## Using Config Files
@@ -73,50 +31,53 @@ Arguments come from:
fastvideo generate --config config.yaml
```
Config files can be JSON or YAML. CLI flags override config-file values.
Config files can be JSON or YAML. Dotted CLI overrides take precedence over
config-file values.
Example `config.yaml`:
```yaml
model_path: "FastVideo/FastHunyuan-diffusers"
prompt: "A capybara lounging in a hammock"
output_path: "outputs/"
num_gpus: 2
sp_size: 2
tp_size: 1
num_frames: 45
height: 720
width: 1280
num_inference_steps: 6
seed: 1024
dit_precision: "bf16"
vae_precision: "fp16"
vae_tiling: true
vae_sp: true
enable_torch_compile: false
generator:
model_path: FastVideo/FastHunyuan-diffusers
engine:
num_gpus: 2
parallelism:
sp_size: 2
tp_size: 1
request:
prompt: A capybara lounging in a hammock
sampling:
num_frames: 45
height: 720
width: 1280
num_inference_steps: 6
seed: 1024
output:
output_path: outputs/
```
Notes:
- Use `dit_precision` / `vae_precision` (not `precision`).
- Nested config objects are supported, for example `vae_config` and
`dit_config`.
- `generator` and `request` are the top-level keys for generation configs.
- `serve` configs use `generator`, `server`, and optional `default_request`.
- Prompt text files belong under `request.inputs.prompt_path`.
## Examples
Simple generation:
```bash
fastvideo generate \
--model-path FastVideo/FastHunyuan-diffusers \
--prompt "A cat playing with a ball of yarn" \
--num-frames 45 --height 720 --width 1280 \
--num-inference-steps 6 --seed 1024 \
--output-path outputs/
fastvideo generate --config config.yaml
```
Config + CLI override:
Config + dotted override:
```bash
fastvideo generate --config config.yaml --prompt "A panda skiing at sunset"
fastvideo generate --config config.yaml --request.prompt "A panda skiing at sunset"
```
Helper wrapper with positional config path:
```bash
bash scripts/inference/run.sh scripts/inference/inference_wan.yaml
```
+27 -19
View File
@@ -73,32 +73,40 @@ if __name__ == '__main__':
## JSON/YAML Config Files (CLI)
The CLI supports `--config` with JSON or YAML. Command-line arguments override
config file values.
By default, `fastvideo generate` uses `return_frames=false` unless you set
`--return-frames` (or `return_frames: true` in config).
The inference CLI is config-first. Use an explicit subcommand with `--config`,
then apply optional dotted overrides on top, matching the training CLI style.
By default, CLI generation uses `return_frames=false` unless you set
`request.output.return_frames: true` in config or via a dotted override.
```bash
fastvideo generate --config config.yaml
```
Use CLI argument names as keys (underscore or hyphen is accepted). Example:
Example nested config:
```yaml
model_path: "FastVideo/FastHunyuan-diffusers"
prompt: "A capybara relaxing in a hammock"
num_gpus: 2
sp_size: 2
num_frames: 45
height: 720
width: 1280
num_inference_steps: 6
seed: 1024
dit_precision: "bf16"
vae_precision: "fp16"
vae_tiling: true
vae_sp: true
enable_torch_compile: false
generator:
model_path: FastVideo/FastHunyuan-diffusers
engine:
num_gpus: 2
parallelism:
sp_size: 2
request:
prompt: A capybara relaxing in a hammock
sampling:
num_frames: 45
height: 720
width: 1280
num_inference_steps: 6
seed: 1024
output:
output_path: outputs/
```
Override individual values from the CLI with dotted paths:
```bash
fastvideo generate --config config.yaml --request.sampling.seed 42
```
## Performance Optimization
+1 -1
View File
@@ -89,7 +89,7 @@ GEN3C defaults in FastVideo:
These values are defined in:
- `fastvideo/configs/sample/gen3c.py`
- `fastvideo/pipelines/basic/gen3c/profiles.py`
- `fastvideo/configs/pipelines/gen3c.py`
and align with the official GEN3C inference defaults in:
+1 -1
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
def main():
@@ -1,5 +1,5 @@
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
def main():
@@ -1,6 +1,6 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
def main():
+1 -1
View File
@@ -2,7 +2,7 @@ import os
import time
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_dmd2"
def main():
+1 -1
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
import json
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_hy15"
def main():
+1 -1
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
import json
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_hy15_1080p"
def main():
@@ -1,7 +1,7 @@
from fastvideo import VideoGenerator
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_lingbotworld"
def main():
# FastVideo will automatically use the optimal default arguments for the
+1 -1
View File
@@ -1,5 +1,5 @@
from fastvideo import VideoGenerator, PipelineConfig
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
def main():
config = PipelineConfig.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
+1 -1
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples"
def main():
@@ -2,7 +2,7 @@
from fastvideo import VideoGenerator, SamplingParam
import json
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_i2v"
def main():
@@ -2,7 +2,7 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_self_forcing_causal_wan2_2_14B_t2v"
def main():
+1 -1
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_wan2_2_14B_t2v"
def main():
+1 -1
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_wan2_1_Fun"
OUTPUT_NAME = "wan2.1_test"
+1 -1
View File
@@ -1,6 +1,6 @@
from fastvideo import VideoGenerator
# from fastvideo.configs.sample import SamplingParam
# from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "video_samples_wan2_2_14B_i2v"
def main():
+100
View File
@@ -0,0 +1,100 @@
"""Generate one LTX2 video and score it with VBench metrics.
The generation block is the same as
``examples/inference/basic/basic_ltx2.py`` — same prompt, same model,
same shape, same num_frames. After ``shutdown()`` the script loads the
mp4 back, builds a single :class:`fastvideo.eval.Evaluator`, and runs
the prompt-aware VBench subset that's meaningful for an arbitrary
text→video sample.
The first run downloads CLIP / DINO / RAFT / AMT / ViCLIP / MUSIQ
weights to ``~/.cache/fastvideo/eval/`` (~few GB total).
GPU memory caveat
-----------------
Scoring 1088×1920×121 with all 8 metrics needs a dedicated GPU (~80 GB).
On a shared GPU, ``vbench.motion_smoothness`` (AMT correlation volume)
will OOM — its memory autoscale reads ``total_memory`` rather than
``mem_get_info()`` free memory and therefore underestimates the
required scale-down. Drop ``motion_smoothness`` from ``METRICS`` if
sharing, or run on a smaller-resolution generation.
"""
import torch
from fastvideo import VideoGenerator
from fastvideo.eval import Evaluator
from fastvideo.eval.io import build_eval_kwargs
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
# VBench sub-metrics meaningful for an arbitrary text→video sample
# (just the generated frames, optionally fps + the source prompt).
# Structured-prompt metrics (vbench.color, vbench.multiple_objects,
# vbench.scene, ...) are excluded — they need prompts built to a
# specific schema.
METRICS = [
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
"vbench.subject_consistency", # DINO frame-to-first cosine
"vbench.background_consistency", # DINO on background patches
"vbench.imaging_quality", # pyiqa MUSIQ
"vbench.temporal_flickering", # pixel-wise frame deltas
"vbench.motion_smoothness", # AMT frame interpolator residual
"vbench.dynamic_degree", # RAFT optical-flow magnitude (needs fps)
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
]
def main() -> None:
# ----- generation (matches examples/inference/basic/basic_ltx2.py) -----
generator = VideoGenerator.from_pretrained(
"Davids048/LTX2-Base-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
num_frames=121,
height=1088,
width=1920,
)
generator.shutdown()
# Free residual CUDA memory the generator left behind so the
# evaluator can grab the largest possible workspace for AMT/RAFT.
torch.cuda.empty_cache()
# ----- scoring -----
print(f"\n[eval] building evaluator: {METRICS}")
evaluator = Evaluator(metrics=METRICS)
# LTX2 outputs at 24 fps by default.
sample = build_eval_kwargs({"prompt": PROMPT}, output_path, fps=24.0)
print(f"[eval] running ({sample['video'].shape[1]} frames @ 24 fps)...")
results = evaluator.evaluate(**sample)
print("\n=== VBench scores ===")
for name in METRICS:
r = results[name]
if r.score is None:
reason = r.details.get("skipped", "no score")
print(f" {name}: SKIPPED ({reason})")
else:
print(f" {name}: {r.score:.4f}")
if __name__ == "__main__":
main()
+157
View File
@@ -0,0 +1,157 @@
"""End-to-end Physics-IQ: dataset → generate → score → aggregate.
Generates one video per take-1 scenario with LTX2 (using the scenario
caption as the prompt), scores each generated video against the take-1
reference and the take-2 "physical-variance" reference, and prints
aggregate scores using :meth:`PhysicsIQMetric.aggregate_components` —
the official scoring recipe from the upstream benchmark.
Reference videos / masks / switch-frames auto-fetch on first miss into
``${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq/``; pass ``--dataset-root``
to point at a pre-downloaded mirror instead.
Quick smoke run on 4 scenarios across 2 GPUs::
python examples/inference/eval/bench_physics_iq.py \\
--limit 4 --num-gpus 2 \\
--videos-dir outputs_video/physics_iq_smoke
Re-score existing generations without regenerating::
python examples/inference/eval/bench_physics_iq.py \\
--videos-dir outputs_video/physics_iq_smoke \\
--skip-generation
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from fastvideo.eval import create_evaluator, get_metric
from fastvideo.eval.datasets import get_dataset
def _expected_filename(row: dict) -> str:
"""Filename Physics-IQ expects for the generated video for *row*.
Uses the dataset's own ``expected_gen_filename`` annotation so the
output filenames match the benchmark's manifest convention.
"""
return row["auxiliary_info"]["expected_gen_filename"]
def _generate_videos(rows: list[dict], videos_dir: Path,
model: str, num_gpus: int,
num_frames: int, height: int, width: int) -> None:
from fastvideo import VideoGenerator
videos_dir.mkdir(parents=True, exist_ok=True)
todo = [(row, videos_dir / _expected_filename(row)) for row in rows]
todo = [(row, out) for (row, out) in todo if not out.is_file()]
if not todo:
print(f"[gen] all {len(rows)} videos already present; skipping.")
return
print(f"[gen] {len(todo)}/{len(rows)} scenarios to render with {model} "
f"({num_frames}x{height}x{width})...")
gen = VideoGenerator.from_pretrained(model, num_gpus=num_gpus)
try:
for row, out_path in todo:
gen.generate_video(
prompt=row["prompt"], output_path=str(out_path), save_video=True,
num_frames=num_frames, height=height, width=width,
)
finally:
gen.shutdown()
def main() -> None:
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--dataset-root", type=Path, default=None,
help="Path to a pre-downloaded Physics-IQ release. "
"Defaults to ${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq, "
"auto-fetching missing assets from the public bucket.")
p.add_argument("--videos-dir", type=Path,
default=Path("outputs_video/bench_physics_iq"),
help="Where to read/write generated videos.")
p.add_argument("--limit", type=int, default=None,
help="Truncate to first N scenarios for smoke runs.")
p.add_argument("--num-gpus", type=int, default=1)
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
help="HF repo id of the text→video generator to use.")
p.add_argument("--num-frames", type=int, default=121)
p.add_argument("--height", type=int, default=1088)
p.add_argument("--width", type=int, default=1920)
p.add_argument("--skip-generation", action="store_true",
help="Re-score existing videos under --videos-dir.")
p.add_argument("--scores-out", type=Path, default=None,
help="Where to write per-scenario scores (JSON). "
"Defaults to <videos-dir>/scores.json.")
args = p.parse_args()
# 1. Walk the Physics-IQ corpus. Pass --limit to the dataset
# constructor so auto-download only fetches the assets we'll use.
ds = get_dataset("physics_iq", dataset_root=args.dataset_root, limit=args.limit)
rows = list(ds)
print(f"[load] Physics-IQ: {len(rows)} scenarios from {ds.dataset_dir}")
# 2. Generate (or reuse) one mp4 per scenario.
if not args.skip_generation:
_generate_videos(
rows, args.videos_dir, args.model, args.num_gpus,
args.num_frames, args.height, args.width,
)
# 3. Score each scenario. The metric reads file paths directly out
# of the row dict (reference, reference_take2, masks), so we
# just attach the generated video path and forward.
evaluator = create_evaluator(metrics=["physics_iq"], num_gpus=args.num_gpus)
samples: list[dict] = []
matched: list[dict] = []
for row in rows:
video_path = args.videos_dir / _expected_filename(row)
if not video_path.is_file():
print(f"[eval] missing {video_path}; skipping.")
continue
# The physics_iq metric accepts file paths via its polymorphic
# input handling — no need to load the tensors here.
samples.append({"video": str(video_path), **row})
matched.append(row)
all_results = evaluator.evaluate(samples=samples)
evaluator.shutdown()
# 4. Aggregate per the upstream scoring recipe.
metric = get_metric("physics_iq")
components = metric.aggregate_components(
[r["physics_iq"] for r in all_results]
)
print()
print("=== Physics-IQ aggregate ===")
for name, value in components.items():
print(f" {name:24s} {value:.4f}")
detailed = [
{
"scenario": row["auxiliary_info"]["scenario_id"],
"view": row["view"],
"scenario_name": row["auxiliary_info"]["scenario_name"],
"score": results["physics_iq"].score,
}
for row, results in zip(matched, all_results)
]
out = args.scores_out or (args.videos_dir / "scores.json")
out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(json.dumps(
{"aggregate": components, "per_scenario": detailed},
indent=2,
))
print(f"\n[done] per-scenario scores → {out}")
if __name__ == "__main__":
main()
+155
View File
@@ -0,0 +1,155 @@
"""End-to-end VBench: dataset → generate → score → aggregate.
Iterates the VBench prompt corpus, generates one video per prompt with
LTX2, scores each generated video against the requested ``vbench.*``
sub-metrics, and prints per-metric averages over the run.
Re-running with ``--skip-generation`` reuses any mp4 already on disk
under ``--videos-dir``, so you can iterate on metric selection without
re-paying the generation cost.
Example — quick smoke run on 4 prompts from the ``aesthetic_quality``
dimension across 2 GPUs::
python examples/inference/eval/bench_vbench.py \\
--dimensions aesthetic_quality \\
--limit 4 --num-gpus 2 \\
--videos-dir outputs_video/vbench_smoke
Full benchmark on a single dimension::
python examples/inference/eval/bench_vbench.py \\
--dimensions subject_consistency --num-gpus 8
"""
from __future__ import annotations
import argparse
import json
import re
from collections import defaultdict
from pathlib import Path
from fastvideo.eval import create_evaluator
from fastvideo.eval.datasets import get_dataset
def _slugify(prompt: str, max_len: int = 100) -> str:
"""Filesystem-safe filename stem; mirrors VBench's official convention."""
s = re.sub(r'[\\/:*?"<>|]', "", prompt[:max_len]).strip().strip(".")
return re.sub(r"\s+", " ", s) or "output"
def _generate_videos(prompts: list[str], videos_dir: Path,
model: str, num_gpus: int,
num_frames: int, height: int, width: int) -> None:
from fastvideo import VideoGenerator
videos_dir.mkdir(parents=True, exist_ok=True)
todo = [(p, videos_dir / f"{_slugify(p)}.mp4") for p in prompts]
todo = [(p, out) for (p, out) in todo if not out.is_file()]
if not todo:
print(f"[gen] all {len(prompts)} videos already present; skipping.")
return
print(f"[gen] {len(todo)}/{len(prompts)} prompts to render with {model} "
f"({num_frames}x{height}x{width})...")
gen = VideoGenerator.from_pretrained(model, num_gpus=num_gpus)
try:
for prompt, out_path in todo:
gen.generate_video(
prompt=prompt, output_path=str(out_path), save_video=True,
num_frames=num_frames, height=height, width=width,
)
finally:
gen.shutdown()
def main() -> None:
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--dimensions", default="aesthetic_quality,subject_consistency",
help="Comma-separated VBench dimensions (or 'all').")
p.add_argument("--limit", type=int, default=None,
help="Truncate to first N prompts for smoke runs.")
p.add_argument("--videos-dir", type=Path,
default=Path("outputs_video/bench_vbench"))
p.add_argument("--num-gpus", type=int, default=1)
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
help="HF repo id of the text→video generator to use.")
p.add_argument("--num-frames", type=int, default=121)
p.add_argument("--height", type=int, default=1088)
p.add_argument("--width", type=int, default=1920)
p.add_argument("--fps", type=float, default=24.0,
help="Frame-rate annotation passed to fps-aware metrics.")
p.add_argument("--skip-generation", action="store_true",
help="Re-score existing videos under --videos-dir without "
"regenerating.")
p.add_argument("--scores-out", type=Path, default=None,
help="Where to dump per-prompt scores as JSON. "
"Defaults to <videos-dir>/scores.json.")
args = p.parse_args()
# 1. Pull prompts from VBench.
dims_arg: list[str] | str = (
args.dimensions if args.dimensions == "all"
else [d.strip() for d in args.dimensions.split(",") if d.strip()]
)
ds = get_dataset("vbench", dimensions=dims_arg)
rows = list(ds)[: args.limit]
print(f"[load] VBench: {len(rows)} prompts across {ds.dimensions}")
# 2. Generate (or reuse) one mp4 per prompt.
if not args.skip_generation:
_generate_videos(
[row["prompt"] for row in rows],
args.videos_dir, args.model, args.num_gpus,
args.num_frames, args.height, args.width,
)
# 3. Score each video against the requested vbench sub-metrics.
metric_names = sorted(set(f"vbench.{d}" for d in ds.dimensions))
print(f"[eval] metrics: {metric_names}")
evaluator = create_evaluator(metrics=metric_names, num_gpus=args.num_gpus)
samples: list[dict] = []
matched_rows: list[dict] = []
for row in rows:
video_path = args.videos_dir / f"{_slugify(row['prompt'])}.mp4"
if not video_path.is_file():
print(f"[eval] missing {video_path}; skipping this row.")
continue
# Pass the path; the worker decodes lazily so memory stays bounded.
samples.append({
"video": str(video_path),
"fps": args.fps,
**row, # prompt / aux / dims
})
matched_rows.append(row)
all_results = evaluator.evaluate(samples=samples)
evaluator.shutdown()
# 4. Aggregate per-metric.
by_metric: dict[str, list[float]] = defaultdict(list)
detailed: list[dict] = []
for row, results in zip(matched_rows, all_results):
scores = {name: r.score for name, r in results.items()}
detailed.append({"prompt": row["prompt"], "scores": scores})
for name, score in scores.items():
if score is not None:
by_metric[name].append(score)
print()
print("=== per-metric averages ===")
for name in sorted(by_metric):
avg = sum(by_metric[name]) / len(by_metric[name])
print(f" {name:42s} {avg:.4f} (n={len(by_metric[name])})")
out = args.scores_out or (args.videos_dir / "scores.json")
out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(json.dumps(detailed, indent=2))
print(f"\n[done] per-prompt scores → {out}")
if __name__ == "__main__":
main()
+167
View File
@@ -0,0 +1,167 @@
"""End-to-end: generate a video with LTX2 and score it with VBench.
Pipeline:
prompt → LTX2-Base → mp4 → fastvideo.eval → vbench scores
Run::
pip install -e .[eval]
git submodule update --init fastvideo/third_party/eval/vbench
python examples/inference/eval/eval_ltx2_vbench.py
# or with 4 GPUs and the distilled checkpoint:
python examples/inference/eval/eval_ltx2_vbench.py \
--model FastVideo/LTX2-Distilled-Diffusers --num-gpus 4
The default metric set covers the vbench sub-metrics that are
meaningful for an arbitrary text→video sample — i.e. those that need
only the generated video (and optionally fps + the source prompt).
Structured-prompt metrics like ``vbench.color``, ``vbench.scene``,
``vbench.multiple_objects`` etc. are *not* on by default — they only
make sense when the prompt is built to a specific schema, and they
require GRiT/detectron2 setup. Pass them via ``--metrics`` if you have
a matching prompt.
First-time runs download CLIP, DINO, RAFT, AMT, ViCLIP, and MUSIQ
weights to ``~/.cache/fastvideo/eval/models/`` and
``~/.cache/torch/hub/`` (~few GB total). Subsequent runs are fast.
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from fastvideo import VideoGenerator
from fastvideo.eval import create_evaluator
from fastvideo.eval.io import load_video
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
DEFAULT_METRICS = [
# No-input metrics: just need the generated frames.
"vbench.aesthetic_quality", # CLIP + LAION aesthetic head
"vbench.subject_consistency", # DINO frame-to-first cosine
"vbench.background_consistency", # DINO on background patches
"vbench.imaging_quality", # pyiqa MUSIQ
"vbench.temporal_flickering", # pixel-wise frame deltas
"vbench.motion_smoothness", # AMT frame interpolator residual
# Need fps annotation:
"vbench.dynamic_degree", # RAFT optical-flow magnitude
# Need the source prompt:
"vbench.overall_consistency", # ViCLIP video↔prompt similarity
]
def parse_args() -> argparse.Namespace:
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--model", default="Davids048/LTX2-Base-Diffusers",
help="HF repo id of the LTX2 checkpoint.")
p.add_argument("--num-gpus", type=int, default=1)
p.add_argument("--output", default="outputs_video/ltx2_eval/clip.mp4",
help="Where to save the generated mp4.")
p.add_argument("--num-frames", type=int, default=121)
p.add_argument("--height", type=int, default=1088)
p.add_argument("--width", type=int, default=1920)
p.add_argument("--prompt", default=PROMPT)
p.add_argument("--fps", type=float, default=24.0,
help="Frame-rate annotation passed to fps-aware metrics "
"(e.g. vbench.dynamic_degree). LTX2 outputs at 24 fps "
"by default.")
p.add_argument("--metrics", default=",".join(DEFAULT_METRICS),
help="Comma-separated metric names. Pass 'all' for every "
"registered metric, or e.g. 'vbench' for the whole group.")
p.add_argument("--scores-out", default="outputs_video/ltx2_eval/scores.json")
p.add_argument("--skip-generation", action="store_true",
help="Reuse an existing --output video instead of regenerating.")
return p.parse_args()
def generate(args: argparse.Namespace) -> Path:
out = Path(args.output)
if args.skip_generation and out.is_file():
print(f"[gen] reusing existing video at {out}")
return out
out.parent.mkdir(parents=True, exist_ok=True)
print(f"[gen] loading {args.model} ({args.num_gpus} GPU)...")
generator = VideoGenerator.from_pretrained(args.model, num_gpus=args.num_gpus)
try:
print(f"[gen] generating to {out}...")
generator.generate_video(
prompt=args.prompt,
output_path=str(out),
save_video=True,
num_frames=args.num_frames,
height=args.height,
width=args.width,
)
finally:
generator.shutdown()
return out
def evaluate_video(video_path: Path, prompt: str, fps: float,
metric_names) -> dict:
print(f"[eval] loading video from {video_path}...")
video = load_video(str(video_path)) # (T, C, H, W) in [0, 1]
video = video.unsqueeze(0) # → (1, T, C, H, W)
print(f"[eval] building evaluator: {metric_names}")
evaluator = create_evaluator(metrics=metric_names, device="cuda")
print(f"[eval] running ({video.shape[1]} frames @ {fps} fps)...")
results = evaluator.evaluate(
video=video,
text_prompt=[prompt],
fps=fps,
)
if isinstance(results, list):
results = results[0] # batch of 1
return {
name: {"score": r.score, "details": r.details}
for name, r in results.items()
}
def main() -> None:
args = parse_args()
if args.metrics.strip() == "all":
metric_names = "all"
else:
metric_names = [m.strip() for m in args.metrics.split(",") if m.strip()]
video_path = generate(args)
scores = evaluate_video(video_path, args.prompt, args.fps, metric_names)
print("\n=== VBench scores ===")
for name, payload in scores.items():
print(f" {name}: {payload['score']}")
out = Path(args.scores_out)
out.parent.mkdir(parents=True, exist_ok=True)
out.write_text(json.dumps(
{"video": str(video_path), "prompt": args.prompt, "scores": scores},
indent=2,
))
print(f"[done] scores written to {out}")
if __name__ == "__main__":
main()
+95
View File
@@ -0,0 +1,95 @@
"""Score a folder of videos in parallel across multiple GPUs.
Uses :meth:`Evaluator.evaluate(samples=[...])`, which round-robins each
sample dict across the GPU replicas the evaluator was built with.
Example::
python examples/inference/eval/score_folder.py \\
--videos generated/ \\
--metrics vbench.aesthetic_quality,vbench.subject_consistency \\
--num-gpus 4 \\
--output scores.json
Pair each generated video with a same-name reference video (e.g.
``ref/<stem>.mp4``) by passing ``--reference-dir``::
python examples/inference/eval/score_folder.py \\
--videos generated/ --reference-dir ref/ \\
--metrics common.psnr,common.ssim,common.lpips \\
--num-gpus 4
"""
from __future__ import annotations
import argparse
import json
from pathlib import Path
from fastvideo.eval import create_evaluator
def _list_videos(directory: Path) -> list[Path]:
exts = {".mp4", ".avi", ".mov", ".mkv", ".gif"}
return sorted(p for p in directory.iterdir() if p.suffix.lower() in exts)
def main() -> None:
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--videos", type=Path, required=True,
help="Directory of generated videos.")
p.add_argument("--reference-dir", type=Path, default=None,
help="Directory of reference videos with matching stems.")
p.add_argument("--metrics", default="vbench.aesthetic_quality")
p.add_argument("--num-gpus", type=int, default=1)
p.add_argument("--fps", type=float, default=None,
help="Frame-rate annotation for fps-aware metrics.")
p.add_argument("--output", type=Path, default=Path("scores.json"))
args = p.parse_args()
video_paths = _list_videos(args.videos)
if not video_paths:
raise SystemExit(f"No videos under {args.videos}")
print(f"Found {len(video_paths)} videos in {args.videos}")
metrics: list[str] | str = (
args.metrics if args.metrics == "all"
else [m.strip() for m in args.metrics.split(",") if m.strip()]
)
evaluator = create_evaluator(metrics=metrics, num_gpus=args.num_gpus)
# Build per-video sample dicts holding *paths*, not pre-loaded
# tensors. Each path is decoded inside the worker thread that picks
# up its sample, so peak resident memory is bounded by num_gpus
# rather than scaling with the size of the folder.
samples: list[dict] = []
for vp in video_paths:
sample: dict = {"video": str(vp)}
if args.reference_dir is not None:
ref_path = args.reference_dir / vp.name
if not ref_path.is_file():
raise FileNotFoundError(f"Missing reference for {vp.name} at {ref_path}")
sample["reference"] = str(ref_path)
if args.fps is not None:
sample["fps"] = args.fps
samples.append(sample)
print(f"Scoring with {len(evaluator.metric_names)} metric(s) "
f"on {evaluator.num_gpus} GPU(s)...")
all_results = evaluator.evaluate(samples=samples)
evaluator.shutdown()
payload = [
{
"video": str(vp),
"scores": {name: r.score for name, r in results.items()},
}
for vp, results in zip(video_paths, all_results)
]
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(payload, indent=2))
print(f"Wrote {args.output}")
if __name__ == "__main__":
main()
+68
View File
@@ -0,0 +1,68 @@
"""Score one video on one GPU.
Smallest possible use of ``fastvideo.eval``: load an mp4, build an
:class:`Evaluator` for the requested metric set, run it.
Examples::
# Reference-free (just the generated video):
python examples/inference/eval/score_video.py \\
--video clip.mp4 \\
--metrics vbench.aesthetic_quality,vbench.imaging_quality
# Reference-paired (compare against ground truth):
python examples/inference/eval/score_video.py \\
--video gen.mp4 --reference ref.mp4 \\
--metrics common.psnr,common.ssim,common.lpips
"""
from __future__ import annotations
import argparse
import json
from fastvideo.eval import create_evaluator
from fastvideo.eval.io import load_video
def main() -> None:
p = argparse.ArgumentParser(description=__doc__,
formatter_class=argparse.RawDescriptionHelpFormatter)
p.add_argument("--video", required=True, help="Path to the generated mp4.")
p.add_argument("--reference", default=None,
help="Optional path to a reference mp4 (for paired metrics).")
p.add_argument("--metrics", default="common.psnr,common.ssim",
help="Comma-separated metric names, or a group name like 'vbench'.")
p.add_argument("--device", default="cuda:0")
p.add_argument("--text-prompt", default=None,
help="Text prompt for prompt-aware metrics "
"(vbench.overall_consistency, etc.).")
p.add_argument("--fps", type=float, default=None,
help="Frame-rate annotation for fps-aware metrics "
"(vbench.dynamic_degree, etc.).")
args = p.parse_args()
metrics: list[str] | str = (
args.metrics if args.metrics in ("all",)
else [m.strip() for m in args.metrics.split(",") if m.strip()]
)
evaluator = create_evaluator(metrics=metrics, device=args.device)
sample: dict = {"video": load_video(args.video)}
if args.reference is not None:
sample["reference"] = load_video(args.reference)
if args.text_prompt is not None:
sample["text_prompt"] = args.text_prompt
if args.fps is not None:
sample["fps"] = args.fps
results = evaluator.evaluate(**sample)
evaluator.shutdown()
print(json.dumps(
{name: r.score for name, r in results.items()},
indent=2,
))
if __name__ == "__main__":
main()
@@ -5,7 +5,7 @@ import time
import gradio as gr
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
from copy import deepcopy
@@ -9,7 +9,7 @@ import tempfile
import gradio as gr
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
MODEL_PATH_MAPPING = {
@@ -185,7 +185,7 @@ class BaseModelDeployment:
def _initialize_generator(self, config: Dict[str, Any]) -> None:
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
print(f"Initializing model: {self.model_path}")
self.generator = VideoGenerator.from_pretrained(
@@ -1,5 +1,5 @@
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "./lora_out"
def main():
@@ -2,7 +2,7 @@
Inference using a LoRA checkpoint from FastVideo trainer.
"""
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
OUTPUT_PATH = "./lora_out"
def main():
+39 -1
View File
@@ -10,6 +10,35 @@ set -ex
echo "Building fastvideo-kernel..."
# ---------------------------------------------------------------------------
# Neutralise conda-injected compiler toolchains.
#
# Conda compiler packages (gcc_linux-aarch64, gxx_linux-64, etc.) set
# CMAKE_ARGS, CFLAGS, CXXFLAGS, and LDFLAGS on activation. When multiple
# toolchains are installed the variables can reference a *cross*-compiler
# that doesn't match the host (e.g. aarch64-conda-linux-gnu-c++ on x86_64).
# Even when the correct toolchain is active, the flags it injects
# (-march=nocona, -mtune=haswell, …) can conflict with nvcc's host-compiler
# expectations. Clear them so CMake discovers the system compiler instead.
# ---------------------------------------------------------------------------
if [[ -n "${CONDA_PREFIX:-}" ]]; then
_need_clean=0
# Detect conda cross-compiler that doesn't match the host.
_host_arch="$(uname -m)"
if [[ "${CXX:-}" == *"conda"* ]] || [[ "${CC:-}" == *"conda"* ]]; then
_need_clean=1
fi
if [[ "${CMAKE_ARGS:-}" == *"conda"* ]]; then
_need_clean=1
fi
if (( _need_clean )); then
echo "NOTE: Clearing conda-injected compiler settings (CC/CXX/CMAKE_ARGS/CFLAGS/...)"
echo " to use the system compiler for CUDA extension builds."
unset CC CXX CMAKE_ARGS CFLAGS CXXFLAGS LDFLAGS
fi
unset _need_clean _host_arch
fi
# Ensure submodules are initialized if needed (tk)
git submodule update --init --recursive
@@ -32,7 +61,16 @@ has_cmake_arg() {
}
detect_with_torch() {
uv run --active --no-project python -c "import torch
# Prefer the active venv's python directly over `uv run --active --no-project`,
# which on some uv versions provisions its own interpreter and misses packages
# installed into VIRTUAL_ENV.
local py
if [[ -n "${VIRTUAL_ENV:-}" && -x "${VIRTUAL_ENV}/bin/python" ]]; then
py="${VIRTUAL_ENV}/bin/python"
else
py="$(command -v python3 || command -v python)"
fi
"${py}" -c "import torch
if not torch.cuda.is_available():
raise RuntimeError('torch.cuda.is_available() is false')
mj, mn = torch.cuda.get_device_capability(0)
@@ -5,6 +5,11 @@ from fastvideo_kernel.ops import (
video_sparse_attn,
)
from fastvideo_kernel.block_sparse_attn import (
block_sparse_attn,
block_sparse_attn_from_indices,
)
from fastvideo_kernel.vmoba import (
moba_attn_varlen,
process_moba_input,
@@ -22,6 +27,8 @@ from fastvideo_kernel.turbodiffusion_ops import (
__all__ = [
"sliding_tile_attention",
"video_sparse_attn",
"block_sparse_attn",
"block_sparse_attn_from_indices",
"moba_attn_varlen",
"process_moba_input",
"process_moba_output",
@@ -1,3 +1,5 @@
"""Autograd-enabled block-sparse attention. Index-native ops with a bool-mask compat shim."""
from __future__ import annotations
import os
@@ -6,6 +8,11 @@ from typing import Tuple
import torch
# ---------------------------------------------------------------------------
# Backend selection helpers
# ---------------------------------------------------------------------------
def _get_sm90_ops():
try:
from fastvideo_kernel._C import fastvideo_kernel_ops # type: ignore
@@ -25,38 +32,66 @@ def _is_sm90() -> bool:
def _force_triton() -> bool:
# Force Triton even on SM90 and even if the compiled extension is available.
# Useful for CI / debugging / parity testing.
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Preferred map->index conversion used by the wrapper.
# ---------------------------------------------------------------------------
# Index helpers
# ---------------------------------------------------------------------------
This wrapper **requires** the Triton implementation.
If Triton (or the Triton map_to_index module) is not available, it raises.
"""
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""Compact a bool block_map to (q2k_idx, q2k_num). Legacy path only."""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
if block_map.dim() != 4:
raise ValueError(f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), got shape={tuple(block_map.shape)}")
raise ValueError(
f"block_map must be [B,H,Q,KV] (or [H,Q,KV]), "
f"got shape={tuple(block_map.shape)}"
)
if block_map.dtype != torch.bool:
block_map = block_map.to(torch.bool)
if not block_map.is_cuda:
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
raise RuntimeError(
"block_map must be a CUDA tensor (Triton map_to_index required)."
)
try:
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
except Exception as e:
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index
except Exception as e: # pragma: no cover - environment issue
raise ImportError(
"Triton map_to_index is required but not available. "
"Ensure Triton is installed and fastvideo_kernel.triton_kernels.index is importable."
"Ensure Triton is installed and "
"fastvideo_kernel.triton_kernels.index is importable."
) from e
return triton_map_to_index(block_map)
def _invert_indices_for_backward(
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
num_kv_blocks: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
from fastvideo_kernel.triton_kernels.index import invert_indices
return invert_indices(q2k_idx, q2k_num, num_kv_blocks=num_kv_blocks)
def _as_int32_contig(t: torch.Tensor, name: str) -> torch.Tensor:
"""Return `t` as a contiguous int32 tensor, raising a clear error on CPU input."""
if not t.is_cuda:
raise RuntimeError(f"{name} must be a CUDA tensor, got device={t.device}")
if t.dtype != torch.int32:
t = t.to(torch.int32)
if not t.is_contiguous():
t = t.contiguous()
return t
# ---------------------------------------------------------------------------
# Triton backend custom ops (index-native)
# ---------------------------------------------------------------------------
@torch.library.custom_op(
"fastvideo_kernel::block_sparse_attn_triton",
mutates_args=(),
@@ -66,34 +101,40 @@ def block_sparse_attn_triton(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
q = q.contiguous()
k = k.contiguous()
v = v.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
triton_block_sparse_attn_forward,
)
o, M = triton_block_sparse_attn_forward(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
o, M = triton_block_sparse_attn_forward(
q.contiguous(),
k.contiguous(),
v.contiguous(),
q2k_idx,
q2k_num,
variable_block_sizes,
)
return o, M
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
o = torch.empty_like(q)
M = torch.empty((q.shape[0], q.shape[1], q.shape[2]), device=q.device, dtype=torch.float32)
M = torch.empty(
(q.shape[0], q.shape[1], q.shape[2]),
device=q.device,
dtype=torch.float32,
)
return o, M
@@ -109,20 +150,32 @@ def block_sparse_attn_backward_triton(
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output = grad_output.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import (
triton_block_sparse_attn_backward,
)
num_kv_blocks = int(variable_block_sizes.numel())
k2q_idx, k2q_num = _invert_indices_for_backward(
q2k_idx, q2k_num, num_kv_blocks
)
# q/k/v are saved from the user-facing inputs and may be non-contiguous;
# o/M are kernel outputs so are already contiguous.
dq, dk, dv = triton_block_sparse_attn_backward(
grad_output, q, k, v, o, M, q2k_idx, q2k_num, k2q_idx, k2q_num, variable_block_sizes
grad_output.contiguous(),
q.contiguous(),
k.contiguous(),
v.contiguous(),
o,
M,
q2k_idx,
q2k_num,
k2q_idx,
k2q_num,
variable_block_sizes,
)
return dq, dk, dv
@@ -135,7 +188,8 @@ def _block_sparse_attn_backward_triton_fake(
v: torch.Tensor,
o: torch.Tensor,
M: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q)
@@ -144,19 +198,28 @@ def _block_sparse_attn_backward_triton_fake(
return dq, dk, dv
def _backward_triton(ctx, grad_o, grad_M):
q, k, v, o, M, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_triton(grad_o, q, k, v, o, M, block_map, variable_block_sizes)
return dq, dk, dv, None, None
def _setup_context_triton(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, M = output
ctx.save_for_backward(q, k, v, o, M, block_map, variable_block_sizes)
ctx.save_for_backward(q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes)
block_sparse_attn_triton.register_autograd(_backward_triton, setup_context=_setup_context_triton)
def _backward_triton(ctx, grad_o, grad_M):
q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_triton(
grad_o, q, k, v, o, M, q2k_idx, q2k_num, variable_block_sizes
)
return dq, dk, dv, None, None, None
block_sparse_attn_triton.register_autograd(
_backward_triton, setup_context=_setup_context_triton
)
# ---------------------------------------------------------------------------
# SM90 backend custom ops (index-native)
# ---------------------------------------------------------------------------
@torch.library.custom_op(
@@ -168,21 +231,21 @@ def block_sparse_attn_sm90(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
block_sparse_fwd, _ = _get_sm90_ops()
if block_sparse_fwd is None:
raise ImportError("fastvideo_kernel_ops.block_sparse_fwd is not available")
q_padded = q_padded.contiguous()
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index(block_map)
o_padded, lse_padded = block_sparse_fwd(
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
q_padded.contiguous(),
k_padded.contiguous(),
v_padded.contiguous(),
q2k_idx,
q2k_num,
variable_block_sizes,
)
return o_padded, lse_padded
@@ -192,11 +255,16 @@ def _block_sparse_attn_sm90_fake(
q_padded: torch.Tensor,
k_padded: torch.Tensor,
v_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
o = torch.empty_like(q_padded)
lse = torch.empty((q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1), device=q_padded.device, dtype=torch.float32)
lse = torch.empty(
(q_padded.shape[0], q_padded.shape[1], q_padded.shape[2], 1),
device=q_padded.device,
dtype=torch.float32,
)
return o, lse
@@ -212,30 +280,34 @@ def block_sparse_attn_backward_sm90(
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
_, block_sparse_bwd = _get_sm90_ops()
if block_sparse_bwd is None:
raise ImportError("fastvideo_kernel_ops.block_sparse_bwd is not available")
grad_output_padded = grad_output_padded.contiguous()
block_map = block_map.to(torch.bool)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
num_kv_blocks = int(variable_block_sizes.numel())
k2q_idx, k2q_num = _invert_indices_for_backward(
q2k_idx, q2k_num, num_kv_blocks
)
# q/k/v are saved from user-facing inputs; o/lse are kernel outputs.
dq, dk, dv = block_sparse_bwd(
q_padded,
k_padded,
v_padded,
q_padded.contiguous(),
k_padded.contiguous(),
v_padded.contiguous(),
o_padded,
lse_padded,
grad_output_padded,
grad_output_padded.contiguous(),
k2q_idx,
k2q_num,
variable_block_sizes.int(),
variable_block_sizes,
)
# C++ kernel returns fp32 grads; cast back to match PyTorch convention if needed
return dq.to(grad_output_padded.dtype), dk.to(grad_output_padded.dtype), dv.to(grad_output_padded.dtype)
# C++ kernel returns fp32 grads; cast back to the input dtype.
out_dtype = grad_output_padded.dtype
return dq.to(out_dtype), dk.to(out_dtype), dv.to(out_dtype)
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_backward_sm90")
@@ -246,7 +318,8 @@ def _block_sparse_attn_backward_sm90_fake(
v_padded: torch.Tensor,
o_padded: torch.Tensor,
lse_padded: torch.Tensor,
block_map: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
dq = torch.empty_like(q_padded)
@@ -255,21 +328,57 @@ def _block_sparse_attn_backward_sm90_fake(
return dq, dk, dv
def _backward_sm90(ctx, grad_o, grad_lse):
q, k, v, o, lse, block_map, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_sm90(
grad_o, q, k, v, o, lse, block_map, variable_block_sizes
)
return dq, dk, dv, None, None
def _setup_context_sm90(ctx, inputs, output):
q, k, v, block_map, variable_block_sizes = inputs
q, k, v, q2k_idx, q2k_num, variable_block_sizes = inputs
o, lse = output
ctx.save_for_backward(q, k, v, o, lse, block_map, variable_block_sizes)
ctx.save_for_backward(q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes)
block_sparse_attn_sm90.register_autograd(_backward_sm90, setup_context=_setup_context_sm90)
def _backward_sm90(ctx, grad_o, grad_lse):
q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes = ctx.saved_tensors
dq, dk, dv = block_sparse_attn_backward_sm90(
grad_o, q, k, v, o, lse, q2k_idx, q2k_num, variable_block_sizes
)
return dq, dk, dv, None, None, None
block_sparse_attn_sm90.register_autograd(
_backward_sm90, setup_context=_setup_context_sm90
)
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def block_sparse_attn_from_indices(
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Block-sparse attention with autograd, taking compact per-row KV indices."""
# Normalize index tensors once at the public boundary so the custom ops
# and their fakes can assume int32/contiguous. No-op on well-formed input.
q2k_idx = _as_int32_contig(q2k_idx, "q2k_idx")
q2k_num = _as_int32_contig(q2k_num, "q2k_num")
variable_block_sizes = _as_int32_contig(variable_block_sizes, "variable_block_sizes")
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
use_sm90 = (
(not _force_triton())
and _is_sm90()
and block_sparse_fwd is not None
and block_sparse_bwd is not None
)
if use_sm90:
return block_sparse_attn_sm90(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
# to a multiple of the block size (64 tokens).
return block_sparse_attn_triton(q, k, v, q2k_idx, q2k_num, variable_block_sizes)
def block_sparse_attn(
@@ -279,16 +388,8 @@ def block_sparse_attn(
block_map: torch.Tensor,
variable_block_sizes: torch.Tensor,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Unified block-sparse attention op with autograd support.
- On SM90 with compiled extension present: uses fastvideo_kernel_ops.block_sparse_fwd/bwd.
- Otherwise: uses Triton implementation (requires q/k/v to have same padded length today).
"""
block_sparse_fwd, block_sparse_bwd = _get_sm90_ops()
if (not _force_triton()) and _is_sm90() and (block_sparse_fwd is not None) and (block_sparse_bwd is not None):
return block_sparse_attn_sm90(q, k, v, block_map, variable_block_sizes)
# Triton path: supports q_seq_len != kv_seq_len as long as both are padded
# to a multiple of the block size (64 tokens).
return block_sparse_attn_triton(q, k, v, block_map, variable_block_sizes)
"""Bool-mask compat wrapper; prefer block_sparse_attn_from_indices."""
q2k_idx, q2k_num = _map_to_index(block_map)
return block_sparse_attn_from_indices(
q, k, v, q2k_idx, q2k_num, variable_block_sizes
)
@@ -1,6 +1,6 @@
import math
import torch
from .block_sparse_attn import block_sparse_attn
from .block_sparse_attn import block_sparse_attn, block_sparse_attn_from_indices
from .triton_kernels.st_attn_triton import sliding_tile_attention_triton
# Try to load the C++ extension
@@ -125,13 +125,18 @@ def video_sparse_attn(
out_c = out_c.repeat(1, 1, 1, block_elements,
1).view(batch, heads, q_seq_len, dim)
# Sparse branch
# Sparse branch: feed top-k indices directly, skipping the bool-mask round-trip.
topk_idx = torch.topk(scores, topk, dim=-1).indices
mask = torch.zeros_like(scores,
dtype=torch.bool).scatter_(-1, topk_idx, True)
# out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
out_s = block_sparse_attn(q, k, v, mask, variable_block_sizes)[0]
q2k_idx = topk_idx.to(torch.int32).contiguous()
q2k_num = torch.full(
(batch, heads, q_num_blocks),
topk,
dtype=torch.int32,
device=q.device,
)
out_s = block_sparse_attn_from_indices(
q, k, v, q2k_idx, q2k_num, variable_block_sizes
)[0]
if compress_attn_weight is not None:
return out_c * compress_attn_weight + out_s
@@ -1,9 +1,10 @@
## pytorch sdpa version of block sparse ##
from typing import Tuple
import triton
import triton.language as tl
import torch
@triton.jit
def topk_index_to_map_kernel(
map_ptr,
@@ -153,3 +154,114 @@ def map_to_index(block_map: torch.Tensor):
)
return index, index_num
@triton.jit
def _invert_indices_kernel(
q2k_idx_ptr,
q2k_num_ptr,
k2q_idx_ptr,
k2q_num_ptr,
q2k_idx_b, q2k_idx_h, q2k_idx_q, q2k_idx_k,
q2k_num_b, q2k_num_h, q2k_num_q,
k2q_idx_b, k2q_idx_h, k2q_idx_k, k2q_idx_q,
k2q_num_b, k2q_num_h, k2q_num_k,
MAX_KV_PER_Q: tl.constexpr,
):
# One program per (b, h, q): reserve a slot in k2q via atomicAdd, write q.
pid_b = tl.program_id(0)
pid_h = tl.program_id(1)
pid_q = tl.program_id(2)
n = tl.load(
q2k_num_ptr
+ pid_b * q2k_num_b
+ pid_h * q2k_num_h
+ pid_q * q2k_num_q
)
q2k_row = (
q2k_idx_ptr
+ pid_b * q2k_idx_b
+ pid_h * q2k_idx_h
+ pid_q * q2k_idx_q
)
for i in tl.range(0, MAX_KV_PER_Q):
if i < n:
kv = tl.load(q2k_row + i * q2k_idx_k)
count_ptr = (
k2q_num_ptr
+ pid_b * k2q_num_b
+ pid_h * k2q_num_h
+ kv * k2q_num_k
)
pos = tl.atomic_add(count_ptr, 1)
tl.store(
k2q_idx_ptr
+ pid_b * k2q_idx_b
+ pid_h * k2q_idx_h
+ kv * k2q_idx_k
+ pos * k2q_idx_q,
pid_q,
)
def invert_indices(
q2k_idx: torch.Tensor,
q2k_num: torch.Tensor,
num_kv_blocks: int,
) -> Tuple[torch.Tensor, torch.Tensor]:
"""Transpose a Q->KV index list into a K->Q one via atomic compaction (GPU)."""
if q2k_idx.dim() != 4:
raise ValueError(
f"q2k_idx must be [B, H, Nq, Mk], got shape={tuple(q2k_idx.shape)}"
)
if q2k_num.dim() != 3:
raise ValueError(
f"q2k_num must be [B, H, Nq], got shape={tuple(q2k_num.shape)}"
)
if not q2k_idx.is_cuda or not q2k_num.is_cuda:
raise RuntimeError("invert_indices requires CUDA tensors.")
B, H, Nq, Mk = q2k_idx.shape
if q2k_num.shape != (B, H, Nq):
raise ValueError(
f"q2k_num shape {tuple(q2k_num.shape)} does not match q2k_idx "
f"[B, H, Nq] = {(B, H, Nq)}"
)
q2k_idx = q2k_idx.contiguous()
q2k_num = q2k_num.contiguous()
if q2k_idx.dtype != torch.int32:
q2k_idx = q2k_idx.to(torch.int32)
if q2k_num.dtype != torch.int32:
q2k_num = q2k_num.to(torch.int32)
# Any KV block is attended by at most Nq Q blocks (one per Q row), so
# `Nq` is a tight upper bound on the compacted K->Q slots.
k2q_idx = torch.empty(
(B, H, num_kv_blocks, Nq),
dtype=torch.int32,
device=q2k_idx.device,
)
k2q_num = torch.zeros(
(B, H, num_kv_blocks),
dtype=torch.int32,
device=q2k_idx.device,
)
grid = (B, H, Nq)
_invert_indices_kernel[grid](
q2k_idx,
q2k_num,
k2q_idx,
k2q_num,
q2k_idx.stride(0), q2k_idx.stride(1), q2k_idx.stride(2), q2k_idx.stride(3),
q2k_num.stride(0), q2k_num.stride(1), q2k_num.stride(2),
k2q_idx.stride(0), k2q_idx.stride(1), k2q_idx.stride(2), k2q_idx.stride(3),
k2q_num.stride(0), k2q_num.stride(1), k2q_num.stride(2),
MAX_KV_PER_Q=Mk,
)
return k2q_idx, k2q_num
+1 -1
View File
@@ -1,5 +1,5 @@
from fastvideo.configs.pipelines import PipelineConfig
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.version import __version__
+32
View File
@@ -7,21 +7,37 @@ from fastvideo.api.schema import (
GenerationPlan,
GenerationRequest,
GeneratorConfig,
GpuPoolConfig,
InputConfig,
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
PlannedStage,
PromptEnhancerConfig,
PromptSafetyConfig,
QuantizationConfig,
RequestRuntimeConfig,
RunConfig,
SamplingConfig,
ServeConfig,
ServerConfig,
StreamingConfig,
WarmupConfig,
)
from fastvideo.api.errors import ConfigValidationError
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
from fastvideo.api.presets import (
InferencePreset,
PresetStageSpec,
get_all_preset_names,
get_preset,
get_presets_for_family,
register_preset,
validate_preset_selection,
validate_stage_names,
validate_stage_overrides,
)
from fastvideo.api.parser import (
config_to_dict,
load_config,
@@ -31,6 +47,7 @@ from fastvideo.api.parser import (
parse_config,
)
from fastvideo.api.results import GenerationResult
from fastvideo.api.sampling_param import SamplingParam
__all__ = [
"CompileConfig",
@@ -42,18 +59,26 @@ __all__ = [
"GenerationPlan",
"GenerationRequest",
"GeneratorConfig",
"GpuPoolConfig",
"InputConfig",
"OffloadConfig",
"OutputConfig",
"ParallelismConfig",
"PipelineSelection",
"PlannedStage",
"PromptEnhancerConfig",
"PromptSafetyConfig",
"QuantizationConfig",
"RequestRuntimeConfig",
"RunConfig",
"SamplingConfig",
"SamplingParam",
"ServeConfig",
"ServerConfig",
"StreamingConfig",
"WarmupConfig",
"InferencePreset",
"PresetStageSpec",
"apply_overrides",
"config_to_dict",
"load_config",
@@ -61,5 +86,12 @@ __all__ = [
"load_run_config",
"load_serve_config",
"parse_cli_overrides",
"get_all_preset_names",
"get_preset",
"get_presets_for_family",
"parse_config",
"register_preset",
"validate_preset_selection",
"validate_stage_names",
"validate_stage_overrides",
]
+216 -94
View File
@@ -7,9 +7,17 @@ from dataclasses import fields, is_dataclass
from pathlib import Path
from typing import Any
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
from fastvideo.api.overrides import apply_overrides, normalize_overrides
from fastvideo.api.parser import config_to_dict, load_raw_config, parse_config
from fastvideo.api.request_metadata import (
EXPLICIT_PATHS_ATTR,
bind_generation_request_raw,
get_explicit_paths,
reset_tracking_roots,
)
from fastvideo.api.schema import (
CompileConfig,
ContinuationState,
GenerationRequest,
GeneratorConfig,
InputConfig,
@@ -17,11 +25,14 @@ from fastvideo.api.schema import (
RequestRuntimeConfig,
SamplingConfig,
)
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.pipelines.basic.ltx2.stage_overrides import (
refine_preset_override_fields,
refine_stage_override_fields,
)
from fastvideo.utils import shallow_asdict
_EXPLICIT_REQUEST_ATTR = "_fastvideo_explicit_request"
_INPUT_FIELD_NAMES = {field.name for field in fields(InputConfig)}
_SAMPLING_FIELD_NAMES = {field.name for field in fields(SamplingConfig)}
_RUNTIME_FIELD_NAMES = {field.name for field in fields(RequestRuntimeConfig)}
@@ -33,6 +44,10 @@ _LEGACY_REQUEST_ALIASES = {
_REQUEST_PIPELINE_OVERRIDE_FIELDS = frozenset({
"embedded_cfg_scale",
})
# torch.compile kwargs that map to first-class CompileConfig fields.
_COMPILE_TYPED_KEYS = ("backend", "fullgraph", "mode", "dynamic")
# LTX-2 refine flat kwargs (init + per-request) known to FastVideoArgs.
_LTX2_REFINE_FLAT_KEYS = (refine_preset_override_fields() | refine_stage_override_fields())
def normalize_generator_config(config: GeneratorConfig | Mapping[str, Any], ) -> GeneratorConfig:
@@ -46,7 +61,7 @@ def load_generator_config_from_file(
overrides: list[str] | Mapping[str, Any] | None = None,
) -> GeneratorConfig:
raw = load_raw_config(path)
normalized_overrides = _normalize_overrides(overrides)
normalized_overrides = normalize_overrides(overrides)
if _looks_like_run_or_serve_config(raw):
if normalized_overrides:
@@ -75,6 +90,8 @@ def legacy_from_pretrained_to_config(
components: dict[str, Any] = {}
quantization: dict[str, Any] = {}
experimental: dict[str, Any] = {}
preset_overrides: dict[str, Any] = {}
preset_refine: dict[str, Any] = {}
for key, value in kwargs.items():
if key == "revision":
@@ -101,8 +118,33 @@ def legacy_from_pretrained_to_config(
offload["pin_cpu_memory"] = value
elif key == "enable_torch_compile":
compile_config["enabled"] = value
elif key == "enable_torch_compile_text_encoder":
compile_config["text_encoder_enabled"] = value
elif key == "torch_compile_kwargs":
compile_config["kwargs"] = deepcopy(value)
remaining: dict[str, Any] = (dict(deepcopy(value)) if isinstance(value, Mapping) else {})
for first_class in _COMPILE_TYPED_KEYS:
if first_class in remaining:
compile_config[first_class] = remaining.pop(first_class)
if remaining:
compile_config["extras"] = remaining
elif key == "ltx2_vae_tiling":
pipeline["vae_tiling"] = value
elif key == "config_model_path":
components["config_root"] = value
elif key == "ltx2_refine_enabled":
preset_refine["enabled"] = value
elif key == "ltx2_refine_upsampler_path":
# Empty string means "no upsampler"; keep typed None.
components["upsampler_weights"] = value or None
elif key == "ltx2_refine_lora_path":
# Empty string means "no refine LoRA"; keep typed None.
components["lora_path"] = value or None
elif key == "ltx2_refine_add_noise":
preset_refine["add_noise"] = value
elif key == "ltx2_refine_num_inference_steps":
preset_refine["num_inference_steps"] = value
elif key == "ltx2_refine_guidance_scale":
preset_refine["guidance_scale"] = value
elif key in {"enable_stage_verification", "use_fsdp_inference", "disable_autocast"}:
engine[key] = value
elif key == "override_text_encoder_quant":
@@ -142,6 +184,10 @@ def legacy_from_pretrained_to_config(
if components:
pipeline["components"] = components
if preset_refine:
preset_overrides["refine"] = preset_refine
if preset_overrides:
pipeline["preset_overrides"] = preset_overrides
if experimental:
pipeline["experimental"] = experimental
if pipeline:
@@ -153,16 +199,12 @@ def legacy_from_pretrained_to_config(
def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, Any], ) -> FastVideoArgs:
normalized = normalize_generator_config(config)
unsupported = []
if normalized.pipeline.profile is not None:
unsupported.append("pipeline.profile")
if normalized.pipeline.profile_version is not None:
unsupported.append("pipeline.profile_version")
if normalized.pipeline.components.config_root is not None:
unsupported.append("pipeline.components.config_root")
if normalized.pipeline.preset is not None:
unsupported.append("pipeline.preset")
if normalized.pipeline.preset_version is not None:
unsupported.append("pipeline.preset_version")
if normalized.pipeline.components.vae_weights is not None:
unsupported.append("pipeline.components.vae_weights")
if normalized.pipeline.components.upsampler_weights is not None:
unsupported.append("pipeline.components.upsampler_weights")
if unsupported:
joined = ", ".join(unsupported)
raise NotImplementedError(f"VideoGenerator compatibility adapter does not support {joined} yet")
@@ -186,13 +228,21 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
"vae_cpu_offload": engine.offload.vae,
"pin_cpu_memory": engine.offload.pin_cpu_memory,
"enable_torch_compile": engine.compile.enabled,
"torch_compile_kwargs": deepcopy(engine.compile.kwargs),
"torch_compile_kwargs": _compile_config_to_torch_kwargs(engine.compile),
"enable_stage_verification": engine.enable_stage_verification,
"use_fsdp_inference": engine.use_fsdp_inference,
"disable_autocast": engine.disable_autocast,
}
if normalized.pipeline.workload_type is not None:
kwargs["workload_type"] = normalized.pipeline.workload_type
if normalized.pipeline.vae_tiling is not None:
kwargs["ltx2_vae_tiling"] = normalized.pipeline.vae_tiling
if engine.compile.text_encoder_enabled is not None:
# ``FastVideoArgs.from_kwargs`` filters to declared fields, so
# this is a no-op on the current legacy path. Emit anyway so the
# realtime runtime (PR 7.6) — which reads from the kwargs dict
# before FastVideoArgs filtering — can pick it up once wired.
kwargs["enable_torch_compile_text_encoder"] = (engine.compile.text_encoder_enabled)
quantization = engine.quantization
if quantization is not None and quantization.text_encoder_quant is not None:
@@ -215,8 +265,18 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
kwargs["init_weights_from_safetensors"] = components.transformer_weights
if components.transformer_2_weights is not None:
kwargs["init_weights_from_safetensors_2"] = components.transformer_2_weights
if components.config_root is not None:
kwargs["config_model_path"] = components.config_root
if components.upsampler_weights is not None:
kwargs["ltx2_refine_upsampler_path"] = components.upsampler_weights
kwargs.update(deepcopy(normalized.pipeline.profile_overrides))
preset_overrides = deepcopy(normalized.pipeline.preset_overrides)
refine = preset_overrides.pop("refine", None)
if isinstance(refine, Mapping):
for key in _LTX2_REFINE_FLAT_KEYS:
if key in refine:
kwargs[f"ltx2_refine_{key}"] = refine[key]
kwargs.update(preset_overrides)
kwargs.update(deepcopy(normalized.pipeline.experimental))
return FastVideoArgs.from_kwargs(**kwargs)
@@ -224,8 +284,10 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An
def normalize_generation_request(request: GenerationRequest | Mapping[str, Any], ) -> GenerationRequest:
normalized = (request if isinstance(request, GenerationRequest) else parse_config(GenerationRequest, request))
if not hasattr(normalized, _EXPLICIT_REQUEST_ATTR):
setattr(normalized, _EXPLICIT_REQUEST_ATTR, _serialize_generation_request(normalized))
if not hasattr(normalized, EXPLICIT_PATHS_ATTR):
# Request wasn't bound through the parser (e.g. constructed
# directly). Treat every currently-set field as explicit.
bind_generation_request_raw(normalized, _serialize_generation_request(normalized))
return normalized
@@ -253,7 +315,7 @@ def legacy_generate_call_to_request(
raw.setdefault("inputs", {})["grid_sizes"] = grid_sizes
normalized = parse_config(GenerationRequest, raw)
setattr(normalized, _EXPLICIT_REQUEST_ATTR, deepcopy(raw))
bind_generation_request_raw(normalized, raw)
return normalized
@@ -264,16 +326,24 @@ def request_to_sampling_param(
) -> SamplingParam:
if request.plan is not None:
raise NotImplementedError("GenerationRequest.plan is not wired into VideoGenerator yet")
if request.state is not None:
raise NotImplementedError("GenerationRequest.state is not wired into VideoGenerator yet")
sampling_param = SamplingParam.from_pretrained(model_path)
updates = _explicit_request_updates(request)
if request.state is not None:
_validate_continuation_state(request.state)
sampling_param.continuation_state = request.state
if request.output.return_state:
sampling_param.return_continuation_state = True
updates = explicit_request_updates(request)
for key, value in updates.items():
if hasattr(sampling_param, key):
setattr(sampling_param, key, deepcopy(value))
elif key in _REQUEST_PIPELINE_OVERRIDE_FIELDS or _is_supported_as_default_only(key, value):
elif key in _REQUEST_PIPELINE_OVERRIDE_FIELDS:
continue
elif value == _SCHEMA_DEFAULT_UPDATES.get(key, _MISSING):
# Schema-default field that isn't on SamplingParam; tolerated
# because direct GenerationRequest(...) construction has no
# way to distinguish "user set" from "schema default".
continue
else:
raise ValueError(f"Request field {key!r} is not supported by sampling params for {model_path}")
@@ -290,10 +360,12 @@ def expand_request_prompt_batch(request: GenerationRequest, ) -> list[Generation
requests: list[GenerationRequest] = []
for index, prompt in enumerate(request.prompt):
single_request = deepcopy(request)
# deepcopy preserves the tracking-root cycle, but re-pin roots
# defensively so that subsequent setattrs record on the copy.
reset_tracking_roots(single_request)
single_request.prompt = prompt
_fan_out_batched_input_value(request, single_request, "image_path", index)
_fan_out_batched_input_value(request, single_request, "video_path", index)
_fan_out_explicit_request_metadata(request, single_request, index, prompt)
requests.append(single_request)
return requests
@@ -302,12 +374,23 @@ def _looks_like_run_or_serve_config(raw: Mapping[str, Any]) -> bool:
return isinstance(raw.get("generator"), Mapping)
def _normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
if not overrides:
return None
if isinstance(overrides, list):
return parse_cli_overrides(overrides)
return dict(overrides)
def _compile_config_to_torch_kwargs(compile_config: CompileConfig, ) -> dict[str, Any]:
"""Flatten typed ``CompileConfig`` back to a ``torch_compile_kwargs``
dict that the legacy ``FastVideoArgs`` path still expects.
Typed first-class fields (:attr:`backend`, :attr:`fullgraph`,
:attr:`mode`, :attr:`dynamic`) are only emitted when the user set
them explicitly (non-``None``). ``extras`` is merged on top for any
uncommon kwargs.
"""
out: dict[str, Any] = {}
for key in _COMPILE_TYPED_KEYS:
value = getattr(compile_config, key)
if value is not None:
out[key] = value
if compile_config.extras:
out.update(deepcopy(compile_config.extras))
return out
def _sampling_param_to_request_raw(sampling_param: SamplingParam | None, ) -> dict[str, Any]:
@@ -348,20 +431,83 @@ def _apply_request_field(
def request_to_pipeline_overrides(request: GenerationRequest) -> dict[str, Any]:
overrides: dict[str, Any] = {}
for key, value in _explicit_request_updates(request).items():
for key, value in explicit_request_updates(request).items():
if key in _REQUEST_PIPELINE_OVERRIDE_FIELDS:
overrides[key] = deepcopy(value)
return overrides
def _explicit_request_updates(request: GenerationRequest) -> dict[str, Any]:
raw = getattr(request, _EXPLICIT_REQUEST_ATTR, None)
if raw is None:
raw = _serialize_generation_request(request)
def explicit_request_updates(request: GenerationRequest) -> dict[str, Any]:
"""Project a ``GenerationRequest`` down to *explicitly set* fields only.
Returns a flat kwargs dict suitable for merging into a generator call.
The projection uses ``_fastvideo_explicit_paths`` (populated during
``parse_config`` / raw binding) so schema defaults on the dataclass
are **not** emitted — only paths the caller/operator actually wrote.
This is what makes ``ServeConfig.default_request`` work as an
operator-pinned baseline rather than a full override: a YAML with just
``sampling.seed: 42`` yields ``{"seed": 42}``, not the full sampling
config with its 15 schema defaults.
Precondition: the request must carry ``_fastvideo_explicit_paths`` —
populated by :func:`fastvideo.api.parser.parse_config` or
:func:`fastvideo.api.compat.normalize_generation_request`. Calling on
a raw ``GenerationRequest()`` asserts.
"""
assert hasattr(request,
EXPLICIT_PATHS_ATTR), ("GenerationRequest reached explicit_request_updates without tracking; "
"every entry point must route through normalize_generation_request "
"or parse_config first")
paths = get_explicit_paths(request)
raw = _build_sparse_raw_from_paths(request, paths)
return _extract_request_updates(raw)
def _build_sparse_raw_from_paths(
request: GenerationRequest,
paths: frozenset[str],
) -> dict[str, Any]:
result: dict[str, Any] = {}
for path in paths:
parts = path.split(".")
value = _read_dotted_path(request, parts)
if value is _MISSING:
continue
_set_dotted_path(result, parts, deepcopy(value))
return result
def _read_dotted_path(obj: Any, parts: list[str]) -> Any:
for part in parts:
if is_dataclass(obj) and not isinstance(obj, type):
if not hasattr(obj, part):
return _MISSING
obj = getattr(obj, part)
elif isinstance(obj, Mapping):
if part not in obj:
return _MISSING
obj = obj[part]
else:
return _MISSING
return obj
def _set_dotted_path(
target: dict[str, Any],
parts: list[str],
value: Any,
) -> None:
cursor = target
for part in parts[:-1]:
nxt = cursor.get(part)
if not isinstance(nxt, dict):
nxt = {}
cursor[part] = nxt
cursor = nxt
cursor[parts[-1]] = value
def _extract_request_updates(raw: Mapping[str, Any]) -> dict[str, Any]:
updates: dict[str, Any] = {}
if "negative_prompt" in raw:
@@ -405,6 +551,40 @@ def _serialize_generation_request(request: GenerationRequest) -> dict[str, Any]:
return deepcopy(config_to_dict(request))
_SCHEMA_DEFAULT_UPDATES = _extract_request_updates(config_to_dict(GenerationRequest()))
_KNOWN_CONTINUATION_KINDS: set[str] = set()
def register_continuation_kind(kind: str) -> None:
"""Register a :class:`ContinuationState.kind` as recognized.
PR 7 wires the envelope through; per-kind payload deserializers live
with each model family (e.g. ``fastvideo.pipelines.basic.ltx2.
continuation.LTX2ContinuationState``). The registry lets the
public-API compat layer validate the kind early, before the state
reaches the pipeline.
"""
if not isinstance(kind, str) or not kind:
raise ValueError("ContinuationState kind must be a non-empty string")
_KNOWN_CONTINUATION_KINDS.add(kind)
def _validate_continuation_state(state: ContinuationState) -> None:
if not isinstance(state.kind, str) or not state.kind:
raise ValueError("GenerationRequest.state.kind must be a non-empty string; got "
f"{state.kind!r}")
if not isinstance(state.payload, Mapping):
raise ValueError(f"GenerationRequest.state.payload must be a mapping; got "
f"{type(state.payload).__name__}")
if state.kind not in _KNOWN_CONTINUATION_KINDS:
known = sorted(_KNOWN_CONTINUATION_KINDS)
raise ValueError(f"Unknown ContinuationState kind {state.kind!r}; registered "
f"kinds: {known}. Import the model family that owns this kind "
"(e.g. `import fastvideo.pipelines.basic.ltx2.continuation`) "
"to register it, or drop the state field.")
def _fan_out_batched_input_value(
source_request: GenerationRequest,
target_request: GenerationRequest,
@@ -418,29 +598,6 @@ def _fan_out_batched_input_value(
setattr(target_request.inputs, field_name, deepcopy(value[index]))
def _fan_out_explicit_request_metadata(
source_request: GenerationRequest,
target_request: GenerationRequest,
index: int,
prompt: str,
) -> None:
raw = getattr(source_request, _EXPLICIT_REQUEST_ATTR, None)
if raw is None:
return
raw = deepcopy(raw)
raw["prompt"] = prompt
inputs = raw.get("inputs")
if isinstance(inputs, dict):
for field_name in ("image_path", "video_path"):
value = inputs.get(field_name)
if isinstance(value, list):
_validate_batched_input_length(source_request.prompt, value, field_name)
inputs[field_name] = deepcopy(value[index])
setattr(target_request, _EXPLICIT_REQUEST_ATTR, raw)
def _validate_batched_input_length(
prompts: str | list[str] | None,
values: list[Any],
@@ -452,50 +609,15 @@ def _validate_batched_input_length(
raise ValueError(f"GenerationRequest.inputs.{field_name} must have the same length as request.prompt")
def _is_supported_as_default_only(key: str, value: Any) -> bool:
default_value = _DEFAULT_REQUEST_UPDATES.get(key, _MISSING)
return default_value is not _MISSING and _values_equal(value, default_value)
def _collect_non_default_fields(
value: Any,
default: Any,
) -> dict[str, Any]:
if not (is_dataclass(value) and is_dataclass(default)):
return {}
result: dict[str, Any] = {}
for field in fields(value):
current = getattr(value, field.name)
default_value = getattr(default, field.name)
if is_dataclass(current) and is_dataclass(default_value):
nested = _collect_non_default_fields(current, default_value)
if nested:
result[field.name] = nested
continue
if not _values_equal(current, default_value):
result[field.name] = deepcopy(current)
return result
def _values_equal(left: Any, right: Any) -> bool:
if left is right:
return True
try:
return bool(left == right)
except Exception:
return False
_DEFAULT_REQUEST_UPDATES = _extract_request_updates(config_to_dict(GenerationRequest()))
__all__ = [
"explicit_request_updates",
"generator_config_to_fastvideo_args",
"legacy_from_pretrained_to_config",
"legacy_generate_call_to_request",
"load_generator_config_from_file",
"normalize_generation_request",
"normalize_generator_config",
"register_continuation_kind",
"request_to_pipeline_overrides",
"request_to_sampling_param",
]
+15 -2
View File
@@ -31,7 +31,7 @@ def parse_cli_overrides(overrides: list[str]) -> dict[str, Any]:
raise ValueError(f"Missing value for override {token!r}")
raw_value = overrides[index]
parsed[key] = _cast_override_value(raw_value)
parsed[_normalize_override_key(key)] = _cast_override_value(raw_value)
index += 1
return parsed
@@ -45,6 +45,15 @@ def apply_overrides(config: Mapping[str, Any], overrides: Mapping[str, Any]) ->
return merged
def normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
"""Normalize a CLI list or mapping of overrides into a flat dict."""
if not overrides:
return None
if isinstance(overrides, list):
return parse_cli_overrides(overrides)
return dict(overrides)
def _apply_single_override(config: dict[str, Any], dotted_key: str, value: Any) -> None:
parts = dotted_key.split(".")
if not all(parts):
@@ -94,4 +103,8 @@ def _cast_override_value(raw: str) -> Any:
return raw
__all__ = ["apply_overrides", "parse_cli_overrides"]
def _normalize_override_key(key: str) -> str:
return key.replace("-", "_")
__all__ = ["apply_overrides", "normalize_overrides", "parse_cli_overrides"]
+16 -12
View File
@@ -11,8 +11,13 @@ from typing import Any, Literal, TypeVar, Union, get_args, get_origin, get_type_
import yaml
from fastvideo.api.errors import ConfigValidationError
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
from fastvideo.api.schema import RunConfig, ServeConfig
from fastvideo.api.overrides import apply_overrides, normalize_overrides
from fastvideo.api.request_metadata import (
bind_generation_request_raw,
bind_run_config_raw,
bind_serve_config_raw,
)
from fastvideo.api.schema import GenerationRequest, RunConfig, ServeConfig
T = TypeVar("T")
_UNION_ORIGINS = {types.UnionType, Union}
@@ -31,7 +36,14 @@ def parse_config(config_type: type[T], raw: Mapping[str, Any] | T) -> T:
return raw
if not isinstance(raw, Mapping):
raise ConfigValidationError("", f"expected mapping for {config_type.__name__}")
return _SchemaParser().parse_dataclass(config_type, raw, "")
parsed = _SchemaParser().parse_dataclass(config_type, raw, "")
if config_type is GenerationRequest:
return bind_generation_request_raw(parsed, raw)
if config_type is RunConfig:
return bind_run_config_raw(parsed, raw)
if config_type is ServeConfig:
return bind_serve_config_raw(parsed, raw)
return parsed
def config_to_dict(config: Any) -> Any:
@@ -52,7 +64,7 @@ def load_config(
) -> T:
"""Load a typed config object from YAML or JSON."""
raw = load_raw_config(path)
normalized_overrides = _normalize_overrides(overrides)
normalized_overrides = normalize_overrides(overrides)
if normalized_overrides:
raw = apply_overrides(raw, normalized_overrides)
return parse_config(config_type, raw)
@@ -96,14 +108,6 @@ def _load_raw_mapping(handle: Any, config_path: Path) -> Any:
raise ValueError(f"Unsupported config file format: {config_path}")
def _normalize_overrides(overrides: list[str] | Mapping[str, Any] | None, ) -> dict[str, Any] | None:
if not overrides:
return None
if isinstance(overrides, list):
return parse_cli_overrides(overrides)
return dict(overrides)
class _SchemaParser:
def parse_dataclass(
+261
View File
@@ -0,0 +1,261 @@
# SPDX-License-Identifier: Apache-2.0
"""Pipeline preset registry.
A *preset* is a named inference preset for a model family. It bundles:
* ``defaults`` — sampling values applied when the user does not
override them (consumed at runtime via ``SamplingParam.from_pretrained``);
* ``stage_schemas`` — **validation-only** metadata describing which
user-facing stage names (``"denoise"``, ``"sr"``) the preset recognises
and which ``stage_overrides`` keys each stage accepts.
The ``stage_schemas`` tuple does **not** drive pipeline execution. The
concrete execution DAG (text encoding, denoising, VAE decoding, …) is
hard-coded per-pipeline in ``create_pipeline_stages()``. Schemas exist
purely so that ``PipelineSelection.preset`` and
``GenerationRequest.stage_overrides`` can be type-checked up front
without touching the pipeline.
Preset base types and the registry API live here (public API surface).
Preset *instances* are defined in pipeline-local ``presets.py`` files
(e.g. ``fastvideo/pipelines/basic/wan/presets.py``) and registered
explicitly from :func:`_register_presets` in ``fastvideo/registry.py``.
"""
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass, field
from typing import Any
from fastvideo.api.errors import ConfigValidationError
# -------------------------------------------------------------------
# Types
# -------------------------------------------------------------------
@dataclass(frozen=True)
class PresetStageSpec:
"""A user-facing stage name within a preset, used only to validate
``stage_overrides`` keys. Not read by pipeline execution — the real
execution DAG lives in each pipeline's ``create_pipeline_stages()``.
"""
name: str
"""Short user-facing name, e.g. ``"denoise"``, ``"sr"``."""
kind: str
"""Semantic kind, e.g. ``"denoising"``, ``"super_resolution"``."""
description: str = ""
allowed_overrides: frozenset[str] = field(default_factory=frozenset)
"""Keys that may appear in ``stage_overrides[name]``."""
@dataclass(frozen=True)
class InferencePreset:
"""A named inference preset for a model family."""
name: str
"""Preset name, e.g. ``"wan_t2v_1_3b"``."""
version: int
"""Preset schema version; bump on breaking schema changes."""
model_family: str
"""Model family key, e.g. ``"wan"``, ``"ltx2"``."""
description: str = ""
workload_type: str | None = None
"""Optional workload hint: ``"t2v"``, ``"i2v"``, etc."""
stage_schemas: tuple[PresetStageSpec, ...] = ()
"""User-facing stage names for ``stage_overrides`` validation.
Validation-only: this tuple is consumed by
:func:`validate_stage_overrides` and is **not** used to drive
pipeline execution. Omit or leave empty if the preset exposes no
per-stage override surface.
"""
defaults: dict[str, Any] = field(default_factory=dict)
"""Preset-level default sampling/runtime values."""
stage_defaults: dict[str, dict[str, Any]] = field(default_factory=dict)
"""Per-stage default overrides, keyed by stage name."""
# -------------------------------------------------------------------
# Registry
# -------------------------------------------------------------------
# Keyed by (model_family, name, version).
_PRESET_REGISTRY: dict[tuple[str, str, int], InferencePreset] = {}
def register_preset(preset: InferencePreset) -> None:
"""Register a preset definition.
Raises :class:`ValueError` on duplicate
``(model_family, name, version)`` keys.
"""
key = (preset.model_family, preset.name, preset.version)
if key in _PRESET_REGISTRY:
raise ValueError(f"Duplicate preset registration: "
f"model_family={key[0]!r}, name={key[1]!r}, "
f"version={key[2]!r}")
_PRESET_REGISTRY[key] = preset
def get_preset(
name: str,
model_family: str,
version: int | None = None,
) -> InferencePreset:
"""Look up a registered preset.
When *version* is ``None`` the highest registered version for the
given *(model_family, name)* pair is returned.
Raises :class:`~fastvideo.api.errors.ConfigValidationError` when the
preset cannot be found.
"""
if version is not None:
key = (model_family, name, version)
preset = _PRESET_REGISTRY.get(key)
if preset is not None:
return preset
raise ConfigValidationError(
"pipeline.preset",
f"unknown preset {name!r} version {version!r} "
f"for model family {model_family!r}; "
f"registered: {_format_registered(model_family)}",
)
# Find the highest version for (model_family, name).
candidates = [prof for (fam, n, _v), prof in _PRESET_REGISTRY.items() if fam == model_family and n == name]
if not candidates:
raise ConfigValidationError(
"pipeline.preset",
f"unknown preset {name!r} for model family "
f"{model_family!r}; "
f"registered: {_format_registered(model_family)}",
)
return max(candidates, key=lambda p: p.version)
def get_presets_for_family(model_family: str, ) -> list[InferencePreset]:
"""Return all presets registered for *model_family*."""
return [prof for (fam, _n, _v), prof in _PRESET_REGISTRY.items() if fam == model_family]
def get_all_preset_names() -> list[str]:
"""Return the sorted list of all registered preset names."""
return sorted({prof.name for prof in _PRESET_REGISTRY.values()})
# -------------------------------------------------------------------
# Validation helpers
# -------------------------------------------------------------------
def validate_stage_names(
preset: InferencePreset,
stage_overrides: Mapping[str, Any],
) -> None:
"""Check that *stage_overrides* keys are valid stage names.
Raises :class:`~fastvideo.api.errors.ConfigValidationError` with a
path-qualified message for unknown stage names.
"""
valid_names = {stage.name for stage in preset.stage_schemas}
for stage_name in stage_overrides:
if stage_name not in valid_names:
raise ConfigValidationError(
f"stage_overrides.{stage_name}",
f"unknown stage for preset {preset.name!r}; "
f"valid stages: {sorted(valid_names)}",
)
def validate_stage_overrides(
preset: InferencePreset,
stage_overrides: Mapping[str, Any],
) -> None:
"""Validate stage override keys against the preset.
Calls :func:`validate_stage_names` first, then checks that each
override key is in the stage's ``allowed_overrides``.
"""
validate_stage_names(preset, stage_overrides)
stages_by_name = {stage.name: stage for stage in preset.stage_schemas}
for stage_name, overrides in stage_overrides.items():
if not isinstance(overrides, Mapping):
raise ConfigValidationError(
f"stage_overrides.{stage_name}",
"must be a mapping",
)
stage_spec = stages_by_name[stage_name]
if not stage_spec.allowed_overrides:
if overrides:
raise ConfigValidationError(
f"stage_overrides.{stage_name}",
f"stage {stage_name!r} does not accept "
f"overrides",
)
continue
for key in overrides:
if key not in stage_spec.allowed_overrides:
raise ConfigValidationError(
f"stage_overrides.{stage_name}.{key}",
f"not an allowed override for stage "
f"{stage_name!r}; allowed: "
f"{sorted(stage_spec.allowed_overrides)}",
)
def validate_preset_selection(
preset_name: str | None,
model_family: str,
*,
preset_version: int | None = None,
stage_overrides: Mapping[str, Any] | None = None,
) -> InferencePreset | None:
"""Resolve and validate a preset selection end-to-end.
Returns the resolved :class:`InferencePreset`, or ``None`` if
*preset_name* is ``None`` (no preset requested).
"""
if preset_name is None:
return None
preset = get_preset(preset_name, model_family, version=preset_version)
if stage_overrides:
validate_stage_overrides(preset, stage_overrides)
return preset
# -------------------------------------------------------------------
# Internal helpers
# -------------------------------------------------------------------
def _format_registered(model_family: str) -> str:
names = sorted({prof.name for (fam, _n, _v), prof in _PRESET_REGISTRY.items() if fam == model_family})
if not names:
return "(none)"
return ", ".join(repr(n) for n in names)
__all__ = [
"InferencePreset",
"PresetStageSpec",
"get_all_preset_names",
"get_preset",
"get_presets_for_family",
"register_preset",
"validate_preset_selection",
"validate_stage_names",
"validate_stage_overrides",
]
+233
View File
@@ -0,0 +1,233 @@
# SPDX-License-Identifier: Apache-2.0
"""Track which GenerationRequest fields the user explicitly provided.
When translating a GenerationRequest into a legacy SamplingParam we must
distinguish user-provided values (which should override model defaults)
from schema defaults (which should NOT override model defaults).
The mechanism: a single ``_fastvideo_explicit_paths`` set stored on the
root ``GenerationRequest``. It holds dotted leaf paths (e.g.
``"sampling.guidance_scale"``) the user has touched, either via raw
config at bind time or via attribute assignment at runtime. A patched
``__setattr__`` on the request dataclass types records assignments into
this set.
The set holds leaf paths only. Nested dataclass or mapping assignments
are flattened to their leaves at record time.
"""
from __future__ import annotations
from collections.abc import Callable, Mapping
import dataclasses
from typing import Any, cast
from fastvideo.api.schema import (
ContinuationState,
GenerationPlan,
GenerationRequest,
InputConfig,
OutputConfig,
PlannedStage,
RequestRuntimeConfig,
RunConfig,
SamplingConfig,
ServeConfig,
)
EXPLICIT_PATHS_ATTR = "_fastvideo_explicit_paths"
_TRACKING_ROOT_ATTR = "_fastvideo_request_tracking_root"
_TRACKING_PATH_ATTR = "_fastvideo_request_tracking_path"
_TRACKING_PATCHED_ATTR = "_fastvideo_request_tracking_patched"
_TRACKED_REQUEST_TYPES = (
GenerationRequest,
InputConfig,
SamplingConfig,
RequestRuntimeConfig,
OutputConfig,
ContinuationState,
PlannedStage,
GenerationPlan,
)
def bind_generation_request_raw(
request: GenerationRequest,
raw: Mapping[str, Any] | None,
) -> GenerationRequest:
"""Install explicit-path tracking on *request*.
*raw* is the parsed config dict (YAML/JSON/kwargs); every leaf key
in it becomes an explicit path. Subsequent attribute assignments on
*request* or its nested dataclasses are recorded automatically via a
patched ``__setattr__``.
"""
_ensure_request_tracking()
# Disable recording while we walk the tree to install roots.
object.__setattr__(request, EXPLICIT_PATHS_ATTR, None)
_set_tracking_roots(request, request, "")
paths: set[str] = set()
_record_value_paths(raw or {}, "", paths)
object.__setattr__(request, EXPLICIT_PATHS_ATTR, paths)
return request
def bind_run_config_raw(
config: RunConfig,
raw: Mapping[str, Any],
) -> RunConfig:
request_raw = raw.get("request")
if isinstance(request_raw, Mapping):
bind_generation_request_raw(config.request, request_raw)
else:
bind_generation_request_raw(config.request, {})
return config
def bind_serve_config_raw(
config: ServeConfig,
raw: Mapping[str, Any],
) -> ServeConfig:
default_request_raw = raw.get("default_request")
if isinstance(default_request_raw, Mapping):
bind_generation_request_raw(config.default_request, default_request_raw)
else:
bind_generation_request_raw(config.default_request, {})
return config
def get_explicit_paths(request: GenerationRequest) -> frozenset[str]:
"""Return a snapshot of the explicit paths set on *request*."""
paths = getattr(request, EXPLICIT_PATHS_ATTR, None)
if isinstance(paths, set | frozenset):
return frozenset(paths)
return frozenset()
def reset_tracking_roots(request: GenerationRequest) -> None:
"""Re-install tracking roots after a deepcopy or manual clone.
The paths set itself deepcopies correctly; we only need to repoint
the tracking root on nested dataclasses at the new root.
"""
_ensure_request_tracking()
_set_tracking_roots(request, request, "")
# ---------------------------------------------------------------------------
# Path recording
# ---------------------------------------------------------------------------
def _record_value_paths(
value: Any,
prefix: str,
out: set[str],
) -> None:
"""Add every leaf path under *value* to *out*.
A leaf is any terminal value (non-dataclass, non-mapping, or empty
mapping/dataclass). ``prefix`` is the dotted path at which *value*
sits. When called with an empty ``prefix`` (the root), leaves are
recorded at their own key.
"""
if dataclasses.is_dataclass(value) and not isinstance(value, type):
dc_fields = dataclasses.fields(value)
if not dc_fields:
if prefix:
out.add(prefix)
return
for field in dc_fields:
child = getattr(value, field.name)
path = f"{prefix}.{field.name}" if prefix else field.name
_record_value_paths(child, path, out)
return
if isinstance(value, Mapping):
if not value:
if prefix:
out.add(prefix)
return
for key, child in value.items():
path = f"{prefix}.{key}" if prefix else key
_record_value_paths(child, path, out)
return
if prefix:
out.add(prefix)
# ---------------------------------------------------------------------------
# __setattr__ patching
# ---------------------------------------------------------------------------
def _ensure_request_tracking() -> None:
for config_type in _TRACKED_REQUEST_TYPES:
_patch_tracking_setattr(config_type)
def _patch_tracking_setattr(config_type: type[Any]) -> None:
if getattr(config_type, _TRACKING_PATCHED_ATTR, False):
return
original_setattr = cast(
Callable[[Any, str, Any], None],
config_type.__setattr__,
)
field_names = {field.name for field in dataclasses.fields(config_type)}
def _tracking_setattr(self: Any, name: str, value: Any) -> None:
if name.startswith("_fastvideo_") or name not in field_names:
original_setattr(self, name, value)
return
original_setattr(self, name, value)
root = getattr(self, _TRACKING_ROOT_ATTR, None)
if root is None:
return
paths = getattr(root, EXPLICIT_PATHS_ATTR, None)
if not isinstance(paths, set):
return
prefix = getattr(self, _TRACKING_PATH_ATTR, "")
path = f"{prefix}.{name}" if prefix else name
# Wholesale dataclass replacement: install roots on the new
# instance so its future mutations are tracked too.
if dataclasses.is_dataclass(value) and not isinstance(value, type):
_set_tracking_roots(root, value, path)
_record_value_paths(value, path, paths)
type.__setattr__(config_type, "__setattr__", _tracking_setattr)
setattr(config_type, _TRACKING_PATCHED_ATTR, True)
# ---------------------------------------------------------------------------
# Tree walk to set tracking root/path on nested dataclasses
# ---------------------------------------------------------------------------
def _set_tracking_roots(
root: GenerationRequest,
obj: Any,
prefix: str,
) -> None:
if not dataclasses.is_dataclass(obj) or isinstance(obj, type):
return
object.__setattr__(obj, _TRACKING_ROOT_ATTR, root)
object.__setattr__(obj, _TRACKING_PATH_ATTR, prefix)
for field in dataclasses.fields(obj):
child = getattr(obj, field.name)
child_path = f"{prefix}.{field.name}" if prefix else field.name
if dataclasses.is_dataclass(child) and not isinstance(child, type):
_set_tracking_roots(root, child, child_path)
__all__ = [
"EXPLICIT_PATHS_ATTR",
"bind_generation_request_raw",
"bind_run_config_raw",
"bind_serve_config_raw",
"get_explicit_paths",
"reset_tracking_roots",
]
@@ -1,10 +1,16 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from typing import Any
from __future__ import annotations
import copy
from dataclasses import dataclass, field, fields
from typing import TYPE_CHECKING, Any
from fastvideo.logger import init_logger
from fastvideo.utils import StoreBoolean
if TYPE_CHECKING:
from fastvideo.api.schema import ContinuationState
logger = init_logger(__name__)
@@ -30,6 +36,16 @@ class SamplingParam:
# Camera control inputs (HYWorld)
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
prompt_attention_mask: list = field(default_factory=list)
negative_attention_mask: list = field(default_factory=list)
# Camera/action control inputs (GameCraft)
camera_states: Any | None = None # Plücker coordinates [B, T_video, 6, H, W]
camera_trajectory: str | None = None
action_list: list[str] | None = None
action_speed_list: list[float] | None = None
gt_latents: Any | None = None # Ground truth latents [B, 16, T, H, W]
conditioning_mask: Any | None = None # Mask [B, 1, T, H, W]
# Camera control inputs (LingBotWorld)
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
@@ -68,6 +84,7 @@ class SamplingParam:
num_inference_steps: int = 50
num_inference_steps_sr: int = 50
guidance_scale: float = 1.0
guidance_scale_2: float | None = None
guidance_rescale: float = 0.0
boundary_ratio: float | None = None
sigmas: list[float] | None = None
@@ -80,6 +97,27 @@ class SamplingParam:
movement_distance: float | None = None
camera_rotation: str | None = None
# LTX-2 multi-modal CFG and STG.
# cfg_scale defaults are 1.0 (CFG off) so ``ForwardBatch.__post_init__``
# doesn't force ``do_classifier_free_guidance`` on non-LTX-2 models that
# never override these fields. LTX-2 presets that need text-CFG on set
# them in their ``defaults`` dict (e.g. ``ltx2_base``).
ltx2_cfg_scale_video: float = 1.0
ltx2_cfg_scale_audio: float = 1.0
ltx2_modality_scale_video: float = 3.0
ltx2_modality_scale_audio: float = 3.0
ltx2_rescale_scale: float = 0.7
ltx2_stg_scale_video: float = 1.0
ltx2_stg_scale_audio: float = 1.0
ltx2_stg_blocks_video: list[int] = field(default_factory=lambda: [29])
ltx2_stg_blocks_audio: list[int] = field(default_factory=lambda: [29])
# Continuation state carried across streaming/multi-segment calls.
continuation_state: ContinuationState | None = None
# When True, the pipeline returns a ContinuationState on the result so
# the caller can resume from the generated segment.
return_continuation_state: bool = False
# Misc
save_video: bool = True
return_frames: bool = True
@@ -94,26 +132,58 @@ class SamplingParam:
raise ValueError("prompt_path must be a txt file")
def update(self, source_dict: dict[str, Any]) -> None:
valid_fields = {f.name for f in fields(self)}
for key, value in source_dict.items():
if hasattr(self, key):
if key in valid_fields:
setattr(self, key, value)
else:
logger.exception("%s has no attribute %s", type(self).__name__, key)
logger.error("%s has no field %s", type(self).__name__, key)
self.__post_init__()
@classmethod
def from_pretrained(cls, model_path: str) -> "SamplingParam":
from fastvideo.registry import get_sampling_param_cls_for_name
sampling_cls = get_sampling_param_cls_for_name(model_path)
if sampling_cls is not None:
sampling_param: SamplingParam = sampling_cls()
else:
logger.warning("Couldn't find an optimal sampling param for %s. Using the default sampling param.",
model_path)
sampling_param = cls()
def from_pretrained(cls, model_path: str) -> SamplingParam:
sampling_param = cls._from_preset(model_path)
if sampling_param is not None:
return sampling_param
return sampling_param
logger.warning(
"Couldn't find a preset for %s."
" Using the default sampling param.",
model_path,
)
return cls()
@classmethod
def _from_preset(
cls,
model_path: str,
) -> SamplingParam | None:
"""Build a SamplingParam from preset defaults.
Returns ``None`` when no preset is configured for
*model_path*, letting the caller fall back to the legacy
subclass lookup.
"""
from fastvideo.registry import get_preset_selection
try:
preset_name, model_family = get_preset_selection(model_path)
except (ValueError, RuntimeError):
return None
if preset_name is None or model_family is None:
return None
from fastvideo.api.presets import get_preset
preset = get_preset(preset_name, model_family)
sp = cls()
valid_fields = {f.name for f in fields(cls)}
for key, value in preset.defaults.items():
if key in valid_fields:
setattr(sp, key, copy.deepcopy(value))
sp.__post_init__()
return sp
@staticmethod
def add_cli_args(parser: Any) -> Any:
+90 -4
View File
@@ -33,8 +33,25 @@ class OffloadConfig:
@dataclass
class CompileConfig:
"""Typed ``torch.compile`` configuration.
``backend``/``fullgraph``/``mode``/``dynamic`` are the four most
common ``torch.compile`` knobs. ``extras`` holds any remaining
``torch.compile`` kwargs (e.g. ``options``, ``disable``).
"""
enabled: bool = False
kwargs: dict[str, Any] = field(default_factory=dict)
text_encoder_enabled: bool | None = None
"""Whether ``torch.compile`` is applied to the text encoder. ``None``
keeps the runtime default. The public ``FastVideoArgs`` adapter does
not yet consume this flag; reserved so the realtime runtime upstream
(PR 7.6) has a typed home for its ``enable_torch_compile_text_encoder``
kwarg without routing through ``pipeline.experimental``."""
backend: str | None = None
fullgraph: bool | None = None
mode: str | None = None
dynamic: bool | None = None
extras: dict[str, Any] = field(default_factory=dict)
@dataclass
@@ -73,10 +90,12 @@ class ComponentConfig:
@dataclass
class PipelineSelection:
workload_type: Literal["t2v", "i2v", "t2i", "i2i"] | None = None
profile: str | None = None
profile_version: str | None = None
preset: str | None = None
preset_version: int | None = None
components: ComponentConfig = field(default_factory=ComponentConfig)
profile_overrides: dict[str, Any] = field(default_factory=dict)
vae_tiling: bool | None = None
"""Tile-based VAE decode. ``None`` keeps the model's default."""
preset_overrides: dict[str, Any] = field(default_factory=dict)
experimental: dict[str, Any] = field(default_factory=dict)
@@ -180,11 +199,73 @@ class RunConfig:
request: GenerationRequest
@dataclass
class WarmupConfig:
enabled: bool = True
prompt: str = ("A cinematic drone shot over coastal cliffs at sunrise, "
"golden light, gentle ocean waves, ultra detailed")
timeout_seconds: int = 2400
@dataclass
class GpuPoolConfig:
num_workers: int | None = None
enable_audio_reencode: bool = True
conditioning_num_frames: int = 9
conditioning_end_offset: int = 0
@dataclass
class PromptEnhancerConfig:
enabled: bool = False
provider: Literal["cerebras", "groq"] = "cerebras"
model: str = "gpt-oss-120b"
timeout_ms: int = 20000
system_prompt_dir: str | None = None
@dataclass
class PromptSafetyConfig:
enabled: bool = False
classifier_path: str | None = None
@dataclass
class StreamingConfig:
session_timeout_seconds: int = 300
generation_segment_cap: int = 6
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
warmup: WarmupConfig = field(default_factory=WarmupConfig)
pool: GpuPoolConfig = field(default_factory=GpuPoolConfig)
prompt: PromptEnhancerConfig = field(default_factory=PromptEnhancerConfig)
safety: PromptSafetyConfig = field(default_factory=PromptSafetyConfig)
@dataclass
class ServeConfig:
"""Typed serve config loaded from ``fastvideo serve --config``.
``default_request`` is a full :class:`GenerationRequest` — the same type
clients POST to ``/v1/videos``. At request time the server merges it into
the incoming body as the operator-pinned baseline.
Important nuance: only fields the operator **explicitly wrote** in the
serve YAML/JSON count as defaults. Although the in-memory object is
fully populated (schema defaults fill every unset field), the merge
walks ``_fastvideo_explicit_paths`` — populated during parse — so
unset fields are *not* forced onto requests. Per-request precedence:
body (client-explicit) > default_request (operator-explicit)
> hardcoded fallback (e.g. ``fps=24``)
See :func:`fastvideo.api.compat.explicit_request_updates` for the
projection and ``entrypoints/openai/video_api.py::_build_generation_kwargs``
for the merge.
"""
generator: GeneratorConfig
server: ServerConfig = field(default_factory=ServerConfig)
default_request: GenerationRequest = field(default_factory=GenerationRequest)
streaming: StreamingConfig | None = None
__all__ = [
@@ -195,16 +276,21 @@ __all__ = [
"GenerationPlan",
"GenerationRequest",
"GeneratorConfig",
"GpuPoolConfig",
"InputConfig",
"OffloadConfig",
"OutputConfig",
"ParallelismConfig",
"PipelineSelection",
"PlannedStage",
"PromptEnhancerConfig",
"PromptSafetyConfig",
"QuantizationConfig",
"RequestRuntimeConfig",
"RunConfig",
"SamplingConfig",
"ServeConfig",
"ServerConfig",
"StreamingConfig",
"WarmupConfig",
]
+1 -1
View File
@@ -5,7 +5,7 @@ from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hunyuangamecraft import HunyuanGameCraftPipelineConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.pipelines.basic.ltx2.pipeline_configs import LTX2T2VConfig
from fastvideo.registry import get_pipeline_config_cls_from_name
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig, WanI2V480PConfig, WanI2V720PConfig,
WanT2V480PConfig, WanT2V720PConfig)
+27 -2
View File
@@ -11,10 +11,34 @@ import torch
from fastvideo.configs.models import DiTConfig, VAEConfig
from fastvideo.configs.models.dits.base import DiTArchConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
from fastvideo.configs.models.encoders.base import TextEncoderArchConfig
from fastvideo.configs.models.encoders.t5 import T5ArchConfig
from fastvideo.configs.models.vaes import WanVAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig
@dataclass
class LongCatT5ArchConfig(T5ArchConfig):
"""T5 arch that pads tokenizer output to ``max_length``.
LongCat's denoising stage concatenates positive and negative
attention masks along the batch dimension for CFG, which requires
uniform seq length. The shared :class:`T5ArchConfig` dropped the
``"padding": "max_length"`` tokenizer kwarg so other DiTs could run
with variable-length masks; LongCat still needs the uniform
contract.
"""
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs["padding"] = "max_length"
@dataclass
class LongCatT5Config(T5Config):
arch_config: TextEncoderArchConfig = field(default_factory=LongCatT5ArchConfig)
@dataclass
class LongCatDiTArchConfig(DiTArchConfig):
"""Extended DiTArchConfig with LongCat-specific fields."""
@@ -103,8 +127,9 @@ class LongCatT2V480PConfig(PipelineConfig):
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(default_factory=lambda: ("bf16", ))
# Text encoding (UMT5 uses T5-like config; postprocess to fixed 512)
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (T5Config(), ))
# UMT5 uses T5-like config; postprocess pads to 512. LongCatT5Config
# restores ``padding="max_length"`` for the CFG concat contract.
text_encoder_configs: tuple[T5Config, ...] = field(default_factory=lambda: (LongCatT5Config(), ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(default_factory=lambda: (longcat_preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda: (umt5_postprocess_text, ))
-13
View File
@@ -1,13 +0,0 @@
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.configs.sample.hunyuangamecraft import (
HunyuanGameCraftSamplingParam,
HunyuanGameCraft65FrameSamplingParam,
HunyuanGameCraft129FrameSamplingParam,
)
__all__ = [
"SamplingParam",
"HunyuanGameCraftSamplingParam",
"HunyuanGameCraft65FrameSamplingParam",
"HunyuanGameCraft129FrameSamplingParam",
]
-18
View File
@@ -1,18 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class Cosmos_Predict2_2B_Video2World_SamplingParam(SamplingParam):
# Video parameters
height: int = 704
width: int = 1280
num_frames: int = 93
fps: int = 16
# Denoising stage
guidance_scale: float = 7.0
negative_prompt: str = "The video captures a series of frames showing ugly scenes, static with no motion, motion blur, over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. Overall, the video is of poor quality."
num_inference_steps: int = 35
-23
View File
@@ -1,23 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class Cosmos25SamplingParamBase(SamplingParam):
height: int = 704
width: int = 1280
num_frames: int = 77
fps: int = 24
seed: int = 0
guidance_scale: float = 7.0
negative_prompt: str = (
"The video captures a series of frames showing ugly scenes, static with no motion, motion blur, "
"over-saturation, shaky footage, low resolution, grainy texture, pixelated images, poorly lit areas, "
"underexposed and overexposed scenes, poor color balance, washed out colors, choppy sequences, jerky movements, "
"low frame rate, artifacting, color banding, unnatural transitions, outdated special effects, fake elements, "
"unconvincing visuals, poorly edited content, jump cuts, visual noise, and flickering. "
"Overall, the video is of poor quality.")
num_inference_steps: int = 35
-24
View File
@@ -1,24 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class Gen3C_Cosmos_7B_SamplingParam(SamplingParam):
"""Defaults for GEN3C (Cosmos-7B) camera-controlled video generation."""
# Video parameters (matching official GEN3C defaults)
height: int = 704
width: int = 1280
num_frames: int = 121
fps: int = 24
# Denoising stage
guidance_scale: float = 1.0
num_inference_steps: int = 35
# GEN3C camera control defaults
trajectory_type: str = "left"
movement_distance: float = 0.3
camera_rotation: str = "center_facing"
-21
View File
@@ -1,21 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class HunyuanSamplingParam(SamplingParam):
num_inference_steps: int = 50
num_frames: int = 125
height: int = 720
width: int = 1280
fps: int = 24
guidance_scale: float = 1.0
@dataclass
class FastHunyuanSamplingParam(HunyuanSamplingParam):
num_inference_steps: int = 6
-55
View File
@@ -1,55 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
import numpy as np
from dataclasses import dataclass, field
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class Hunyuan15_480P_SamplingParam(SamplingParam):
num_inference_steps: int = 50
num_frames: int = 121
height: int = 480
width: int = 848
fps: int = 24
guidance_scale: float = 6.0
sigmas: list[float] | None = field(default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
negative_prompt: str = ""
def __post_init__(self):
super().__post_init__()
self.sigmas = list(np.linspace(1.0, 0.0, self.num_inference_steps + 1)[:-1])
@dataclass
class Hunyuan15_480P_StepDistilled_I2V_SamplingParam(Hunyuan15_480P_SamplingParam):
num_inference_steps: int = 12
height: int = 720
width: int = 1280
guidance_scale: float = 1.0
@dataclass
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
height: int = 720
width: int = 1280
@dataclass
class Hunyuan15_720P_Distilled_I2V_SamplingParam(Hunyuan15_720P_SamplingParam):
guidance_scale: float = 1.0
@dataclass
class Hunyuan15_SR_1080P_SamplingParam(Hunyuan15_480P_SamplingParam):
height_sr: int = 1072
width_sr: int = 1920
num_inference_steps: int = 12
num_inference_steps_sr: int = 8
guidance_scale: float = 1.0
@@ -1,92 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
Sampling parameters for HunyuanGameCraft video generation.
GameCraft generates game-like videos with camera/action control.
Default parameters are based on the official implementation.
"""
from dataclasses import dataclass
from typing import Any
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class HunyuanGameCraftSamplingParam(SamplingParam):
"""Sampling parameters for HunyuanGameCraft video generation.
Supports camera/action conditioning via:
- camera_trajectory: Plücker coordinates for camera motion
- action_list: List of actions (e.g., ["forward", "left", "right"])
- action_speed_list: Speed multipliers for each action
Default resolution is 704x1280 (same as HunyuanVideo).
Default frame count is 33 video frames -> 9 latent frames.
"""
# Number of denoising steps
num_inference_steps: int = 50
# Video dimensions
# 33 video frames -> 9 latent frames (4x temporal compression)
num_frames: int = 33
height: int = 704
width: int = 1280
fps: int = 24
# Guidance scale - official GameCraft uses CFG with guidance_scale=6.0
guidance_scale: float = 6.0
# Negative prompt for CFG (empty string = unconditional)
negative_prompt: str = ""
# Camera/Action conditioning
# Camera states as Plücker coordinates [B, T_video, 6, H, W]
camera_states: Any | None = None
# Camera trajectory file/identifier (alternative to camera_states)
camera_trajectory: str | None = None
# Action list for camera motion (e.g., ["forward", "left"])
action_list: list[str] | None = None
# Speed multipliers for each action
action_speed_list: list[float] | None = None
# History frame conditioning (for autoregressive generation)
# Ground truth latents for conditioning [B, 16, T, H, W]
gt_latents: Any | None = None
# Mask for conditioning (1=use gt, 0=generate) [B, 1, T, H, W]
conditioning_mask: Any | None = None
# Number of conditioning frames (for autoregressive) - maps to num_cond_frames
num_cond_frames: int = 0
def __post_init__(self) -> None:
super().__post_init__()
# Validate action lists
if (self.action_list is not None and self.action_speed_list is not None
and len(self.action_list) != len(self.action_speed_list)):
raise ValueError(f"action_list length ({len(self.action_list)}) must match "
f"action_speed_list length ({len(self.action_speed_list)})")
@dataclass
class HunyuanGameCraft65FrameSamplingParam(HunyuanGameCraftSamplingParam):
"""Sampling parameters for 65-frame GameCraft generation.
65 video frames -> 17 latent frames (with first frame as key frame).
This is useful for longer video generation.
"""
num_frames: int = 65
@dataclass
class HunyuanGameCraft129FrameSamplingParam(HunyuanGameCraftSamplingParam):
"""Sampling parameters for 129-frame GameCraft generation.
129 video frames -> 33 latent frames.
This is the maximum supported by the official implementation.
"""
num_frames: int = 129
-25
View File
@@ -1,25 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.sample.base import SamplingParam
import numpy as np
@dataclass
class HYWorld_SamplingParam(SamplingParam):
num_inference_steps: int = 50
num_frames: int = 125
height: int = 480
width: int = 832
fps: int = 24
# Camera trajectory: pose string (e.g., 'w-31' means generating [1 + 31] latents) or JSON file path
pose: str = 'w-31'
guidance_scale: float = 6.0
prompt_attention_mask: list = field(default_factory=list)
negative_attention_mask: list = field(default_factory=list)
sigmas: list[float] | None = field(default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
negative_prompt: str = ""
-20
View File
@@ -1,20 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.wan import Wan2_2_I2V_A14B_SamplingParam
@dataclass
class LingBotWorld_SamplingParam(Wan2_2_I2V_A14B_SamplingParam):
guidance_scale: float = 5.0 # high_noise
guidance_scale_2: float = 5.0 # low_noise
num_inference_steps: int = 70
boundary_ratio: float | None = 0.947
negative_prompt: str | None = ("画面突变,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,"
"最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,"
"畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走,"
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
"皮肤,肢体,面部特征,汽车,电线")
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
-69
View File
@@ -1,69 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class LTX2BaseSamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 base one-stage T2V.
Values follow the official LTX-2 one-stage defaults.
Multi-modal CFG params are read by ``LTX2DenoisingStage``.
"""
seed: int = 10
num_frames: int = 121
height: int = 512
width: int = 768
fps: int = 24
num_inference_steps: int = 40
guidance_scale: float = 3.0
# Copied/following official LTX-2 DEFAULT_NEGATIVE_PROMPT.
negative_prompt: str = ("blurry, out of focus, overexposed, underexposed, low contrast, "
"washed out colors, excessive noise, grainy texture, poor lighting, "
"flickering, motion blur, distorted proportions, unnatural skin "
"tones, deformed facial features, asymmetrical face, missing facial "
"features, extra limbs, disfigured hands, wrong hand count, "
"artifacts around text, inconsistent perspective, camera shake, "
"incorrect depth of field, background too sharp, background clutter, "
"distracting reflections, harsh shadows, inconsistent lighting "
"direction, color banding, cartoonish rendering, 3D CGI look, "
"unrealistic materials, uncanny valley effect, incorrect ethnicity, "
"wrong gender, exaggerated expressions, wrong gaze direction, "
"mismatched lip sync, silent or muted audio, distorted voice, "
"robotic voice, echo, background noise, off-sync audio, incorrect "
"dialogue, added dialogue, repetitive speech, jittery movement, "
"awkward pauses, incorrect timing, unnatural transitions, "
"inconsistent framing, tilted camera, flat lighting, inconsistent "
"tone, cinematic oversaturation, stylized filters, or AI artifacts.")
# Official LTX-2 multi-modal CFG defaults.
ltx2_cfg_scale_video: float = 3.0
ltx2_cfg_scale_audio: float = 7.0
ltx2_modality_scale_video: float = 3.0
ltx2_modality_scale_audio: float = 3.0
ltx2_rescale_scale: float = 0.7
# STG (Spatio-Temporal Guidance) defaults from official LTX-2.
ltx2_stg_scale_video: float = 1.0
ltx2_stg_scale_audio: float = 1.0
ltx2_stg_blocks_video: list[int] = field(default_factory=lambda: [29])
ltx2_stg_blocks_audio: list[int] = field(default_factory=lambda: [29])
@dataclass
class LTX2DistilledSamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled one-stage T2V."""
seed: int = 10
num_frames: int = 121
height: int = 1024
width: int = 1536
fps: int = 24
num_inference_steps: int = 8
guidance_scale: float = 1.0
# No default negative_prompt for distilled models
negative_prompt: str = ""
# Backward compatibility alias.
LTX2SamplingParam = LTX2DistilledSamplingParam
-25
View File
@@ -1,25 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class SD35SamplingParam(SamplingParam):
prompt: str | None = "a photo of a cat"
negative_prompt: str = ""
num_videos_per_prompt: int = 1
seed: int = 0
num_frames: int = 1
height: int = 512
width: int = 512
fps: int = 1
num_inference_steps: int = 28
guidance_scale: float = 6.0
@@ -1,73 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
"""
TurboDiffusion sampling parameters.
TurboDiffusion uses RCM (recurrent Consistency Model) scheduler for
1-4 step video generation with no classifier-free guidance.
"""
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class TurboDiffusionT2V_1_3B_SamplingParam(SamplingParam):
"""Sampling parameters for TurboDiffusion T2V 1.3B model.
Uses 4-step RCM sampling with guidance_scale=1.0 (no CFG).
"""
# Video parameters
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
guidance_scale: float = 1.0
num_inference_steps: int = 4
# No negative prompt needed for TurboDiffusion (no CFG)
negative_prompt: str | None = None
@dataclass
class TurboDiffusionT2V_14B_SamplingParam(SamplingParam):
"""Sampling parameters for TurboDiffusion T2V 14B model.
Uses 4-step RCM sampling with guidance_scale=1.0 (no CFG).
"""
# Video parameters (720p for 14B)
height: int = 720
width: int = 1280
num_frames: int = 81
fps: int = 16
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
guidance_scale: float = 1.0
num_inference_steps: int = 4
# No negative prompt needed for TurboDiffusion (no CFG)
negative_prompt: str | None = None
@dataclass
class TurboDiffusionI2V_A14B_SamplingParam(SamplingParam):
"""Sampling parameters for TurboDiffusion I2V A14B model.
Uses 4-step RCM sampling with dual-model switching (high/low noise).
"""
# Video parameters (720p for A14B I2V)
height: int = 720
width: int = 1280
num_frames: int = 81
fps: int = 16
# Denoising stage - TurboDiffusion uses 1-4 steps with no CFG
guidance_scale: float = 1.0
num_inference_steps: int = 4
# Note: boundary_ratio is set in the pipeline config (TurboDiffusionI2VConfig),
# not here. This keeps sampling params and pipeline config separate.
# No negative prompt needed for TurboDiffusion (no CFG)
negative_prompt: str | None = None
-154
View File
@@ -1,154 +0,0 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class WanT2V_1_3B_SamplingParam(SamplingParam):
# Video parameters
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
# Denoising stage
guidance_scale: float = 3.0
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
num_inference_steps: int = 50
@dataclass
class WanT2V_14B_SamplingParam(SamplingParam):
# Video parameters
height: int = 720
width: int = 1280
num_frames: int = 81
fps: int = 16
# Denoising stage
guidance_scale: float = 5.0
negative_prompt: str = "Bright tones, overexposed, static, blurred details, subtitles, style, works, paintings, images, static, overall gray, worst quality, low quality, JPEG compression residue, ugly, incomplete, extra fingers, poorly drawn hands, poorly drawn faces, deformed, disfigured, misshapen limbs, fused fingers, still picture, messy background, three legs, many people in the background, walking backwards"
num_inference_steps: int = 50
@dataclass
class WanI2V_14B_480P_SamplingParam(WanT2V_1_3B_SamplingParam):
# Denoising stage
guidance_scale: float = 5.0
num_inference_steps: int = 40
@dataclass
class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
# Denoising stage
guidance_scale: float = 5.0
num_inference_steps: int = 40
@dataclass
class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
# DMD parameters
# dmd_denoising_steps: list[int] | None = field(default_factory=lambda: [1000, 757, 522])
num_inference_steps: int = 3
num_frames: int = 61
height: int = 448
width: int = 832
fps: int = 16
# =============================================
# ============= Wan2.1 Fun Models =============
# =============================================
@dataclass
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale: float = 6.0
num_inference_steps: int = 50
@dataclass
class Wan2_1_Fun_1_3B_Control_SamplingParam(SamplingParam):
fps: int = 16
num_frames: int = 49
height: int = 832
width: int = 480
guidance_scale: float = 6.0
# =============================================
# ============= Wan2.2 TI2V Models =============
# =============================================
@dataclass
class Wan2_2_Base_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.2 TI2V 5B model."""
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
@dataclass
class Wan2_2_TI2V_5B_SamplingParam(Wan2_2_Base_SamplingParam):
"""Sampling parameters for Wan2.2 TI2V 5B model."""
height: int = 704
width: int = 1280
num_frames: int = 121
fps: int = 24
guidance_scale: float = 5.0
num_inference_steps: int = 50
@dataclass
class Wan2_2_T2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale: float = 4.0 # high_noise
guidance_scale_2: float = 3.0 # low_noise
num_inference_steps: int = 40
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
@dataclass
class Wan2_2_I2V_A14B_SamplingParam(Wan2_2_Base_SamplingParam):
guidance_scale: float = 3.5 # high_noise
guidance_scale_2: float = 3.5 # low_noise
num_inference_steps: int = 40
fps: int = 16
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
# can be overridden during sampling
@dataclass
class Wan2_2_Fun_A14B_Control_SamplingParam(Wan2_1_Fun_1_3B_Control_SamplingParam):
num_frames: int = 81
# =============================================
# ============= Causal Self-Forcing =============
# =============================================
@dataclass
class SelfForcingWan2_1_T2V_1_3B_480P_SamplingParam(Wan2_1_Fun_1_3B_InP_SamplingParam):
pass
@dataclass
class SelfForcingWan2_2_T2V_A14B_480P_SamplingParam(Wan2_2_T2V_A14B_SamplingParam):
num_inference_steps: int = 8
num_frames: int = 81
height: int = 448
width: int = 832
fps: int = 16
@dataclass
class MatrixGame2_SamplingParam(SamplingParam):
height: int = 352
width: int = 640
num_frames: int = 57
fps: int = 25
guidance_scale: float = 1.0
num_inference_steps: int = 3
negative_prompt: str | None = None
+1 -1
View File
@@ -7,7 +7,7 @@ Example usage:
# launch a server and benchmark on it
# T2V or T2I or any other multimodal generation model
fastvideo serve --model-path Wan-AI/Wan2.1-T2V-1.3B-Diffusers --port 8000
fastvideo serve --config serve.yaml
# benchmark it and make sure the port is the same as the server's port
fastvideo bench --dataset vbench --num-prompts 20 --port 8000
+216
View File
@@ -0,0 +1,216 @@
"""``fastvideo eval`` CLI: list registered eval metrics and run them
against a set of videos.
This is a thin wrapper around :mod:`fastvideo.eval`. Heavy lifting
(metric loading, GPU handling, batching) lives in
:func:`fastvideo.eval.create_evaluator`.
Examples::
fastvideo eval list
fastvideo eval list --group vbench
fastvideo eval run --videos path/to/videos/*.mp4 \\
--metrics common.ssim --reference path/to/refs/
fastvideo eval run --videos clip.mp4 --metrics vbench.aesthetic_quality \\
--output scores.json
"""
from __future__ import annotations
import argparse
import glob
import json
from pathlib import Path
from typing import cast
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.logger import init_logger
from fastvideo.utils import FlexibleArgumentParser
logger = init_logger(__name__)
class EvalSubcommand(CLISubcommand):
"""The ``eval`` subcommand — entry point for the eval suite."""
def __init__(self) -> None:
self.name = "eval"
super().__init__()
def cmd(self, args: argparse.Namespace) -> None:
action = getattr(args, "eval_action", None)
if action == "list":
_cmd_list(args)
elif action == "run":
_cmd_run(args)
else:
# Re-print help if no action was given.
self._parser.print_help() # type: ignore[attr-defined]
def validate(self, args: argparse.Namespace) -> None:
action = getattr(args, "eval_action", None)
if action == "run" and not args.videos:
raise SystemExit("`fastvideo eval run` requires --videos")
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
eval_parser = subparsers.add_parser(
"eval",
help="Run video-gen evaluation metrics",
usage="fastvideo eval {list,run} [...]",
)
sub = eval_parser.add_subparsers(dest="eval_action", required=False)
# `eval list`
list_p = sub.add_parser("list", help="List registered metrics")
list_p.add_argument("--group", type=str, default=None, help="Filter to a metric group (e.g. 'vbench').")
# `eval run`
run_p = sub.add_parser("run", help="Evaluate videos against one or more metrics")
run_p.add_argument("--videos",
type=str,
nargs="+",
required=False,
help="Path, glob, or directory of generated videos.")
run_p.add_argument("--reference",
type=str,
default=None,
help="Path / glob / dir of reference videos (for paired metrics).")
run_p.add_argument("--metrics", type=str, default="all", help="Comma-separated metric names, or 'all'.")
run_p.add_argument("--device", type=str, default="cuda", help="Torch device (e.g. 'cuda', 'cuda:0', 'cpu').")
run_p.add_argument("--text-prompt",
type=str,
nargs="*",
default=None,
help="Prompt(s) for text-conditioned metrics. One per video.")
run_p.add_argument("--fps", type=float, default=None, help="Frame-rate annotation passed to fps-aware metrics.")
run_p.add_argument("--output",
type=str,
default=None,
help="Write results as JSON to this path (default: stdout).")
# Stash the parser so cmd() can re-print help on no-action.
self._parser = eval_parser # type: ignore[attr-defined]
return cast(FlexibleArgumentParser, eval_parser)
def _cmd_list(args: argparse.Namespace) -> None:
from fastvideo.eval import list_metrics
names = list_metrics()
if args.group:
prefix = args.group.rstrip(".") + "."
names = [n for n in names if n == args.group or n.startswith(prefix)]
if not names:
print(f"(no metrics matched group {args.group!r})")
return
for name in names:
print(name)
def _cmd_run(args: argparse.Namespace) -> None:
from fastvideo.eval import create_evaluator
from fastvideo.eval.io import load_video
video_paths = _expand_paths(args.videos)
if not video_paths:
raise SystemExit(f"No videos matched: {args.videos}")
ref_paths = _expand_paths([args.reference]) if args.reference else None
metrics_arg: list[str] | str = ("all" if args.metrics == "all" else
[m.strip() for m in args.metrics.split(",") if m.strip()])
evaluator = create_evaluator(metrics=metrics_arg, device=args.device)
all_results: list[dict] = []
for i, vp in enumerate(video_paths):
logger.info("Evaluating %s (%d/%d)", vp, i + 1, len(video_paths))
kwargs: dict = {"video": load_video(vp)}
if ref_paths is not None:
ref = ref_paths[i] if i < len(ref_paths) else ref_paths[0]
kwargs["reference"] = load_video(ref)
if args.text_prompt is not None:
prompt = (args.text_prompt[i] if i < len(args.text_prompt) else args.text_prompt[0])
kwargs["text_prompt"] = [prompt]
if args.fps is not None:
kwargs["fps"] = args.fps
results = evaluator.evaluate(**kwargs)
all_results.append({
"video": str(vp),
"scores": _serialize_results(results),
})
payload = json.dumps(all_results, indent=2, default=_jsonable)
if args.output:
Path(args.output).write_text(payload)
logger.info("Wrote results to %s", args.output)
else:
print(payload)
def _expand_paths(patterns: list[str]) -> list[str]:
out: list[str] = []
for pat in patterns:
p = Path(pat)
if p.is_dir():
for ext in (".mp4", ".avi", ".mov", ".mkv", ".gif"):
out.extend(sorted(str(f) for f in p.iterdir() if f.suffix.lower() == ext))
elif any(c in pat for c in "*?["):
out.extend(sorted(glob.glob(pat)))
else:
out.append(pat)
# de-dup, preserve order
seen: set[str] = set()
deduped: list[str] = []
for x in out:
if x not in seen:
seen.add(x)
deduped.append(x)
return deduped
def _serialize_results(results) -> dict | list:
"""Turn evaluator output (dict or list-of-dicts) into JSON-friendly form."""
if isinstance(results, list):
return [_serialize_results(r) for r in results]
if isinstance(results, dict):
return {k: _serialize_metric_result(v) for k, v in results.items()}
return _serialize_metric_result(results)
def _serialize_metric_result(mr) -> dict:
return {
"name": getattr(mr, "name", None),
"score": getattr(mr, "score", None),
"details": getattr(mr, "details", None),
}
def _jsonable(obj):
"""``json.dumps(default=...)`` coercer for metric outputs.
Metrics frequently land numpy scalars / arrays, torch tensors, and
pathlib paths inside ``MetricResult.details`` (e.g. ``optical_flow``
populates ``per_frame_metrics`` with numpy floats). The stdlib JSON
encoder rejects all of those by default — this callback walks the
leaves and coerces them to native Python types.
"""
import numpy as np
import torch
if isinstance(obj, np.integer):
return int(obj)
if isinstance(obj, np.floating):
return float(obj)
if isinstance(obj, np.bool_):
return bool(obj)
if isinstance(obj, np.ndarray):
return obj.tolist()
if isinstance(obj, torch.Tensor):
return obj.detach().cpu().tolist()
if isinstance(obj, Path):
return str(obj)
raise TypeError(f"Object of type {type(obj).__name__} is not JSON serializable")
def cmd_init() -> list[CLISubcommand]:
return [EvalSubcommand()]
+25 -69
View File
@@ -2,19 +2,17 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
import argparse
import dataclasses
import os
from typing import cast
from fastvideo import VideoGenerator
from fastvideo.configs.sample.base import SamplingParam
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.entrypoints.cli.utils import RaiseNotImplementedAction
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.entrypoints.cli.inference_config import build_generate_run_config
from fastvideo.logger import init_logger
from fastvideo.utils import FlexibleArgumentParser
logger = init_logger(__name__)
_VALIDATED_RUN_CONFIG_ATTR = "_fastvideo_validated_run_config"
class GenerateSubcommand(CLISubcommand):
@@ -23,89 +21,47 @@ class GenerateSubcommand(CLISubcommand):
def __init__(self) -> None:
self.name = "generate"
super().__init__()
self.init_arg_names = self._get_init_arg_names()
self.generation_arg_names = self._get_generation_arg_names()
def _get_init_arg_names(self) -> list[str]:
"""Get names of arguments for VideoGenerator initialization"""
return ["num_gpus", "tp_size", "sp_size", "model_path"]
def _get_generation_arg_names(self) -> list[str]:
"""Get names of arguments for generate_video method"""
return [field.name for field in dataclasses.fields(SamplingParam)]
def cmd(self, args: argparse.Namespace) -> None:
excluded_args = ['subparser', 'config', 'dispatch_function']
run_config = getattr(args, _VALIDATED_RUN_CONFIG_ATTR, None)
if run_config is None:
run_config = build_generate_run_config(
args,
overrides=getattr(args, "_unknown", None),
)
logger.info("CLI generate config: %s", run_config)
provided_args = {}
for k, v in vars(args).items():
if (k not in excluded_args and v is not None and hasattr(args, '_provided') and k in args._provided):
provided_args[k] = v
if 'model_path' in vars(args) and args.model_path is not None:
provided_args['model_path'] = args.model_path
if 'prompt' in vars(args) and args.prompt is not None:
provided_args['prompt'] = args.prompt
merged_args = {**provided_args}
logger.info('CLI Args: %s', merged_args)
if 'model_path' not in merged_args or not merged_args['model_path']:
raise ValueError("model_path must be provided either in config file or via --model-path")
# Check if either prompt or prompt_txt is provided
has_prompt = 'prompt' in merged_args and merged_args['prompt']
has_prompt_txt = 'prompt_txt' in merged_args and merged_args['prompt_txt']
if not (has_prompt or has_prompt_txt):
raise ValueError("Either prompt or prompt_txt must be provided")
if has_prompt and has_prompt_txt:
raise ValueError("Cannot provide both 'prompt' and 'prompt_txt'. Use only one of them.")
init_args = {k: v for k, v in merged_args.items() if k not in self.generation_arg_names}
generation_args = {k: v for k, v in merged_args.items() if k in self.generation_arg_names}
generation_args.setdefault("return_frames", False)
model_path = init_args.pop('model_path')
prompt = generation_args.pop('prompt', None)
generator = VideoGenerator.from_pretrained(model_path=model_path, **init_args)
# Call generate_video - it handles both single and batch modes
generator.generate_video(prompt=prompt, **generation_args)
generator = VideoGenerator.from_config(run_config.generator)
generator.generate(run_config.request)
def validate(self, args: argparse.Namespace) -> None:
"""Validate the arguments for this command"""
if args.num_gpus is not None and args.num_gpus <= 0:
raise ValueError("Number of gpus must be positive")
if args.config and not os.path.exists(args.config):
if not args.config:
raise ValueError("fastvideo generate requires --config PATH; use a nested "
"run config plus optional dotted overrides")
if not os.path.exists(args.config):
raise ValueError(f"Config file not found: {args.config}")
setattr(
args,
_VALIDATED_RUN_CONFIG_ATTR,
build_generate_run_config(
args,
overrides=getattr(args, "_unknown", None),
),
)
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
generate_parser = subparsers.add_parser(
"generate",
help="Run inference on a model",
usage="fastvideo generate (--model-path MODEL_PATH_OR_ID --prompt PROMPT) | --config CONFIG_FILE [OPTIONS]")
usage="fastvideo generate --config RUN_CONFIG [--dotted.override VALUE]")
generate_parser.add_argument(
"--config",
type=str,
default='',
required=False,
help="Read CLI options from a config JSON or YAML file. If provided, --model-path and --prompt are optional."
)
generate_parser = FastVideoArgs.add_cli_args(generate_parser)
generate_parser = SamplingParam.add_cli_args(generate_parser)
generate_parser.add_argument(
"--text-encoder-configs",
action=RaiseNotImplementedAction,
help="JSON array of text encoder configurations (NOT YET IMPLEMENTED)",
help="Path to a nested run config JSON or YAML file. Required.",
)
return cast(FlexibleArgumentParser, generate_parser)
@@ -0,0 +1,111 @@
# SPDX-License-Identifier: Apache-2.0
from __future__ import annotations
import argparse
from collections.abc import Mapping
from copy import deepcopy
from typing import Any
from fastvideo.api.overrides import apply_overrides, parse_cli_overrides
from fastvideo.api.parser import load_raw_config, parse_config
from fastvideo.api.schema import RunConfig, ServeConfig
_GENERATE_OVERRIDE_PREFIXES = ("generator.", "request.")
_SERVE_OVERRIDE_PREFIXES = (
"generator.",
"server.",
"default_request.",
)
def build_generate_run_config(
args: argparse.Namespace,
overrides: list[str] | None = None,
) -> RunConfig:
raw = _load_nested_config(getattr(args, "config", None))
raw.setdefault("request", {})
raw = _apply_dotted_overrides(
raw,
overrides,
allowed_prefixes=_GENERATE_OVERRIDE_PREFIXES,
)
_ensure_generate_cli_defaults(raw)
config = parse_config(RunConfig, raw)
_validate_num_gpus(config.generator.engine.num_gpus)
_validate_generate_prompt_sources(config)
return config
def build_serve_config(
args: argparse.Namespace,
overrides: list[str] | None = None,
) -> ServeConfig:
raw = _load_nested_config(getattr(args, "config", None))
raw.setdefault("server", {})
raw.setdefault("default_request", {})
raw = _apply_dotted_overrides(
raw,
overrides,
allowed_prefixes=_SERVE_OVERRIDE_PREFIXES,
)
config = parse_config(ServeConfig, raw)
_validate_num_gpus(config.generator.engine.num_gpus)
return config
def _load_nested_config(path: str | None) -> dict[str, Any]:
if not path:
raise ValueError("Inference CLI requires --config PATH; use a nested config file "
"plus optional dotted overrides")
raw = load_raw_config(path)
if not isinstance(raw.get("generator"), Mapping):
raise ValueError("Inference config must use the nested schema with a top-level "
"'generator' mapping")
return deepcopy(dict(raw))
def _apply_dotted_overrides(
raw: Mapping[str, Any],
overrides: list[str] | None,
*,
allowed_prefixes: tuple[str, ...],
) -> dict[str, Any]:
if not overrides:
return deepcopy(dict(raw))
parsed = parse_cli_overrides(overrides)
for key in parsed:
if "." not in key:
raise ValueError("CLI overrides must use dotted config paths like "
"--request.sampling.seed 42")
if not key.startswith(allowed_prefixes):
allowed = ", ".join(allowed_prefixes)
raise ValueError(f"Unsupported override path {key!r}. Allowed prefixes: {allowed}")
return apply_overrides(raw, parsed)
def _ensure_generate_cli_defaults(raw: dict[str, Any]) -> None:
request = raw.setdefault("request", {})
output = request.setdefault("output", {})
output.setdefault("return_frames", False)
def _validate_generate_prompt_sources(config: RunConfig) -> None:
has_prompt = config.request.prompt is not None
has_prompt_path = config.request.inputs.prompt_path is not None
if not (has_prompt or has_prompt_path):
raise ValueError("Either request.prompt or request.inputs.prompt_path must be provided")
if has_prompt and has_prompt_path:
raise ValueError("Cannot provide both request.prompt and request.inputs.prompt_path")
def _validate_num_gpus(num_gpus: int) -> None:
if num_gpus <= 0:
raise ValueError(f"generator.engine.num_gpus must be > 0; got {num_gpus}")
__all__ = [
"build_generate_run_config",
"build_serve_config",
]
+10 -6
View File
@@ -1,11 +1,11 @@
# SPDX-License-Identifier: Apache-2.0
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/main.py
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.entrypoints.cli.generate import cmd_init as generate_cmd_init
from fastvideo.utils import FlexibleArgumentParser
from fastvideo.entrypoints.cli.serve import cmd_init as serve_cmd_init
from fastvideo.entrypoints.cli.bench import cmd_init as bench_cmd_init
from fastvideo.entrypoints.cli.eval import cmd_init as eval_cmd_init
def cmd_init() -> list[CLISubcommand]:
@@ -14,6 +14,7 @@ def cmd_init() -> list[CLISubcommand]:
commands.extend(generate_cmd_init())
commands.extend(serve_cmd_init())
commands.extend(bench_cmd_init())
commands.extend(eval_cmd_init())
return commands
@@ -27,14 +28,17 @@ def main() -> None:
for cmd in cmd_init():
cmd.subparser_init(subparsers).set_defaults(dispatch_function=cmd.cmd)
cmds[cmd.name] = cmd
args = parser.parse_args()
args, unknown = parser.parse_known_args()
if unknown and args.subparser not in {"generate", "serve"}:
parser.error(f"unrecognized arguments: {' '.join(unknown)}")
args._unknown = unknown
if args.subparser in cmds:
cmds[args.subparser].validate(args)
if hasattr(args, "dispatch_function"):
args.dispatch_function(args)
else:
parser.print_help()
return
parser.print_help()
if __name__ == "__main__":
+47 -69
View File
@@ -2,14 +2,17 @@
# adapted from vllm: https://github.com/vllm-project/vllm/blob/v0.7.3/vllm/entrypoints/cli/serve.py
import argparse
import os
from typing import cast
from fastvideo.api.compat import generator_config_to_fastvideo_args
from fastvideo.entrypoints.cli.cli_types import CLISubcommand
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.entrypoints.cli.inference_config import build_serve_config
from fastvideo.logger import init_logger
from fastvideo.utils import FlexibleArgumentParser
logger = init_logger(__name__)
_VALIDATED_SERVE_CONFIG_ATTR = "_fastvideo_validated_serve_config"
class ServeSubcommand(CLISubcommand):
@@ -20,94 +23,69 @@ class ServeSubcommand(CLISubcommand):
super().__init__()
def cmd(self, args: argparse.Namespace) -> None:
excluded_args = {
"subparser",
"config",
"dispatch_function",
"host",
"port",
"output_dir",
}
serve_config = getattr(args, _VALIDATED_SERVE_CONFIG_ATTR, None)
if serve_config is None:
serve_config = build_serve_config(
args,
overrides=getattr(args, "_unknown", None),
)
provided: set[str] = getattr(args, '_provided', set())
cli_kwargs = {}
for k, v in vars(args).items():
if k in excluded_args:
continue
if k == '_provided':
continue
if k in provided and v is not None:
cli_kwargs[k] = v
logger.info("CLI serve config: %s", serve_config)
if 'model_path' not in cli_kwargs and args.model_path is not None:
cli_kwargs['model_path'] = args.model_path
if not cli_kwargs.get('model_path'):
raise ValueError("model_path must be provided via --model-path")
# A `streaming:` block selects the WebSocket/Dynamo runtime;
# its deps stay out of REST-only deployments via lazy import.
if serve_config.streaming is not None:
from fastvideo.entrypoints.streaming.server import (
run_server as run_streaming_server, )
run_streaming_server(serve_config)
return
from fastvideo.entrypoints.openai.api_server import (
DEFAULT_HOST,
DEFAULT_OUTPUT_DIR,
DEFAULT_PORT,
run_server,
run_server, )
logger.info(
"Server will listen on %s:%d",
serve_config.server.host,
serve_config.server.port,
)
host = getattr(args, "host", DEFAULT_HOST)
port = getattr(args, "port", DEFAULT_PORT)
output_dir = getattr(args, "output_dir", DEFAULT_OUTPUT_DIR)
logger.info("CLI serve args: %s", cli_kwargs)
logger.info("Server will listen on %s:%d", host, port)
fastvideo_args = FastVideoArgs.from_kwargs(**cli_kwargs)
run_server(fastvideo_args, host=host, port=port, output_dir=output_dir)
fastvideo_args = generator_config_to_fastvideo_args(serve_config.generator)
run_server(
fastvideo_args,
host=serve_config.server.host,
port=serve_config.server.port,
output_dir=serve_config.server.output_dir,
default_request=serve_config.default_request,
)
def validate(self, args: argparse.Namespace) -> None:
if args.num_gpus is not None and args.num_gpus <= 0:
raise ValueError("Number of gpus must be positive")
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
from fastvideo.entrypoints.openai.api_server import (
DEFAULT_HOST,
DEFAULT_OUTPUT_DIR,
DEFAULT_PORT,
if not args.config:
raise ValueError("fastvideo serve requires --config PATH; use a nested "
"serve config plus optional dotted overrides")
if not os.path.exists(args.config):
raise ValueError(f"Config file not found: {args.config}")
setattr(
args,
_VALIDATED_SERVE_CONFIG_ATTR,
build_serve_config(
args,
overrides=getattr(args, "_unknown", None),
),
)
def subparser_init(self, subparsers: argparse._SubParsersAction) -> FlexibleArgumentParser:
serve_parser = subparsers.add_parser(
"serve",
help="Start an OpenAI-compatible HTTP server",
usage=("fastvideo serve --model-path MODEL_PATH_OR_ID "
"[--host HOST] [--port PORT] [OPTIONS]"),
)
serve_parser.add_argument(
"--host",
type=str,
default=DEFAULT_HOST,
help=f"Host to bind the server to (default: {DEFAULT_HOST})",
)
serve_parser.add_argument(
"--port",
type=int,
default=DEFAULT_PORT,
help=f"Port to listen on (default: {DEFAULT_PORT})",
)
serve_parser.add_argument(
"--output-dir",
type=str,
default=DEFAULT_OUTPUT_DIR,
help=("Directory for generated outputs "
f"(default: {DEFAULT_OUTPUT_DIR})"),
usage="fastvideo serve --config SERVE_CONFIG [--dotted.override VALUE]",
)
serve_parser.add_argument(
"--config",
type=str,
default="",
required=False,
help="Read CLI options from a config JSON or YAML file.",
help="Path to a nested config JSON or YAML file. Required.",
)
serve_parser = FastVideoArgs.add_cli_args(serve_parser)
return cast(FlexibleArgumentParser, serve_parser)
+38 -2
View File
@@ -8,6 +8,8 @@ import uvicorn
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from fastvideo.api.presets import validate_preset_selection
from fastvideo.api.schema import GenerationRequest
from fastvideo.entrypoints.openai.state import (
DEFAULT_OUTPUT_DIR,
clear_state,
@@ -16,6 +18,7 @@ from fastvideo.entrypoints.openai.state import (
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.registry import get_preset_selection
logger = init_logger(__name__)
@@ -23,17 +26,40 @@ DEFAULT_HOST = "0.0.0.0"
DEFAULT_PORT = 8000
def _validate_default_request_against_preset(
default_request: GenerationRequest,
model_path: str,
) -> None:
"""Validate ``default_request.stage_overrides`` against the model's preset.
Called once at server startup from :func:`run_server`. The
``default_request`` is static server config, so validation results are
invariant across requests — there's no reason to re-run per request.
"""
if not default_request.stage_overrides:
return
preset_name, model_family = get_preset_selection(model_path)
if preset_name is None or model_family is None:
return
validate_preset_selection(
preset_name,
model_family,
stage_overrides=default_request.stage_overrides,
)
@asynccontextmanager
async def lifespan(app: FastAPI) -> AsyncIterator[None]:
"""Load model on startup, clean up on shutdown"""
args: FastVideoArgs = app.state.fastvideo_args
output_dir: str = app.state.output_dir
default_request: GenerationRequest | None = getattr(app.state, "default_request", None)
logger.info("Loading model from %s ...", args.model_path)
generator = VideoGenerator.from_fastvideo_args(args)
logger.info("Model loaded successfully.")
set_state(generator, args, output_dir)
set_state(generator, args, output_dir, default_request=default_request)
yield # server is running
@@ -46,6 +72,7 @@ async def lifespan(app: FastAPI) -> AsyncIterator[None]:
def create_app(
fastvideo_args: FastVideoArgs,
output_dir: str = DEFAULT_OUTPUT_DIR,
default_request: GenerationRequest | None = None,
) -> FastAPI:
"""Build the FastAPI application with all routers mounted"""
@@ -56,6 +83,7 @@ def create_app(
)
app.state.fastvideo_args = fastvideo_args
app.state.output_dir = output_dir
app.state.default_request = default_request
app.add_middleware(
CORSMiddleware,
@@ -108,9 +136,17 @@ def run_server(
host: str = DEFAULT_HOST,
port: int = DEFAULT_PORT,
output_dir: str = DEFAULT_OUTPUT_DIR,
default_request: GenerationRequest | None = None,
):
"""Create the app and run it with uvicorn"""
app = create_app(fastvideo_args, output_dir=output_dir)
if default_request is not None:
_validate_default_request_against_preset(default_request, fastvideo_args.model_path)
app = create_app(
fastvideo_args,
output_dir=output_dir,
default_request=default_request,
)
logger.info("Starting FastVideo server on %s:%d", host, port)
logger.info("Model: %s", fastvideo_args.model_path)
+12 -2
View File
@@ -10,6 +10,7 @@ from __future__ import annotations
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from fastvideo.api.schema import GenerationRequest
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.fastvideo_args import FastVideoArgs
@@ -18,6 +19,7 @@ DEFAULT_OUTPUT_DIR = "outputs"
_generator: VideoGenerator | None = None
_fastvideo_args: FastVideoArgs | None = None
_output_dir: str = DEFAULT_OUTPUT_DIR
_default_request: GenerationRequest | None = None
def get_generator() -> VideoGenerator:
@@ -37,20 +39,28 @@ def get_output_dir() -> str:
return _output_dir
def get_default_request() -> GenerationRequest | None:
"""Return the ServeConfig.default_request set at startup, if any."""
return _default_request
def set_state(
generator: VideoGenerator,
fastvideo_args: FastVideoArgs,
output_dir: str,
default_request: GenerationRequest | None = None,
) -> None:
"""Set all server state at once (called from lifespan)."""
global _generator, _fastvideo_args, _output_dir
global _generator, _fastvideo_args, _output_dir, _default_request
_generator = generator
_fastvideo_args = fastvideo_args
_output_dir = output_dir
_default_request = default_request
def clear_state() -> None:
"""Clear server state on shutdown."""
global _generator, _fastvideo_args
global _generator, _fastvideo_args, _default_request
_generator = None
_fastvideo_args = None
_default_request = None
+53 -21
View File
@@ -19,7 +19,10 @@ from fastapi import (
)
from fastapi.responses import FileResponse
from fastvideo.api.compat import explicit_request_updates
from fastvideo.api.schema import GenerationRequest
from fastvideo.entrypoints.openai.state import (
get_default_request,
get_generator,
get_output_dir,
get_server_args,
@@ -42,49 +45,73 @@ logger = init_logger(__name__)
router = APIRouter(prefix="/v1/videos", tags=["videos"])
def _build_generation_kwargs(request_id: str, req: VideoGenerationsRequest) -> dict[str, Any]:
def _build_generation_kwargs(
request_id: str,
req: VideoGenerationsRequest,
default_request: GenerationRequest | None = None,
) -> dict[str, Any]:
"""Build a flat kwargs dict for ``generator.generate_video``.
Precedence (highest to lowest):
1. Request body — only fields the client explicitly sent
(``req.model_fields_set``, Pydantic v2).
2. ``default_request`` — only fields the operator explicitly set in
the serve YAML, projected via ``explicit_request_updates``. Schema
defaults on the dataclass are *not* treated as defaults here.
3. Hardcoded fallback (e.g. ``fps=24`` when neither side set it).
Why gate on ``model_fields_set`` / explicit paths? Both the request
Pydantic model and the ``GenerationRequest`` dataclass carry schema
defaults (e.g. ``seed=1024``, ``num_frames=125``). Without the gate
those would masquerade as intent and shadow the other side — the
gate preserves "operator pinned it" vs. "dataclass happened to have
that default."
"""
kwargs: dict[str, Any] = {}
if default_request is not None:
kwargs.update(explicit_request_updates(default_request))
body_set = req.model_fields_set
kwargs["prompt"] = req.prompt
# Resolution
if req.size:
if "size" in body_set and req.size:
w, h = parse_size(req.size)
if w is not None and h is not None:
kwargs["width"] = w
kwargs["height"] = h
# Frame count / duration
fps = req.fps if req.fps is not None else 24
kwargs["fps"] = fps
if "fps" in body_set and req.fps is not None:
kwargs["fps"] = req.fps
if req.num_frames is not None:
if "num_frames" in body_set and req.num_frames is not None:
kwargs["num_frames"] = req.num_frames
elif req.seconds is not None:
elif "seconds" in body_set and req.seconds is not None:
fps = kwargs.get("fps", 24)
kwargs["num_frames"] = fps * req.seconds
# Sampling parameters
if req.seed is not None:
if "seed" in body_set and req.seed is not None:
kwargs["seed"] = req.seed
if req.num_inference_steps is not None:
if ("num_inference_steps" in body_set and req.num_inference_steps is not None):
kwargs["num_inference_steps"] = req.num_inference_steps
if req.guidance_scale is not None:
if "guidance_scale" in body_set and req.guidance_scale is not None:
kwargs["guidance_scale"] = req.guidance_scale
if req.guidance_scale_2 is not None:
if "guidance_scale_2" in body_set and req.guidance_scale_2 is not None:
kwargs["guidance_scale_2"] = req.guidance_scale_2
if req.negative_prompt is not None:
if "negative_prompt" in body_set and req.negative_prompt is not None:
kwargs["negative_prompt"] = req.negative_prompt
if req.enable_teacache:
if "enable_teacache" in body_set and req.enable_teacache:
kwargs["enable_teacache"] = True
if req.true_cfg_scale is not None:
if "true_cfg_scale" in body_set and req.true_cfg_scale is not None:
kwargs["true_cfg_scale"] = req.true_cfg_scale
# Image-to-video input
if req.input_reference is not None:
if "input_reference" in body_set and req.input_reference is not None:
kwargs["image_path"] = req.input_reference
# Output path
output_dir = req.output_path or os.path.join(get_output_dir(), "videos")
kwargs.setdefault("fps", 24)
default_output_path = kwargs.pop("output_path", None)
body_output_dir = req.output_path if "output_path" in body_set else None
output_dir = body_output_dir or default_output_path or os.path.join(get_output_dir(), "videos")
os.makedirs(output_dir, exist_ok=True)
kwargs["output_path"] = os.path.join(output_dir, f"{request_id}.mp4")
kwargs["save_video"] = True
@@ -272,7 +299,12 @@ async def create_video(
logger.info("Video generation request %s: prompt=%s", request_id, req.prompt[:100])
gen_kwargs = _build_generation_kwargs(request_id, req)
# default_request was validated at server startup (run_server) and is
# read-only on the request hot path — _build_generation_kwargs and
# explicit_request_updates only read, so no per-request deepcopy needed.
default_request = get_default_request()
gen_kwargs = _build_generation_kwargs(request_id, req, default_request=default_request)
job = _make_video_job(request_id, req, gen_kwargs)
await VIDEO_STORE.upsert(request_id, job)
@@ -0,0 +1,31 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.entrypoints.streaming.server import build_app, run_server
from fastvideo.entrypoints.streaming.session import (
Session,
SessionManager,
SessionState,
)
from fastvideo.entrypoints.streaming.session_store import (
BlobStore,
InMemoryBlobStore,
InMemorySessionStore,
SessionStore,
)
from fastvideo.entrypoints.streaming.stream import (
FragmentedMP4Chunk,
FragmentedMP4Encoder,
)
__all__ = [
"BlobStore",
"FragmentedMP4Chunk",
"FragmentedMP4Encoder",
"InMemoryBlobStore",
"InMemorySessionStore",
"Session",
"SessionManager",
"SessionState",
"SessionStore",
"build_app",
"run_server",
]
+252
View File
@@ -0,0 +1,252 @@
# SPDX-License-Identifier: Apache-2.0
"""JSON WebSocket protocol schemas for the streaming server.
Every control message shares the envelope ``{"type": <str>, ...}``.
Pydantic models live here so the server can parse / validate incoming
frames and emit well-typed outgoing frames without hand-rolled dicts.
The message catalogue matches the contract in
``docs/design/server_contracts/streaming.md``; additions must land in
both places in the same PR.
"""
from __future__ import annotations
from typing import Annotated, Any, Literal, Union
from pydantic import BaseModel, ConfigDict, Field
# ---------------------------------------------------------------------------
# Client → server
# ---------------------------------------------------------------------------
class SessionInitV2(BaseModel):
"""Opening frame the client sends after the WebSocket handshake."""
model_config = ConfigDict(extra="allow")
type: Literal["session_init_v2"]
client_id: str | None = None
preset: str | None = None
preset_label: str | None = None
curated_prompts: list[str] = Field(default_factory=list)
initial_image: dict[str, Any] | None = None
enhancement_enabled: bool = False
auto_extension_enabled: bool = False
loop_generation_enabled: bool = False
single_clip_mode: bool = False
stream_mode: Literal["av_fmp4", "legacy_jpeg"] = "av_fmp4"
continuation_state: dict[str, Any] | None = None
"""Optional ``{kind, payload}`` dict; hydrated into
:class:`fastvideo.api.ContinuationState` server-side."""
class SegmentPromptSource(BaseModel):
"""Request a new segment using a specific prompt."""
type: Literal["segment_prompt_source"]
prompt: str
negative_prompt: str | None = None
source: Literal["curated", "enhanced", "user", "auto_extension"] = "user"
seed: int | None = None
num_inference_steps: int | None = None
guidance_scale: float | None = None
class SeedPromptsUpdated(BaseModel):
type: Literal["seed_prompts_updated"]
seed_prompts: list[str] = Field(default_factory=list)
class EnhancementUpdated(BaseModel):
type: Literal["enhancement_updated"]
enabled: bool
class AutoExtensionUpdated(BaseModel):
type: Literal["auto_extension_updated"]
enabled: bool
class LoopGenerationUpdated(BaseModel):
type: Literal["loop_generation_updated"]
enabled: bool
class GenerationPausedUpdated(BaseModel):
type: Literal["generation_paused_updated"]
paused: bool
class SnapshotState(BaseModel):
"""Request the current ``ContinuationState`` for export."""
type: Literal["snapshot_state"]
ClientMessage = Annotated[
Union[ # noqa: UP007 - Annotated requires Union for discriminator
SessionInitV2,
SegmentPromptSource,
SeedPromptsUpdated,
EnhancementUpdated,
AutoExtensionUpdated,
LoopGenerationUpdated,
GenerationPausedUpdated,
SnapshotState,
],
Field(discriminator="type"),
]
# ---------------------------------------------------------------------------
# Server → client
# ---------------------------------------------------------------------------
class QueueStatus(BaseModel):
type: Literal["queue_status"] = "queue_status"
position: int
queue_depth: int
class GpuAssigned(BaseModel):
type: Literal["gpu_assigned"] = "gpu_assigned"
gpu_id: int
session_timeout: int
class Ltx2StreamStart(BaseModel):
type: Literal["ltx2_stream_start"] = "ltx2_stream_start"
preset: str | None = None
width: int
height: int
fps: int
num_frames: int
class Ltx2SegmentStart(BaseModel):
type: Literal["ltx2_segment_start"] = "ltx2_segment_start"
segment_idx: int
prompt: str
total_steps: int
class StepComplete(BaseModel):
type: Literal["step_complete"] = "step_complete"
segment_idx: int
step: int
total_steps: int
stage: str = "denoise"
class MediaInit(BaseModel):
"""Descriptor for the fMP4 initialization segment that follows."""
type: Literal["media_init"] = "media_init"
segment_idx: int
mime: str = "video/mp4; codecs=\"avc1.64001f, mp4a.40.2\""
stream_id: str
mode: Literal["av_fmp4"] = "av_fmp4"
class MediaSegmentComplete(BaseModel):
type: Literal["media_segment_complete"] = "media_segment_complete"
segment_idx: int
stream_id: str
chunks: int
duration_ms: float | None = None
pts_base_ms: float | None = None
class Ltx2SegmentComplete(BaseModel):
type: Literal["ltx2_segment_complete"] = "ltx2_segment_complete"
segment_idx: int
generation_time_ms: float
e2e_latency_ms: float | None = None
class Ltx2StreamComplete(BaseModel):
type: Literal["ltx2_stream_complete"] = "ltx2_stream_complete"
reason: Literal["segment_cap", "stop_requested", "error"] = "stop_requested"
class SessionTimeout(BaseModel):
type: Literal["session_timeout"] = "session_timeout"
timeout_seconds: int
class ContinuationStateSnapshot(BaseModel):
type: Literal["continuation_state_snapshot"] = "continuation_state_snapshot"
state: dict[str, Any]
"""``{kind, payload}`` dict matching
:class:`fastvideo.api.ContinuationState`."""
class ErrorMessage(BaseModel):
type: Literal["error"] = "error"
code: Literal[
"session_rejected",
"invalid_message",
"preset_mismatch",
"gpu_unavailable",
"worker_failed",
"upstream_timeout",
"internal_error",
] = "internal_error"
message: str
retryable: bool = False
ServerMessage = Union[ # noqa: UP007 - pydantic Union handling
QueueStatus,
GpuAssigned,
Ltx2StreamStart,
Ltx2SegmentStart,
StepComplete,
MediaInit,
MediaSegmentComplete,
Ltx2SegmentComplete,
Ltx2StreamComplete,
SessionTimeout,
ContinuationStateSnapshot,
ErrorMessage,
]
def parse_client_message(raw: dict[str, Any]) -> ClientMessage:
"""Parse an incoming WebSocket dict into a typed client message.
Unknown ``type`` values raise :class:`pydantic.ValidationError`; the
server handler turns that into an ``error`` frame with
``code="invalid_message"``.
"""
from pydantic import TypeAdapter
return TypeAdapter(ClientMessage).validate_python(raw)
__all__ = [
"AutoExtensionUpdated",
"ClientMessage",
"ContinuationStateSnapshot",
"EnhancementUpdated",
"ErrorMessage",
"GenerationPausedUpdated",
"GpuAssigned",
"Ltx2SegmentComplete",
"Ltx2SegmentStart",
"Ltx2StreamComplete",
"Ltx2StreamStart",
"LoopGenerationUpdated",
"MediaInit",
"MediaSegmentComplete",
"QueueStatus",
"SeedPromptsUpdated",
"SegmentPromptSource",
"ServerMessage",
"SessionInitV2",
"SessionTimeout",
"SnapshotState",
"StepComplete",
"parse_client_message",
]
+531
View File
@@ -0,0 +1,531 @@
# SPDX-License-Identifier: Apache-2.0
"""Single-generator FastAPI + WebSocket streaming server."""
from __future__ import annotations
import asyncio
import contextlib
import os
import time
from dataclasses import dataclass
from typing import Any, Protocol
from fastapi import FastAPI, WebSocket, WebSocketDisconnect
from fastapi.responses import JSONResponse
from fastvideo.api.schema import (
ContinuationState,
GenerationRequest,
InputConfig,
OutputConfig,
SamplingConfig,
ServeConfig,
)
from fastvideo.entrypoints.streaming.protocol import (
AutoExtensionUpdated,
ContinuationStateSnapshot,
EnhancementUpdated,
ErrorMessage,
GenerationPausedUpdated,
GpuAssigned,
LoopGenerationUpdated,
Ltx2SegmentComplete,
Ltx2SegmentStart,
Ltx2StreamComplete,
Ltx2StreamStart,
MediaInit,
MediaSegmentComplete,
QueueStatus,
SeedPromptsUpdated,
SegmentPromptSource,
SessionInitV2,
SnapshotState,
StepComplete,
parse_client_message,
)
from fastvideo.entrypoints.streaming.session import (
InvalidSessionTransition,
Session,
SessionManager,
SessionRejected,
SessionState,
)
from fastvideo.entrypoints.streaming.session_init_image import (
persist_session_init_image, )
from fastvideo.entrypoints.streaming.session_store import (
InMemorySessionStore,
SessionStore,
)
from fastvideo.entrypoints.streaming.stream import FragmentedMP4Encoder
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# RFC 6455 WebSocket close codes used by the server.
_WS_CLOSE_UNSUPPORTED_DATA = 1003
_WS_CLOSE_TRY_AGAIN_LATER = 1013
class _GeneratorProto(Protocol):
"""Subset of :class:`fastvideo.VideoGenerator` the server calls."""
def generate(self, request: GenerationRequest) -> Any:
...
@dataclass
class ServerState:
serve_config: ServeConfig
generator: _GeneratorProto
sessions: SessionManager
session_store: SessionStore
def build_app(
serve_config: ServeConfig,
generator: _GeneratorProto,
*,
session_store: SessionStore | None = None,
) -> FastAPI:
"""Build the FastAPI app used by :func:`run_server`.
Exposed so tests can drive the WebSocket endpoint in-process via
``starlette.testclient.TestClient(app).websocket_connect(...)``.
"""
if serve_config.streaming is None:
raise ValueError("ServeConfig.streaming must be set to launch the streaming "
"server; got None. Add a `streaming:` block to your serve config.")
sessions = SessionManager(
segment_cap=serve_config.streaming.generation_segment_cap,
session_timeout_seconds=serve_config.streaming.session_timeout_seconds,
)
state = ServerState(
serve_config=serve_config,
generator=generator,
sessions=sessions,
session_store=session_store or InMemorySessionStore(),
)
app = FastAPI(title="FastVideo Streaming")
@app.get("/health")
async def _health() -> JSONResponse:
return JSONResponse({
"status": "ok",
"sessions": len(state.sessions),
"stream_mode": state.serve_config.streaming.stream_mode,
})
@app.websocket("/v1/stream")
async def _stream(websocket: WebSocket) -> None:
await websocket.accept()
try:
session = state.sessions.create()
except SessionRejected as exc:
await _send_error(websocket, "session_rejected", str(exc), retryable=False)
await websocket.close(code=_WS_CLOSE_TRY_AGAIN_LATER, reason="session_rejected")
return
try:
await _handle_session(websocket, session, state)
except WebSocketDisconnect:
logger.info("session %s: client disconnected", session.id[:8])
except Exception: # pragma: no cover - defensive catch-all
logger.exception("session %s: unhandled error", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
finally:
_cleanup_session(session, state)
app.state.server_state = state
return app
def run_server(serve_config: ServeConfig, *, generator: _GeneratorProto | None = None) -> None:
"""Launch the streaming server.
Boots a :class:`fastvideo.VideoGenerator` from
``serve_config.generator`` unless ``generator`` is provided, then
serves ``build_app(...)`` via uvicorn.
"""
if serve_config.streaming is None:
raise ValueError("ServeConfig.streaming must be set to launch the streaming server; "
"got None. Add a `streaming:` block to your serve config.")
import uvicorn
if generator is None:
from fastvideo import VideoGenerator # lazy to avoid boot cost
generator = VideoGenerator.from_pretrained(config=serve_config.generator)
app = build_app(serve_config, generator)
uvicorn.run(
app,
host=serve_config.server.host,
port=serve_config.server.port,
)
async def _handle_session(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> None:
init = await _read_init_message(websocket, session, state)
if init is None:
return
await _apply_session_init(session, init, state)
await _send_json(websocket, QueueStatus(position=0, queue_depth=0))
session.transition(SessionState.GPU_BINDING)
await _send_json(websocket, GpuAssigned(
gpu_id=0,
session_timeout=state.sessions.session_timeout_seconds,
))
session.transition(SessionState.ACTIVE)
await _send_json(websocket, _build_stream_start(session, state))
try:
await _run_segment_loop(websocket, session, state)
finally:
with contextlib.suppress(RuntimeError):
await _send_json(websocket, Ltx2StreamComplete(reason="stop_requested"))
async def _read_init_message(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> SessionInitV2 | None:
try:
raw = await asyncio.wait_for(
websocket.receive_json(),
timeout=state.sessions.session_timeout_seconds,
)
except asyncio.TimeoutError:
logger.info("session %s: init timeout", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.TIMEOUT)
return None
except WebSocketDisconnect:
return None
try:
parsed = parse_client_message(raw)
except Exception as exc:
await _reject_init(websocket, session, f"opening frame failed validation: {exc}", "invalid_init")
return None
if not isinstance(parsed, SessionInitV2):
await _reject_init(websocket, session, "first frame must be session_init_v2", "expected_session_init_v2")
return None
return parsed
async def _reject_init(
websocket: WebSocket,
session: Session,
message: str,
close_reason: str,
) -> None:
await _send_error(websocket, "invalid_message", message, retryable=False)
await websocket.close(code=_WS_CLOSE_UNSUPPORTED_DATA, reason=close_reason)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.REJECTED)
async def _apply_session_init(
session: Session,
init: SessionInitV2,
state: ServerState,
) -> None:
session.client_id = init.client_id
session.preset = init.preset
session.preset_label = init.preset_label
session.curated_prompts = list(init.curated_prompts)
session.enhancement_enabled = init.enhancement_enabled
session.auto_extension_enabled = init.auto_extension_enabled
session.loop_generation_enabled = init.loop_generation_enabled
session.single_clip_mode = init.single_clip_mode
session.stream_mode = init.stream_mode
if init.initial_image is not None:
# Decode + disk write off the event loop; payload is up to 32 MiB.
image = await asyncio.to_thread(persist_session_init_image, init.initial_image)
if image is not None:
session.metadata["session_init_image"] = image.path
if init.continuation_state is not None:
session.continuation_state = _coerce_state(init.continuation_state)
if session.continuation_state is not None:
state.session_store.store(session.id, session.continuation_state)
async def _run_segment_loop(
websocket: WebSocket,
session: Session,
state: ServerState,
) -> None:
cap = state.sessions.segment_cap
while True:
if session.segment_cap_reached(cap):
logger.info("session %s: segment cap (%d) reached", session.id[:8], cap)
return
try:
raw = await asyncio.wait_for(
websocket.receive_json(),
timeout=state.sessions.session_timeout_seconds,
)
except asyncio.TimeoutError:
logger.info("session %s: idle timeout", session.id[:8])
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.TIMEOUT)
return
except WebSocketDisconnect:
return
session.touch()
try:
parsed = parse_client_message(raw)
except Exception as exc:
await _send_error(websocket, "invalid_message", str(exc), retryable=True)
continue
if isinstance(parsed, SnapshotState):
snap = state.session_store.snapshot(session.id)
if snap is None:
await _send_error(websocket,
"internal_error",
"no continuation state available for session",
retryable=False)
continue
await _send_json(websocket, ContinuationStateSnapshot(state={"kind": snap.kind, "payload": snap.payload}, ))
continue
if isinstance(parsed, SegmentPromptSource):
await _run_segment(websocket, session, state, parsed)
continue
# Silently ignore unknown-but-valid types (additive-evolution
# rule in streaming.md).
_apply_toggle(session, parsed)
async def _run_segment(
websocket: WebSocket,
session: Session,
state: ServerState,
message: SegmentPromptSource,
) -> None:
request = _build_generation_request(session, message, state)
segment_idx = session.segment_idx
await _send_json(
websocket,
Ltx2SegmentStart(
segment_idx=segment_idx,
prompt=message.prompt,
total_steps=request.sampling.num_inference_steps,
))
start = time.perf_counter()
loop = asyncio.get_running_loop()
# TODO: executor-wrapped generate() cannot be cancelled, so a
# client disconnect mid-segment leaves the GPU work running to
# completion. Real cancellation needs the generate_async API.
try:
result = await loop.run_in_executor(None, state.generator.generate, request)
except Exception as exc:
logger.exception("session %s: generator failed", session.id[:8])
await _send_error(websocket, "worker_failed", f"generator.generate failed: {exc}", retryable=True)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
return
elapsed_ms = (time.perf_counter() - start) * 1000.0
frames = _extract_frames(result)
if not frames:
await _send_error(websocket, "worker_failed", "generator returned no frames", retryable=True)
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ERROR)
return
# Synchronous generator call has no per-step hook; emit one
# terminal StepComplete so observability wiring still sees the
# segment finish.
total = request.sampling.num_inference_steps
await _send_json(websocket, StepComplete(
segment_idx=segment_idx,
step=total,
total_steps=total,
stage="denoise",
))
encoder = FragmentedMP4Encoder(
width=request.sampling.width,
height=request.sampling.height,
fps=request.sampling.fps,
segment_idx=segment_idx,
)
chunks_relayed = 0
async with encoder:
init_sent = False
async for chunk in encoder.encode(frames):
if chunk.kind == "init":
await _send_json(websocket, MediaInit(
segment_idx=segment_idx,
stream_id=chunk.stream_id,
))
init_sent = True
await websocket.send_bytes(chunk.data)
if init_sent and chunk.kind == "media":
chunks_relayed += 1
await _send_json(
websocket,
MediaSegmentComplete(
segment_idx=segment_idx,
stream_id=encoder.stream_id,
chunks=chunks_relayed,
duration_ms=float(request.sampling.num_frames) / request.sampling.fps * 1000.0,
))
new_state = _extract_state(result)
if new_state is not None:
session.continuation_state = new_state
state.session_store.store(session.id, new_state)
session.segment_idx += 1
with contextlib.suppress(InvalidSessionTransition):
session.transition(SessionState.ACTIVE)
await _send_json(
websocket,
Ltx2SegmentComplete(
segment_idx=segment_idx,
generation_time_ms=elapsed_ms,
e2e_latency_ms=elapsed_ms,
))
def _build_stream_start(
session: Session,
state: ServerState,
) -> Ltx2StreamStart:
default = state.serve_config.default_request
return Ltx2StreamStart(
preset=session.preset,
width=default.sampling.width,
height=default.sampling.height,
fps=default.sampling.fps,
num_frames=default.sampling.num_frames,
)
def _build_generation_request(
session: Session,
message: SegmentPromptSource,
state: ServerState,
) -> GenerationRequest:
# Start from the operator-pinned default_request to pick up the
# preset-selected sampling knobs; override with per-message values.
base = state.serve_config.default_request
sampling_kwargs: dict[str, Any] = {
"num_videos_per_prompt":
base.sampling.num_videos_per_prompt,
"seed":
message.seed if message.seed is not None else base.sampling.seed,
"num_frames":
base.sampling.num_frames,
"height":
base.sampling.height,
"width":
base.sampling.width,
"fps":
base.sampling.fps,
"num_inference_steps":
(message.num_inference_steps if message.num_inference_steps is not None else base.sampling.num_inference_steps),
"guidance_scale":
(message.guidance_scale if message.guidance_scale is not None else base.sampling.guidance_scale),
}
request = GenerationRequest(
prompt=message.prompt,
negative_prompt=message.negative_prompt or base.negative_prompt,
inputs=InputConfig(image_path=session.metadata.get("session_init_image"), ),
sampling=SamplingConfig(**sampling_kwargs),
output=OutputConfig(save_video=False, return_frames=True, return_state=True),
state=session.continuation_state,
)
return request
def _coerce_state(raw: dict[str, Any]) -> ContinuationState | None:
kind = raw.get("kind")
payload = raw.get("payload")
if not isinstance(kind, str) or not isinstance(payload, dict):
return None
return ContinuationState(kind=kind, payload=payload)
def _apply_toggle(session: Session, message: Any) -> None:
if isinstance(message, EnhancementUpdated):
session.enhancement_enabled = message.enabled
elif isinstance(message, AutoExtensionUpdated):
session.auto_extension_enabled = message.enabled
elif isinstance(message, LoopGenerationUpdated):
session.loop_generation_enabled = message.enabled
elif isinstance(message, GenerationPausedUpdated):
session.generation_paused = message.paused
elif isinstance(message, SeedPromptsUpdated):
session.curated_prompts = list(message.seed_prompts)
def _extract_frames(result: Any) -> list:
if hasattr(result, "frames"):
return list(result.frames or [])
if isinstance(result, dict):
return list(result.get("frames") or [])
return []
def _extract_state(result: Any) -> ContinuationState | None:
state = getattr(result, "state", None)
if state is None and isinstance(result, dict):
state = result.get("state")
if isinstance(state, ContinuationState):
return state
if isinstance(state, dict):
return _coerce_state(state)
return None
async def _send_json(websocket: WebSocket, message: Any) -> None:
payload = (message.model_dump(mode="json", exclude_none=True) if hasattr(message, "model_dump") else message)
await websocket.send_json(payload)
async def _send_error(
websocket: WebSocket,
code: str,
message: str,
*,
retryable: bool,
) -> None:
await _send_json(
websocket,
ErrorMessage(code=code, message=message, retryable=retryable),
)
def _cleanup_session(session: Session, state: ServerState) -> None:
state.sessions.close(session.id)
state.session_store.drop(session.id)
init_image_path = session.metadata.get("session_init_image")
if isinstance(init_image_path, str):
with contextlib.suppress(FileNotFoundError):
os.unlink(init_image_path)
__all__ = [
"ServerState",
"build_app",
"run_server",
]
+214
View File
@@ -0,0 +1,214 @@
# SPDX-License-Identifier: Apache-2.0
"""Per-connection session lifecycle for the streaming server.
Each WebSocket opens exactly one :class:`Session`. :class:`SessionManager`
enforces the ``generation_segment_cap`` and ``session_timeout_seconds``
budgets from :class:`fastvideo.api.StreamingConfig`.
"""
from __future__ import annotations
import enum
import time
import uuid
from dataclasses import dataclass, field
from typing import Any
from fastvideo.api.schema import ContinuationState
class SessionState(enum.Enum):
"""State-machine positions for a streaming session.
Transitions are server-owned. See
``docs/design/server_contracts/streaming.md`` for the full diagram.
"""
INITIALIZING = "initializing"
QUEUED = "queued"
GPU_BINDING = "gpu_binding"
ACTIVE = "active"
COMPLETE = "complete"
ERROR = "error"
TIMEOUT = "timeout"
REJECTED = "rejected"
_VALID_TRANSITIONS: dict[SessionState, frozenset[SessionState]] = {
SessionState.INITIALIZING:
frozenset({
SessionState.QUEUED,
SessionState.GPU_BINDING,
SessionState.REJECTED,
SessionState.ERROR,
}),
SessionState.QUEUED:
frozenset({
SessionState.GPU_BINDING,
SessionState.ERROR,
SessionState.TIMEOUT,
SessionState.REJECTED,
}),
SessionState.GPU_BINDING:
frozenset({
SessionState.ACTIVE,
SessionState.ERROR,
SessionState.TIMEOUT,
}),
SessionState.ACTIVE:
frozenset({
SessionState.ACTIVE,
SessionState.COMPLETE,
SessionState.ERROR,
SessionState.TIMEOUT,
}),
SessionState.COMPLETE:
frozenset(),
SessionState.ERROR:
frozenset(),
SessionState.TIMEOUT:
frozenset(),
SessionState.REJECTED:
frozenset(),
}
class InvalidSessionTransition(RuntimeError):
"""Raised when a session is asked to transition along an illegal edge."""
@dataclass
class Session:
id: str = field(default_factory=lambda: uuid.uuid4().hex)
state: SessionState = SessionState.INITIALIZING
created_at: float = field(default_factory=time.monotonic)
last_activity: float = field(default_factory=time.monotonic)
client_id: str | None = None
preset: str | None = None
preset_label: str | None = None
curated_prompts: list[str] = field(default_factory=list)
segment_idx: int = 0
enhancement_enabled: bool = False
auto_extension_enabled: bool = False
loop_generation_enabled: bool = False
single_clip_mode: bool = False
generation_paused: bool = False
stream_mode: str = "av_fmp4"
gpu_id: int | None = None
continuation_state: ContinuationState | None = None
metadata: dict[str, Any] = field(default_factory=dict)
def transition(self, target: SessionState) -> None:
"""Move to ``target`` if the edge is allowed.
Raises :class:`InvalidSessionTransition` on illegal moves. The
self-loop on ``ACTIVE`` is legal so the server can re-assert
ACTIVE on segment completion without special casing.
"""
allowed = _VALID_TRANSITIONS.get(self.state, frozenset())
if target not in allowed and target is not self.state:
raise InvalidSessionTransition(f"{self.state.value} -> {target.value} is not a valid "
f"session transition")
self.state = target
self.last_activity = time.monotonic()
def touch(self) -> None:
self.last_activity = time.monotonic()
def is_active(self) -> bool:
return self.state is SessionState.ACTIVE
def segment_cap_reached(self, cap: int) -> bool:
return self.segment_idx >= cap
class SessionManager:
"""Registers sessions and enforces per-server session limits."""
def __init__(
self,
*,
segment_cap: int,
session_timeout_seconds: int,
max_sessions: int = 1,
) -> None:
self._segment_cap = segment_cap
self._session_timeout_seconds = session_timeout_seconds
self._max_sessions = max_sessions
self._sessions: dict[str, Session] = {}
@property
def segment_cap(self) -> int:
return self._segment_cap
@property
def session_timeout_seconds(self) -> int:
return self._session_timeout_seconds
def create(self) -> Session:
if len(self._sessions) >= self._max_sessions:
raise SessionRejected(f"max sessions reached ({self._max_sessions})")
session = Session()
self._sessions[session.id] = session
return session
def get(self, session_id: str) -> Session | None:
return self._sessions.get(session_id)
def close(self, session_id: str) -> None:
self._sessions.pop(session_id, None)
def __contains__(self, session_id: str) -> bool:
return session_id in self._sessions
def __len__(self) -> int:
return len(self._sessions)
def active_sessions(self) -> list[Session]:
return [s for s in self._sessions.values() if s.is_active()]
def reap_timed_out(self, now: float | None = None) -> list[str]:
"""Return the ids of sessions that have exceeded the idle timeout.
The caller is responsible for actually closing them — this
method only *identifies* dead sessions so the server can emit
``session_timeout`` frames before dropping the WebSocket.
TODO: unused until a background driver calls it. Per-connection
idle enforcement currently happens via asyncio.wait_for on
receive_json; this helper catches sessions stuck before any
receive (e.g. future QUEUED state) and is expected to be wired
into the GPU-pool reaper.
"""
now = now if now is not None else time.monotonic()
dead: list[str] = []
for sid, session in self._sessions.items():
if session.state in {
SessionState.COMPLETE,
SessionState.ERROR,
SessionState.TIMEOUT,
SessionState.REJECTED,
}:
continue
if now - session.last_activity > self._session_timeout_seconds:
dead.append(sid)
return dead
class SessionRejected(RuntimeError):
"""Raised when session creation fails (queue full, auth, etc.)."""
__all__ = [
"InvalidSessionTransition",
"Session",
"SessionManager",
"SessionRejected",
"SessionState",
]
@@ -0,0 +1,103 @@
# SPDX-License-Identifier: Apache-2.0
"""Persist the initial-image blob attached to a streaming session."""
from __future__ import annotations
import base64
import binascii
import contextlib
import os
import tempfile
from dataclasses import dataclass
from typing import Any
_ACCEPTED_MIMES = {
"image/png": ".png",
"image/jpeg": ".jpg",
"image/jpg": ".jpg",
"image/webp": ".webp",
}
_MAX_IMAGE_BYTES = 32 * 1024 * 1024 # 32 MiB cap
@dataclass(frozen=True)
class SessionInitImage:
"""Location of the persisted init image.
Callers pass ``path`` to ``InputConfig.image_path``; ``display_name``
is only used for logs.
"""
path: str
display_name: str
mime: str
def persist_session_init_image(
payload: Any,
*,
output_dir: str | None = None,
) -> SessionInitImage | None:
"""Decode a client init-image blob and persist it to disk.
``payload`` shape (matches the internal UI protocol)::
{
"mime": "image/png",
"name": "ref.png",
"data": "<base64 bytes>",
}
Returns ``None`` when ``payload`` is falsy (no init image). Raises
:class:`ValueError` on schema / size / decode errors so the caller
can surface a user-facing ``error`` frame.
"""
if not payload:
return None
if not isinstance(payload, dict):
raise ValueError("session init image must be an object")
mime = payload.get("mime")
if mime not in _ACCEPTED_MIMES:
raise ValueError(f"session init image mime {mime!r} is not one of "
f"{sorted(_ACCEPTED_MIMES)}")
data_b64 = payload.get("data")
if not isinstance(data_b64, str):
raise ValueError("session init image data must be a base64 string")
try:
data = base64.b64decode(data_b64, validate=True)
except (binascii.Error, ValueError) as exc:
raise ValueError(f"session init image data is not valid base64: {exc}") from exc
if len(data) > _MAX_IMAGE_BYTES:
raise ValueError(f"session init image is {len(data)} bytes; limit is "
f"{_MAX_IMAGE_BYTES}")
if len(data) == 0:
raise ValueError("session init image data is empty")
ext = _ACCEPTED_MIMES[mime]
display_name = _sanitize_display_name(payload.get("name")) or f"init{ext}"
fd, path = tempfile.mkstemp(prefix="fastvideo-init-", suffix=ext, dir=output_dir)
try:
with os.fdopen(fd, "wb") as f:
f.write(data)
except Exception:
with contextlib.suppress(FileNotFoundError):
os.unlink(path)
raise
return SessionInitImage(path=path, display_name=display_name, mime=mime)
def _sanitize_display_name(name: Any) -> str | None:
if not isinstance(name, str):
return None
name = name.strip()
if not name:
return None
# Strip any path components — we only keep the leaf for logging.
return os.path.basename(name)
__all__ = [
"SessionInitImage",
"persist_session_init_image",
]
@@ -0,0 +1,206 @@
# SPDX-License-Identifier: Apache-2.0
"""Session state store for the FastVideo streaming server.
The streaming server keeps continuation state (decoded frames + audio
latents from the previous segment) server-side so the client doesn't
re-upload multi-megabyte tensors each WebSocket message. Two operations
are needed:
* ``snapshot(session_id) -> ContinuationState`` — serialize the current
state so it can be exported (e.g. over HTTP) or migrated to a
different server.
* ``hydrate(state) -> session_id`` — load a previously serialized state
into a new session (for resume-after-disconnect flows).
The store is an ABC with an :class:`InMemorySessionStore` default; Redis
or other backends can drop in without touching the pipeline.
Large tensor payloads (video frames, audio latents) are kept out of the
JSON payload via an accompanying :class:`BlobStore`. Both stores share a
process today; they are separate types so that a future implementation
can put blobs on S3 while keeping session metadata in Redis.
"""
from __future__ import annotations
import threading
import uuid
from abc import ABC, abstractmethod
from collections.abc import Iterator
from dataclasses import dataclass
from fastvideo.api.schema import ContinuationState
class BlobStore(ABC):
"""Opaque byte-blob storage keyed by id.
A :class:`ContinuationState` payload can reference large tensors
stored in a :class:`BlobStore` rather than inlining them, so the
JSON payload stays small when the state travels over the wire.
"""
@abstractmethod
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
"""Store ``data`` and return a blob id for later retrieval."""
@abstractmethod
def get(self, blob_id: str) -> bytes:
"""Load a previously stored blob. Raises ``KeyError`` if absent."""
@abstractmethod
def drop(self, blob_id: str) -> None:
"""Remove a blob. Missing ids are a no-op."""
@abstractmethod
def __contains__(self, blob_id: str) -> bool:
...
@dataclass(frozen=True)
class _BlobRecord:
data: bytes
mime: str
class InMemoryBlobStore(BlobStore):
"""Thread-safe in-memory :class:`BlobStore` for single-process servers.
No eviction policy — callers are responsible for calling
:meth:`drop` when a blob's owning state is replaced or a session
ends. A redis- or filesystem-backed :class:`BlobStore` should
replace this when the streaming server lands as a real service
(PR 7.5+).
"""
def __init__(self) -> None:
self._blobs: dict[str, _BlobRecord] = {}
self._lock = threading.Lock()
def put(self, data: bytes, *, mime: str = "application/octet-stream") -> str:
blob_id = uuid.uuid4().hex
with self._lock:
self._blobs[blob_id] = _BlobRecord(data=data, mime=mime)
return blob_id
def get(self, blob_id: str) -> bytes:
with self._lock:
record = self._blobs.get(blob_id)
if record is None:
raise KeyError(f"Unknown blob id: {blob_id}")
return record.data
def drop(self, blob_id: str) -> None:
with self._lock:
self._blobs.pop(blob_id, None)
def __contains__(self, blob_id: str) -> bool:
with self._lock:
return blob_id in self._blobs
def __len__(self) -> int:
with self._lock:
return len(self._blobs)
class SessionStore(ABC):
"""Keyed store for per-session continuation state.
Implementations own the session-id → state mapping. The streaming
server calls :meth:`store` after each segment and :meth:`snapshot`
when a client explicitly asks for an exportable state handle.
"""
@abstractmethod
def store(self, session_id: str, state: ContinuationState) -> None:
"""Persist ``state`` for ``session_id``, replacing any prior value."""
@abstractmethod
def snapshot(self, session_id: str) -> ContinuationState | None:
"""Return the current state for ``session_id`` (or ``None``)."""
@abstractmethod
def hydrate(
self,
state: ContinuationState,
*,
session_id: str | None = None,
) -> str:
"""Install ``state`` as the starting point for a session.
When ``session_id`` is ``None`` the store allocates a fresh id
(UUID4); when provided the store uses it verbatim, overwriting
any prior state at that id.
"""
@abstractmethod
def drop(self, session_id: str) -> None:
"""Forget a session. Missing ids are a no-op."""
@abstractmethod
def __contains__(self, session_id: str) -> bool:
...
@abstractmethod
def __iter__(self) -> Iterator[str]:
...
class InMemorySessionStore(SessionStore):
"""Thread-safe in-memory :class:`SessionStore`.
Default implementation used by single-process deployments; a future
Redis-backed store can be dropped in without changes to the server.
No eviction / TTL / bounded capacity — sessions only leave via
:meth:`drop`. The live streaming server (PR 7.5+) is responsible
for bounding growth and for dropping any :class:`BlobStore` blobs
referenced by a state when that state is replaced or a session
ends; this class does not know about blobs.
"""
def __init__(self) -> None:
self._sessions: dict[str, ContinuationState] = {}
self._lock = threading.Lock()
def store(self, session_id: str, state: ContinuationState) -> None:
with self._lock:
self._sessions[session_id] = state
def snapshot(self, session_id: str) -> ContinuationState | None:
with self._lock:
return self._sessions.get(session_id)
def hydrate(
self,
state: ContinuationState,
*,
session_id: str | None = None,
) -> str:
sid = session_id or uuid.uuid4().hex
with self._lock:
self._sessions[sid] = state
return sid
def drop(self, session_id: str) -> None:
with self._lock:
self._sessions.pop(session_id, None)
def __contains__(self, session_id: str) -> bool:
with self._lock:
return session_id in self._sessions
def __iter__(self) -> Iterator[str]:
with self._lock:
return iter(list(self._sessions))
def __len__(self) -> int:
with self._lock:
return len(self._sessions)
__all__ = [
"BlobStore",
"InMemoryBlobStore",
"InMemorySessionStore",
"SessionStore",
]
+213
View File
@@ -0,0 +1,213 @@
# SPDX-License-Identifier: Apache-2.0
"""fMP4 stream encoder used by the streaming server.
The client's Media Source Extensions player needs a continuous fMP4
byte stream: first an *initialization segment* (``ftyp`` + ``moov``),
then one or more *media segments* (``moof`` + ``mdat``). We pipe raw
RGB frames into an ffmpeg subprocess configured for fragmented output
via ``-movflags empty_moov+default_base_moof+frag_keyframe+faststart``
and stream the bytes back out.
"""
from __future__ import annotations
import asyncio
import contextlib
import subprocess
import uuid
from collections.abc import AsyncIterator
from dataclasses import dataclass
from typing import TYPE_CHECKING, Literal
if TYPE_CHECKING:
import numpy as np
@dataclass
class FragmentedMP4Chunk:
"""A single fMP4 byte chunk emitted by :class:`FragmentedMP4Encoder`.
``kind`` identifies whether the chunk is the init segment (must be
fed into the client's ``SourceBuffer`` first) or a media fragment.
"""
kind: Literal["init", "media"]
data: bytes
stream_id: str
segment_idx: int
class FragmentedMP4Encoder:
"""Stream RGB frames in, fMP4 chunks out.
One encoder covers one segment. The server creates a new encoder
per :class:`ltx2_segment_start`` boundary so each segment becomes
one media fragment the client can append independently.
Example::
encoder = FragmentedMP4Encoder(width=1024, height=576, fps=24,
segment_idx=0)
async with encoder:
async for chunk in encoder.encode(frames):
await websocket.send_bytes(chunk.data)
"""
def __init__(
self,
*,
width: int,
height: int,
fps: int,
segment_idx: int,
stream_id: str | None = None,
ffmpeg_path: str = "ffmpeg",
preset: str = "ultrafast",
pixel_format_out: str = "yuv420p",
extra_args: list[str] | None = None,
) -> None:
self.width = width
self.height = height
self.fps = fps
self.segment_idx = segment_idx
self.stream_id = stream_id or uuid.uuid4().hex
self._ffmpeg_path = ffmpeg_path
self._preset = preset
self._pixel_format_out = pixel_format_out
self._extra_args = list(extra_args or [])
self._proc: subprocess.Popen | None = None
self._init_emitted = False
async def __aenter__(self) -> FragmentedMP4Encoder:
self._spawn()
return self
async def __aexit__(self, exc_type, exc, tb) -> None:
await self.close()
def _spawn(self) -> None:
args = [
self._ffmpeg_path,
"-hide_banner",
"-loglevel",
"error",
"-f",
"rawvideo",
"-pix_fmt",
"rgb24",
"-s",
f"{self.width}x{self.height}",
"-r",
str(self.fps),
"-i",
"-",
"-c:v",
"libx264",
"-preset",
self._preset,
"-tune",
"zerolatency",
"-pix_fmt",
self._pixel_format_out,
"-movflags",
"empty_moov+default_base_moof+frag_keyframe+faststart",
"-f",
"mp4",
*self._extra_args,
"-",
]
# stderr → DEVNULL: with -loglevel error on, the only thing
# stderr would carry is unsolicited warnings. Piping without a
# reader deadlocks ffmpeg once the pipe buffer (~64 KiB) fills.
self._proc = subprocess.Popen( # noqa: S603
args,
stdin=subprocess.PIPE,
stdout=subprocess.PIPE,
stderr=subprocess.DEVNULL,
bufsize=0,
)
async def encode(
self,
frames: list[np.ndarray] | AsyncIterator[np.ndarray],
) -> AsyncIterator[FragmentedMP4Chunk]:
"""Feed frames into ffmpeg and yield fMP4 chunks as they appear."""
if self._proc is None:
self._spawn()
assert self._proc is not None and self._proc.stdin is not None
proc = self._proc
loop = asyncio.get_running_loop()
async def _writer() -> None:
try:
if hasattr(frames, "__aiter__"):
async for frame in frames: # type: ignore[union-attr]
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
else:
for frame in frames: # type: ignore[assignment]
await loop.run_in_executor(None, _write_frame, proc.stdin, frame)
finally:
with contextlib.suppress(BrokenPipeError):
proc.stdin.close()
writer_task = asyncio.create_task(_writer())
try:
reader = proc.stdout
assert reader is not None
# Read in reasonably-sized chunks; MSE tolerates any size
# but we don't want to starve the event loop.
chunk_size = 64 * 1024
while True:
data = await loop.run_in_executor(None, reader.read, chunk_size)
if not data:
break
kind: Literal["init", "media"] = "init" if not self._init_emitted else "media"
self._init_emitted = True
yield FragmentedMP4Chunk(
kind=kind,
data=bytes(data),
stream_id=self.stream_id,
segment_idx=self.segment_idx,
)
finally:
await writer_task
async def close(self) -> None:
if self._proc is None:
return
proc = self._proc
self._proc = None
try:
if proc.stdin and not proc.stdin.closed:
proc.stdin.close()
except BrokenPipeError:
pass
loop = asyncio.get_running_loop()
try:
await asyncio.wait_for(
loop.run_in_executor(None, proc.wait),
timeout=5.0,
)
except asyncio.TimeoutError:
proc.kill()
await loop.run_in_executor(None, proc.wait)
def _write_frame(stdin, frame: np.ndarray) -> None:
import numpy as np
if not isinstance(frame, np.ndarray):
raise TypeError("fMP4 encoder frames must be numpy.ndarray")
if frame.dtype != np.uint8:
frame = frame.astype(np.uint8)
if frame.ndim != 3 or frame.shape[-1] != 3:
raise ValueError("fMP4 encoder frames must be HxWx3 uint8 RGB; got "
f"shape={frame.shape}, dtype={frame.dtype}")
with contextlib.suppress(BrokenPipeError):
stdin.write(frame.tobytes())
__all__ = [
"FragmentedMP4Chunk",
"FragmentedMP4Encoder",
]
+1 -1
View File
@@ -8,7 +8,7 @@ import torch
import torchvision
from einops import rearrange
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.entrypoints.video_generator import VideoGenerator
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
+1 -1
View File
@@ -35,7 +35,7 @@ from fastvideo.api.compat import (
)
from fastvideo.api.results import GenerationResult
from fastvideo.api.schema import GenerationRequest, GeneratorConfig
from fastvideo.configs.sample import SamplingParam
from fastvideo.api.sampling_param import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines import ForwardBatch
+214
View File
@@ -0,0 +1,214 @@
# `fastvideo.eval`
In-process evaluation suite for video generations. Includes pixel
metrics (SSIM, PSNR, LPIPS), optical-flow comparisons, the full VBench
suite, Physics-IQ, and a VLM scorer (VideoScore-2) behind a single
registry-driven API.
## Install
| Use case | Install |
|---|---|
| Default (common, optical_flow, vbench-light, physics_iq, videoscore2) | `uv pip install -e .[eval]` |
| Just VBench (12 of 16 sub-metrics) | `uv pip install -e .[eval-vbench]` |
| Just Physics-IQ (covered by `[eval]`) | `uv pip install -e .[eval-physics-iq]` |
| Plus `vbench.scene` (AVoCaDO) | `uv pip install -e .[eval-full]` |
| Plus `vbench.{color, multiple_objects, object_class, spatial_relationship}` (GRiT) | `uv pip install -e .[eval-vbench]` then `uv pip install --no-build-isolation 'git+https://github.com/facebookresearch/detectron2.git'` |
To use VBench, also pull the upstream submodule:
```bash
git submodule update --init --recursive # fetches vbench + kernel deps
```
The submodule is a clean upstream pin. Compat with current
transformers/numpy/timm versions is applied at import time in
`fastvideo/eval/metrics/vbench/__init__.py` via attribute-level
monkey-patches; the submodule files are unchanged.
## Public API
```python
from fastvideo.eval import (
create_evaluator, # build a reusable Evaluator
evaluate, # one-shot helper
Evaluator, # the class itself
BaseMetric, MetricResult,
register, list_metrics, get_metric,
ensure_checkpoint, get_cache_dir,
)
ev = create_evaluator(metrics=["common.ssim", "vbench.aesthetic_quality"],
device="cuda")
scores = ev.evaluate(video=tensor, reference=ref, fps=8.0)
```
`evaluate` accepts either a pre-loaded `(T, C, H, W)` tensor or a path
string for `video` and `reference`. Paths are decoded inside the worker
that picks up the sample, so peak memory stays bounded by `num_gpus`
when scoring large batches.
### CLI
```bash
fastvideo eval list # list registered metrics
fastvideo eval list --group vbench # filter by group
fastvideo eval run --videos clip.mp4 \
--metrics vbench.aesthetic_quality \
--output scores.json
fastvideo eval run --videos generated/*.mp4 \
--reference reference/ \
--metrics common.ssim,common.lpips
```
### Generate-then-score example
`examples/inference/eval/eval_ltx2_vbench.py` runs an LTX2 prompt
through `VideoGenerator` and scores the resulting mp4 with
`vbench.aesthetic_quality` and `vbench.subject_consistency`. Use it as
a template for end-to-end "generate then score" pipelines.
## Layout
```
fastvideo/
├── eval/
│ ├── api.py, evaluator.py, registry.py, models.py, ...
│ ├── io/ # video loading helpers
│ ├── datasets/ # prompt corpora (vbench, physics_iq)
│ └── metrics/
│ ├── base.py # BaseMetric + @register contract
│ ├── common/ # SSIM, PSNR, LPIPS
│ ├── optical_flow/ # gt_optical_flow, synthetic_optical_flow
│ ├── videoscore2/ # VideoScore-2 (Qwen2.5-VL)
│ ├── physics_iq/ # PhysicsIQ + sub-metrics
│ └── vbench/ # adapter: sys.path bootstrap + shims
│ ├── __init__.py
│ └── <16 sub-metric pkgs>
└── third_party/
└── eval/
└── vbench/ # git submodule (Vchitect/VBench)
```
### Prompt datasets
```python
from fastvideo.eval.datasets import get_dataset, list_datasets
list_datasets() # ['physics_iq', 'vbench']
ds = get_dataset("physics_iq", limit=4) # auto-fetches assets on first miss
for row in ds:
# row contains 'prompt', 'reference', 'reference_take2', and
# metric-specific aux fields. Drop straight into Evaluator.evaluate(**row).
...
```
The Physics-IQ manifest CSV is vendored at
`fastvideo/eval/metrics/physics_iq/_vendored/descriptions.csv`.
Per-scenario videos, masks, and switch-frames auto-fetch on first use
into `${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq/`. For air-gapped
runs, pass `auto_download=False` or `dataset_root=` a pre-downloaded
copy. Set `FASTVIDEO_PHYSICS_IQ_BUCKET_URL` to redirect the fetch to
an internal mirror.
## Adding a new metric
The full porting guide is at
[`docs/contributing/eval-metrics.md`](../../docs/contributing/eval-metrics.md).
Summary below.
### Native metric (no submodule)
```python
# fastvideo/eval/metrics/common/<your_metric>/metric.py
from fastvideo.eval.metrics.base import BaseMetric
from fastvideo.eval.registry import register
from fastvideo.eval.types import MetricResult
@register("common.your_metric")
class YourMetric(BaseMetric):
name = "common.your_metric"
requires_reference = True
needs_gpu = False
dependencies: list[str] = [] # e.g. ["pyiqa"] if relevant
def compute(self, sample) -> list[MetricResult]:
...
```
The metric is auto-discovered by `fastvideo/eval/metrics/__init__.py`,
which walks all non-underscore subdirectories and imports their
`metric` module.
### Wrapping upstream code via a submodule
See `fastvideo/eval/metrics/vbench/` for a worked example. The
contract is:
1. Upstream lives as a git submodule under
`fastvideo/third_party/eval/<bench>/`, pinned to a SHA in repo-root
`.gitmodules`.
2. The metric package's `__init__.py`
(`fastvideo/eval/metrics/<bench>/__init__.py`) inserts that
submodule path on `sys.path` and installs any compat shims for
modern torch/transformers/numpy. Do not modify upstream files on
disk.
3. Per-sub-metric `metric.py` files use `@register("<bench>.<name>")`.
Patches live as Python in the metric's `__init__.py` so they are
grep-able and reviewable.
## Caches
Eval cache root: `${FASTVIDEO_CACHE_ROOT}/eval/`, default
`~/.cache/fastvideo/eval/`. Override with `FASTVIDEO_EVAL_CACHE`.
```
${FASTVIDEO_CACHE_ROOT}/eval/
├── models/ # URL-fetched checkpoints (LAION head, AMT, GRiT)
├── torch/ # redirected TORCH_HOME (DINO via torch.hub, lpips)
├── clip/ # passed as download_root= to clip.load callsites
└── datasets/ # auto-fetched dataset assets, one subdir per benchmark
# (e.g. datasets/physics_iq/{split-videos,switch-frames,...})
```
HF-hosted models stay in HF's default cache
(`~/.cache/huggingface/hub/`) so they dedupe with other ML projects on
the same host.
### Convention for new metrics
If your metric wraps a third-party loader that has its own cache
directory, route it through `get_cache_dir()` so users get one knob
to redirect everything.
```python
# CLIP: pass download_root explicitly
import clip
from fastvideo.eval.models import get_cache_dir
model, _ = clip.load("ViT-B/32", device=device,
download_root=str(get_cache_dir() / "clip"))
# torch.hub is already redirected by fastvideo.eval.__init__ via
# TORCH_HOME; no per-callsite work needed.
# transformers / huggingface_hub: leave alone. HF's default cache is
# shared with other tools.
```
For libraries that do not honour any env var or kwarg (pyiqa, funasr),
their cache lands in the library's own dir. Document the exception in
the metric's docstring if it matters.
## Out of scope (follow-up PRs)
- **MIND** metrics. Depend on a separate `vipe` upstream submodule.
- **VBench-2.0**. Sibling vbench2 package; needs its own port.
- **FVD as a registered metric**. Currently still at `benchmarks/fvd/`.
FVD is a set-vs-set distribution distance and does not fit the
per-sample `BaseMetric.compute` API without a stateful accumulator;
conversion is a designed follow-up.
- **Training-time eval callback** (`EvalCallback`) and the
`RolloutEvaluator` helper.
+45
View File
@@ -0,0 +1,45 @@
from fastvideo.eval.models import ensure_checkpoint, get_cache_dir
def _redirect_third_party_caches() -> None:
"""Point libraries that respect env vars at the eval cache root.
Run before metric modules import torch.hub (and friends) so the
redirect actually takes effect. We intentionally leave ``HF_HOME``
alone — HF's default cache (``~/.cache/huggingface/hub``) is widely
shared with other ML projects, and isolating it for eval would force
users to re-download already-cached transformers weights.
Other libraries with non-standard caches (CLIP, pyiqa) don't honour
env vars at all; their callsites in metric.py files pass
``download_root=str(get_cache_dir() / "<library>")`` directly.
"""
import os
root = get_cache_dir()
os.environ.setdefault("TORCH_HOME", str(root / "torch"))
_redirect_third_party_caches()
from fastvideo.eval.types import MetricResult, Video # noqa: E402
from fastvideo.eval.metrics.base import BaseMetric # noqa: E402
from fastvideo.eval.registry import register, list_metrics, get_metric # noqa: E402
from fastvideo.eval.api import evaluate # noqa: E402
from fastvideo.eval.evaluator import Evaluator, create_evaluator # noqa: E402
# Trigger metric auto-discovery
import fastvideo.eval.metrics # noqa: F401, E402
__all__ = [
"evaluate",
"Evaluator",
"create_evaluator",
"MetricResult",
"Video",
"BaseMetric",
"register",
"list_metrics",
"get_metric",
"ensure_checkpoint",
"get_cache_dir",
]
+36
View File
@@ -0,0 +1,36 @@
from __future__ import annotations
from pathlib import Path
import torch
from fastvideo.eval.evaluator import create_evaluator
from fastvideo.eval.types import MetricResult
def evaluate(
generated: torch.Tensor | str | Path,
reference: torch.Tensor | str | Path | None = None,
metrics: list[str] | str = "all",
device: str = "cuda",
**kwargs,
) -> dict[str, MetricResult] | list[dict[str, MetricResult]]:
"""One-shot evaluation. For repeated use, prefer :func:`create_evaluator`.
Parameters
----------
generated : Tensor | str | Path
Generated video. Either a pre-loaded ``(T, C, H, W)`` tensor or a
path to an mp4/avi/etc. — paths are decoded by the worker.
reference : Tensor | str | Path | None
Reference video (same accepted shapes as *generated*).
metrics : list[str] | str
Metric names, or ``"all"``.
device : str
PyTorch device string.
"""
ev = create_evaluator(metrics=metrics, device=device)
kw: dict = {"video": generated, **kwargs}
if reference is not None:
kw["reference"] = reference
return ev.evaluate(**kw)
+51
View File
@@ -0,0 +1,51 @@
"""Prompt-corpus datasets for end-to-end benchmark evaluation.
Public API mirrors :mod:`fastvideo.eval` (metrics side):
from fastvideo.eval.datasets import (
PromptDataset, Sample,
register_dataset, get_dataset, list_datasets,
)
A dataset is an iterable of plain dicts (one per sample). Built-in
datasets self-register at import time. To add one, drop a module into
this package that subclasses :class:`PromptDataset` and decorates with
``@register_dataset("name")`` — auto-discovery picks it up.
"""
from fastvideo.eval.datasets.base import (BasePromptDataset, PromptDataset, Sample)
from fastvideo.eval.datasets.registry import (get_dataset, list_datasets, register_dataset)
def _autodiscover() -> None:
"""Import every non-underscore .py module / subpackage in this package
so the ``@register_dataset`` decorators fire."""
import importlib
import os
for entry in os.listdir(os.path.dirname(__file__)):
if entry.startswith("_") or entry.startswith("."):
continue
if entry in {"base.py", "registry.py"}:
continue
if entry.endswith(".py"):
importlib.import_module(f"{__name__}.{entry[:-3]}")
elif os.path.isdir(os.path.join(os.path.dirname(__file__), entry)) \
and os.path.exists(os.path.join(
os.path.dirname(__file__), entry, "__init__.py")):
importlib.import_module(f"{__name__}.{entry}")
_autodiscover()
# Re-export the canonical class for typed imports.
from fastvideo.eval.datasets.vbench import VBenchPromptDataset # noqa: E402
__all__ = [
"PromptDataset",
"BasePromptDataset",
"Sample",
"register_dataset",
"get_dataset",
"list_datasets",
"VBenchPromptDataset",
]
+81
View File
@@ -0,0 +1,81 @@
"""Prompt-corpus datasets.
A :class:`PromptDataset` is an iterable of *sample dicts* describing the
prompts and conditions for a benchmark. Each sample is a plain dict —
no dataclass, no schema enforcement — that flows directly into both
generation (``VideoGenerator.generate_video(**sample)``) and scoring
(``Evaluator.evaluate(**eval_kwargs)``). The runner picks well-known
keys (``prompt``, ``n_samples``, ``dimensions``, ``auxiliary_info``,
...) and passes the rest through.
This matches the surrounding FastVideo style:
* :class:`fastvideo.dataset.validation_dataset.ValidationDataset` yields dicts.
* :meth:`fastvideo.VideoGenerator.generate_video` consumes ``**kwargs``.
* :meth:`fastvideo.eval.Evaluator.evaluate` consumes ``**kwargs``.
To add a new benchmark:
1. Subclass :class:`PromptDataset`, populate ``self._rows`` with dicts in
``__init__``.
2. Decorate with ``@register_dataset("my_bench")``.
Convention for ``auxiliary_info``: a *flat* dict of metric-keyed values
(e.g. ``{"color": "red"}``). Benchmarks with nested aux schemas (VBench's
``{dim: {key: val}}``) flatten at load time so every consumer sees the
same shape.
"""
from __future__ import annotations
from typing import TypedDict
from collections.abc import Iterator
class Sample(TypedDict, total=False):
"""Documented schema for a row yielded by :class:`PromptDataset`.
Only ``prompt`` is required. Extra keys beyond these are forwarded to
the runner's eval-kwargs builder verbatim, so action-conditioned or
audio-bearing benchmarks can add their own fields without changing
the base class.
"""
prompt: str
n_samples: int
dimensions: list[str]
auxiliary_info: dict
image_path: str
reference_video: str
class PromptDataset:
"""Iterable corpus of sample dicts. Subclasses populate ``self._rows``."""
name: str = ""
description: str = ""
supports_dimensions: bool = False
requires_reference_image: bool = False
requires_reference_video: bool = False
def __init__(self) -> None:
self._rows: list[dict] = []
def __iter__(self) -> Iterator[dict]:
return iter(self._rows)
def __len__(self) -> int:
return len(self._rows)
def __getitem__(self, i: int) -> dict:
return self._rows[i]
def by_dimension(self) -> dict[str, list[dict]]:
"""Group samples by dimension. A multi-dim sample appears under each."""
out: dict[str, list[dict]] = {}
for s in self._rows:
for d in s.get("dimensions", ()):
out.setdefault(d, []).append(s)
return out
# Back-compat alias for callers still importing the old class name.
BasePromptDataset = PromptDataset
+444
View File
@@ -0,0 +1,444 @@
"""Physics-IQ benchmark prompt corpus.
Yields one sample dict per take-1 scenario, paired with its take-2
reference and both takes' real motion masks. Each row drops straight
into :meth:`fastvideo.eval.Evaluator.evaluate` for the ``physics_iq``
metric:
{
"prompt": <description>,
"reference": "<take-1 mp4>",
"reference_take2": "<take-2 mp4>",
"reference_mask": "<take-1 mask mp4>",
"reference_take2_mask": "<take-2 mask mp4>",
"scenario": <scenario_id>,
"view": <camera view>,
"auxiliary_info": { ... metadata ... },
}
Self-contained dataset: the manifest CSV is vendored under
``fastvideo/eval/metrics/physics_iq/_vendored/descriptions.csv``;
per-scenario videos/masks/switch-frames auto-fetch on first use from the public
DeepMind bucket into ``${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq/``.
Pass ``auto_download=False`` (or ``dataset_root=`` pointing at a
pre-downloaded copy) to opt out of network fetches.
"""
from __future__ import annotations
import csv
import os
from dataclasses import dataclass
from pathlib import Path
from urllib.request import urlretrieve
import cv2
import numpy as np
from fastvideo.eval.datasets.base import PromptDataset
from fastvideo.eval.datasets.registry import register_dataset
from fastvideo.eval.models import get_cache_dir
VIEWS = ("perspective-left", "perspective-center", "perspective-right")
TAKE1_TOKEN = "take-1"
TAKE2_TOKEN = "take-2"
# FPS the dataset rows should resolve to. Source release is recorded at
# 30 FPS; if a different value is requested the loader transcodes once
# into a per-repo-root cache directory.
_DEFAULT_FPS = 30
_DEFAULT_DURATION_SECONDS = 5
# Vendored manifest under ``fastvideo/eval/metrics/physics_iq/_vendored/``
# is the same file shipped by upstream's git repo. The ``_vendored/``
# subdir is the project-wide convention for upstream-provenance files
# (matches the ``_``-prefixed auto-discovery skip and a single
# codespell skip glob).
_VENDORED_DESCRIPTIONS_CSV = (Path(__file__).resolve().parent.parent / "metrics" / "physics_iq" / "_vendored" /
"descriptions.csv")
# Public DeepMind bucket; HTTPS-readable, no auth. Override via
# ``FASTVIDEO_PHYSICS_IQ_BUCKET_URL`` (e.g. for an internal mirror).
_DEFAULT_BUCKET_URL = "https://storage.googleapis.com/physics-iq-benchmark"
def _bucket_url() -> str:
return os.environ.get("FASTVIDEO_PHYSICS_IQ_BUCKET_URL", _DEFAULT_BUCKET_URL)
def _default_dataset_root() -> Path:
"""Sibling to ``models/torch/clip/`` under the eval cache root."""
return get_cache_dir() / "datasets" / "physics_iq"
@dataclass(frozen=True)
class PhysicsIQScenario:
"""One row of the Physics-IQ manifest, fully resolved on disk."""
scenario_id: str
view: str
scenario_name: str
take1_video_path: str
take2_video_path: str
switch_frame_path: str
caption: str
expected_gen_filename: str
generated_video_path: str | None = None
take1_mask_path: str | None = None
take2_mask_path: str | None = None
@register_dataset("physics_iq")
class PhysicsIQPromptDataset(PromptDataset):
"""Physics-IQ benchmark prompt corpus.
Self-contained: ``get_dataset("physics_iq")`` works with no kwargs.
The manifest CSV is vendored next to the metric, and per-scenario
assets auto-fetch on first miss from the public bucket into
``${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq/``.
Args:
dataset_root: path to a pre-downloaded copy of the Physics-IQ
release. Defaults to ``${FASTVIDEO_EVAL_CACHE}/datasets/physics_iq``;
override only if you already have a local mirror.
fps: target frame rate. The release ships at 30 FPS; other rates
transcode once on first access into ``<root>/.physics_iq_cache/``.
limit: optional truncation for quick smoke runs. Apply this kwarg
(not a post-construction slice) so we only fetch the assets
for the scenarios actually requested.
generated_dir: optional directory of pre-generated videos —
attaches each manifest row's expected output path to the
sample dict under ``auxiliary_info["generated_video_path"]``.
auto_download: when True (the default), missing testing videos,
masks, and switch frames are fetched from the public bucket
into ``dataset_root``. Set False for air-gapped runs; the
loader will then raise ``FileNotFoundError`` on miss.
"""
description = ("Physics-IQ benchmark, 396 take-1 scenarios across 66 unique physics "
"setups × 3 perspective views, each paired with a take-2 reference.")
requires_reference_video = True
def __init__(
self,
dataset_root: str | Path | None = None,
*,
fps: int = _DEFAULT_FPS,
limit: int | None = None,
generated_dir: str | Path | None = None,
auto_download: bool = True,
) -> None:
super().__init__()
repo_root = Path(dataset_root or _default_dataset_root()).expanduser().resolve()
self.repo_root = repo_root
self.dataset_dir = _resolve_dataset_dir(repo_root)
self.descriptions_path = _resolve_descriptions_path(repo_root, self.dataset_dir)
self.cache_dir = repo_root / ".physics_iq_cache"
self.fps = fps
self.auto_download = auto_download
self.bucket_url = _bucket_url()
scenarios = self._iter_scenarios(
fps=fps,
generated_dir=generated_dir,
limit=limit,
)
self._rows = [_scenario_to_row(s) for s in scenarios]
def _iter_scenarios(
self,
*,
fps: int,
generated_dir: str | Path | None,
limit: int | None,
) -> list[PhysicsIQScenario]:
with self.descriptions_path.open("r", newline="") as handle:
rows = list(csv.DictReader(handle))
take2_by_suffix = {_scenario_suffix(row["scenario"]): row for row in rows if TAKE2_TOKEN in row["scenario"]}
take1_rows = [row for row in rows if TAKE1_TOKEN in row["scenario"]]
if limit is not None:
take1_rows = take1_rows[:limit]
generated_dir_path = (Path(generated_dir).expanduser().resolve() if generated_dir else None)
scenarios: list[PhysicsIQScenario] = []
for row in take1_rows:
scenario_filename = row["scenario"]
scenario_id, view, _, scenario_name = _parse_scenario_filename(scenario_filename)
take2_row = take2_by_suffix.get(_scenario_suffix(scenario_filename))
if take2_row is None:
raise FileNotFoundError(f"Could not find take-2 row matching {scenario_filename}")
take2_id, _, _, _ = _parse_scenario_filename(take2_row["scenario"])
take1_video_path = self._resolve_testing_video_path(
scenario_id=scenario_id,
view=view,
take=TAKE1_TOKEN,
scenario_name=scenario_name,
fps=fps,
)
take2_video_path = self._resolve_testing_video_path(
scenario_id=take2_id,
view=view,
take=TAKE2_TOKEN,
scenario_name=scenario_name,
fps=fps,
)
switch_frame_path = self._resolve_switch_frame_path(
scenario_id=scenario_id,
view=view,
scenario_name=scenario_name,
)
take1_mask_path = self._resolve_real_mask_path(
scenario_id=scenario_id,
view=view,
take=TAKE1_TOKEN,
scenario_name=scenario_name,
fps=fps,
)
take2_mask_path = self._resolve_real_mask_path(
scenario_id=take2_id,
view=view,
take=TAKE2_TOKEN,
scenario_name=scenario_name,
fps=fps,
)
generated_video_path = (str(generated_dir_path /
row["generated_video_name"]) if generated_dir_path is not None else None)
scenarios.append(
PhysicsIQScenario(
scenario_id=scenario_id,
view=view,
scenario_name=scenario_name,
take1_video_path=str(take1_video_path),
take2_video_path=str(take2_video_path),
switch_frame_path=str(switch_frame_path),
caption=row["description"],
expected_gen_filename=row["generated_video_name"],
generated_video_path=generated_video_path,
take1_mask_path=str(take1_mask_path),
take2_mask_path=str(take2_mask_path),
))
return scenarios
def _resolve_testing_video_path(
self,
*,
scenario_id: str,
view: str,
take: str,
scenario_name: str,
fps: int,
) -> Path:
target_dir = self.dataset_dir / "split-videos" / "testing" / f"{fps}FPS"
target_name = (f"{scenario_id}_testing-videos_{fps}FPS_{view}_{take}_{scenario_name}.mp4")
target_path = target_dir / target_name
if target_path.exists():
return target_path
# 30-FPS source: either present locally or auto-fetchable.
source_name = (f"{scenario_id}_testing-videos_30FPS_{view}_{take}_{scenario_name}.mp4")
source_rel = f"split-videos/testing/30FPS/{source_name}"
source_path = self.dataset_dir / source_rel
self._ensure_remote_asset(source_rel, source_path)
if fps == _DEFAULT_FPS:
return source_path
# FPS-convert and cache so repeat runs are free.
cache_dir = self.cache_dir / "split-videos" / "testing" / f"{fps}FPS"
cache_dir.mkdir(parents=True, exist_ok=True)
cached_path = cache_dir / target_name
if not cached_path.exists():
_convert_video_fps(source_path, cached_path, fps_new=fps)
return cached_path
def _resolve_switch_frame_path(
self,
*,
scenario_id: str,
view: str,
scenario_name: str,
) -> Path:
rel = (f"switch-frames/{scenario_id}_switch-frames_anyFPS_{view}_{scenario_name}.jpg")
target_path = self.dataset_dir / rel
self._ensure_remote_asset(rel, target_path)
return target_path
def _resolve_real_mask_path(
self,
*,
scenario_id: str,
view: str,
take: str,
scenario_name: str,
fps: int,
) -> Path:
# Source release ships masks at 30 FPS only; non-30 rates are
# regenerated downstream from the (downsampled) real videos by
# the metric — see upstream ``run_physics_iq.py::ensure_binary_mask_structure``.
# We only auto-fetch 30 FPS here.
rel = (f"video-masks/real/30FPS/"
f"{scenario_id}_video-masks_30FPS_{view}_{take}_{scenario_name}.mp4")
target_path = self.dataset_dir / rel
self._ensure_remote_asset(rel, target_path)
if fps == _DEFAULT_FPS:
return target_path
# Caller asked for a non-30 rate; metric layer handles the
# regeneration. Return the canonical 30 FPS path so the metric
# always sees a valid mp4 it can transcode.
return target_path
def _ensure_remote_asset(self, rel_path: str, target_path: Path) -> Path:
"""Download ``<bucket_url>/<rel_path>`` into *target_path* on miss.
Atomic via a sibling ``.part`` file; safe under concurrent runs
because the final ``rename`` is atomic on POSIX. Raises
``FileNotFoundError`` if the file is missing and ``auto_download``
is False.
"""
if target_path.exists():
return target_path
if not self.auto_download:
raise FileNotFoundError(f"Physics-IQ asset missing: {target_path}. "
"Set auto_download=True or pass dataset_root= a pre-downloaded copy.")
target_path.parent.mkdir(parents=True, exist_ok=True)
url = f"{self.bucket_url}/{rel_path.lstrip('/')}"
tmp_path = target_path.with_suffix(target_path.suffix + ".part")
try:
urlretrieve(url, tmp_path)
except Exception as exc:
if tmp_path.exists():
tmp_path.unlink()
raise FileNotFoundError(f"Failed to fetch Physics-IQ asset {url} -> {target_path}: "
f"{type(exc).__name__}: {exc}") from exc
tmp_path.rename(target_path)
return target_path
# ---------------------------------------------------------------------------
# Internal helpers
# ---------------------------------------------------------------------------
def _resolve_dataset_dir(repo_root: Path) -> Path:
nested = repo_root / "physics-IQ-benchmark"
if nested.exists():
return nested
return repo_root
def _resolve_descriptions_path(repo_root: Path, dataset_dir: Path) -> Path:
"""Prefer a co-located CSV under the user's dataset_root; fall back
to the copy vendored in this repo so ``get_dataset("physics_iq")``
works without external setup.
"""
candidates = (
repo_root / "descriptions" / "descriptions.csv",
dataset_dir / "descriptions" / "descriptions.csv",
)
for path in candidates:
if path.exists():
return path
if _VENDORED_DESCRIPTIONS_CSV.is_file():
return _VENDORED_DESCRIPTIONS_CSV
raise FileNotFoundError("Could not locate Physics-IQ descriptions/descriptions.csv "
f"(checked {[str(c) for c in candidates]} and vendored "
f"{_VENDORED_DESCRIPTIONS_CSV})")
def _parse_scenario_filename(filename: str) -> tuple[str, str, str, str]:
stem = Path(filename).name
if stem.endswith(".mp4"):
stem = stem[:-4]
parts = stem.split("_")
if len(parts) < 4:
raise ValueError(f"Unexpected Physics-IQ filename format: {filename}")
return parts[0], parts[1], parts[2], "_".join(parts[3:])
def _scenario_suffix(filename: str) -> str:
_, view, _, scenario_name = _parse_scenario_filename(filename)
return f"{view}_{scenario_name}"
def _scenario_to_row(scenario: PhysicsIQScenario) -> dict:
"""Flatten a :class:`PhysicsIQScenario` into the public sample-dict shape."""
aux: dict = {
"scenario_id": scenario.scenario_id,
"scenario_name": scenario.scenario_name,
"switch_frame_path": scenario.switch_frame_path,
"expected_gen_filename": scenario.expected_gen_filename,
}
if scenario.generated_video_path is not None:
aux["generated_video_path"] = scenario.generated_video_path
row: dict = {
"prompt": scenario.caption,
"reference": scenario.take1_video_path,
"reference_take2": scenario.take2_video_path,
"scenario": scenario.scenario_id,
"view": scenario.view,
"auxiliary_info": aux,
}
if scenario.take1_mask_path is not None:
row["reference_mask"] = scenario.take1_mask_path
if scenario.take2_mask_path is not None:
row["reference_take2_mask"] = scenario.take2_mask_path
return row
def _convert_video_fps(input_path: str | Path, output_path: str | Path, *, fps_new: int) -> None:
"""Trim *input_path* to ``_DEFAULT_DURATION_SECONDS`` and re-encode at
*fps_new*, writing the result to *output_path*. Used to materialize
Physics-IQ's 30-FPS source release at user-requested rates.
"""
input_path = Path(input_path)
output_path = Path(output_path)
cap = cv2.VideoCapture(str(input_path))
if not cap.isOpened():
raise FileNotFoundError(f"Could not open video for FPS conversion: {input_path}")
fps_original = cap.get(cv2.CAP_PROP_FPS)
frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
width = int(cap.get(cv2.CAP_PROP_FRAME_WIDTH))
height = int(cap.get(cv2.CAP_PROP_FRAME_HEIGHT))
duration = frame_count / fps_original if fps_original else 0.0
width, height = width - width % 2, height - height % 2
subclip_duration = min(_DEFAULT_DURATION_SECONDS, duration)
frames: list[np.ndarray] = []
for _ in range(int(subclip_duration * fps_original)):
ret, frame = cap.read()
if not ret:
break
frames.append(frame)
cap.release()
if not frames:
raise ValueError(f"No frames decoded from {input_path}")
frame_count_new = int(subclip_duration * fps_new)
output_path.parent.mkdir(parents=True, exist_ok=True)
writer = cv2.VideoWriter(
str(output_path),
cv2.VideoWriter_fourcc(*"avc1"),
fps_new,
(width, height),
)
if frame_count_new <= 1:
writer.write(frames[0])
writer.release()
return
frame_count_original = len(frames)
for j in range(frame_count_new):
alpha = j * (frame_count_original - 1) / (frame_count_new - 1)
idx = int(alpha)
alpha -= idx
f1 = frames[idx].astype(np.float32)
f2 = frames[min(idx + 1, frame_count_original - 1)].astype(np.float32)
writer.write(((1.0 - alpha) * f1 + alpha * f2).astype(np.uint8))
writer.release()
+41
View File
@@ -0,0 +1,41 @@
"""Registry for prompt-corpus datasets, mirroring :mod:`fastvideo.eval.registry`."""
from __future__ import annotations
from typing import Any, TYPE_CHECKING
if TYPE_CHECKING:
from fastvideo.eval.datasets.base import BasePromptDataset
_REGISTRY: dict[str, type[BasePromptDataset]] = {}
def register_dataset(name: str):
"""Decorator to register a prompt-dataset class.
Usage::
@register_dataset("vbench")
class VBenchPromptDataset(BasePromptDataset):
...
"""
def wrapper(cls):
cls.name = name
_REGISTRY[name] = cls
return cls
return wrapper
def get_dataset(name: str, **kwargs: Any) -> BasePromptDataset:
"""Instantiate a registered dataset by name."""
cls = _REGISTRY.get(name)
if cls is None:
available = ", ".join(sorted(_REGISTRY.keys()))
raise KeyError(f"Unknown dataset '{name}'. Available: {available}")
return cls(**kwargs)
def list_datasets() -> list[str]:
"""Return sorted list of all registered dataset names."""
return sorted(_REGISTRY.keys())
+126
View File
@@ -0,0 +1,126 @@
"""VBench prompt corpus.
Single source of truth: upstream's ``VBench_full_info.json`` (946 entries,
each with ``prompt_en``, a ``dimension`` list, optional ``auxiliary_info``
keyed by dimension).
"""
from __future__ import annotations
import json
import os
from pathlib import Path
from fastvideo.eval.datasets.base import PromptDataset
from fastvideo.eval.datasets.registry import register_dataset
# VBench's official sampling protocol: 5 generations per prompt, except
# temporal_flickering which requires 25 (averaging over 5 is too noisy
# for a high-frequency-noise metric). See upstream prompts/README.md.
TEMPORAL_FLICKERING_SAMPLES = 25
DEFAULT_SAMPLES = 5
_FULL_INFO_REL = "fastvideo/third_party/eval/vbench/vbench/VBench_full_info.json"
def _locate_full_info() -> Path:
env = os.environ.get("VBENCH_FULL_INFO_JSON")
if env:
p = Path(env)
if p.is_file():
return p
raise FileNotFoundError(f"VBENCH_FULL_INFO_JSON={env} does not point at a file")
here = Path(__file__).resolve()
for ancestor in here.parents:
candidate = ancestor / _FULL_INFO_REL
if candidate.is_file():
return candidate
if (ancestor / ".git").exists():
break
raise FileNotFoundError("Could not locate VBench_full_info.json. Initialize the upstream "
"submodule (`git submodule update --init "
"fastvideo/third_party/eval/vbench`) or set VBENCH_FULL_INFO_JSON.")
@register_dataset("vbench")
class VBenchPromptDataset(PromptDataset):
"""VBench prompts filtered by evaluation dimension.
Args:
dimensions: List of dimension names, or ``"all"``. Unknown
dimensions raise ``ValueError``.
full_info_path: Optional override for ``VBench_full_info.json``;
defaults to autodetection.
A prompt that belongs to several requested dimensions is yielded once;
its ``dimensions`` list carries all matches so the scorer can route.
"""
description = ("VBench (Vchitect) prompt corpus, 946 prompts across 16 "
"evaluation dimensions.")
supports_dimensions = True
def __init__(
self,
dimensions: list[str] | str = "all",
full_info_path: str | Path | None = None,
) -> None:
super().__init__()
path = Path(full_info_path) if full_info_path else _locate_full_info()
with path.open() as f:
entries = json.load(f)
all_dims = sorted({d for e in entries for d in e["dimension"]})
if dimensions == "all":
self.dimensions: list[str] = all_dims
else:
unknown = set(dimensions) - set(all_dims)
if unknown:
raise ValueError(f"Unknown VBench dimensions: {sorted(unknown)}. "
f"Available: {all_dims}")
self.dimensions = list(dimensions)
wanted = set(self.dimensions)
for entry in entries:
relevant = [d for d in entry["dimension"] if d in wanted]
if not relevant:
continue
n = (TEMPORAL_FLICKERING_SAMPLES if "temporal_flickering" in relevant else DEFAULT_SAMPLES)
# Strip the outer {dim_name: ...} wrapper from upstream's aux
# schema so every metric reads its inputs from a flat dict.
#
# This unwraps exactly one level — the dimension key. Whatever
# shape lives inside is the metric's contract:
#
# color: {"color": {"color": "red"}}
# → flat: {"color": "red"} (scalar)
#
# object_class: {"object_class": {"object": "person"}}
# → flat: {"object": "person"} (scalar)
#
# multiple_objects: {"multiple_objects": {"object": "a and b"}}
# → flat: {"object": "a and b"} (scalar)
#
# spatial_relationship: {"spatial_relationship":
# {"spatial_relationship":
# {"object_a": ..., "object_b": ...,
# "relationship": ...}}}
# → flat: {"spatial_relationship": {object_a,object_b,relationship}}
#
# Note the spatial_relationship case keeps a nested inner dict
# by design — upstream double-wraps it, the SpatialRelationship
# metric reads ``aux["spatial_relationship"]`` expecting that
# inner dict. Don't "simplify" the wrapping away.
raw_aux = entry.get("auxiliary_info") or {}
flat_aux: dict = {}
for v in raw_aux.values():
if isinstance(v, dict):
flat_aux.update(v)
self._rows.append({
"prompt": entry["prompt_en"],
"n_samples": n,
"dimensions": relevant,
"auxiliary_info": flat_aux,
})
self.full_info_path = path

Some files were not shown because too many files have changed in this diff Show More