Compare commits

...
Author SHA1 Message Date
Will LinandClaude Fable 5 9ba740e125 [docs] inference optimizations: document the regional fullgraph compile knob
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:55:03 +00:00
Will LinandClaude Fable 5 34346779d1 [perf] inference: regional fullgraph torch.compile of the DiT blocks (port of #1718 to the inference path)
Extend the training-side regional-compile port (internal 863e87342, GO
verdict vsa_gate/compile_ab/VERDICT.md) to inference so every H3 run can
compile the 52 transformer blocks. Opt-in and off by default:

- New FastVideoArgs.inference_torch_compile (CLI --inference-torch-compile),
  env FASTVIDEO_INFERENCE_TORCH_COMPILE=1 folded in __post_init__ (the
  attention_backend parse-once pattern), reachable through
  PipelineSelection.experimental {"inference_torch_compile": true} exactly
  like the VSA_sparsity / VSA_tile_size knobs.
- maybe_load_fsdp_model applies the compile right after the transformer
  loads: per-_compile_conditions block, fullgraph=True + inductor
  options.emulate_precision_casts injected, no user kwargs needed (the
  compile-A/B verdict recipe). prepare_for_compile still runs first, so the
  H3 fusion-inertness warning of #1735 is preserved.
- _regional_compile_unsupported_reason ported from the training loader:
  VSA / VSA-H3 backends and the FASTVIDEO_DISABLE_ATTENTION_COMPILE=1
  escape hatch degrade the transformer to eager with one warning; FLASH_ATTN
  on flash-attn 3 is rejected with an actionable message. FA4/FA2/SDPA
  compile through their existing custom-op boundaries (forward-only is
  enough at inference; the training port's backward custom ops are not
  needed under no_grad).
- FASTVIDEO_DISABLE_ATTENTION_COMPILE default flipped to 0 (attention is
  traced by default), matching the training port and upstream #1718 —
  eager runs are unaffected (torch.compiler.disable is inert outside
  dynamo).
- enable_torch_compile + inference_torch_compile together skip the
  pipeline-level DiT compile (the loader already owns those forwards).
- examples: --inference-torch-compile on basic_minimax_h3_t2v.py (dense
  target) and basic_fasth3.py (exercises the VSA degrade path).
- CPU-safe contract tests for the guard and the kwargs injection.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
2026-08-21 17:46:26 +00:00
10 changed files with 334 additions and 4 deletions
+25
View File
@@ -337,6 +337,31 @@ Only DiT submodules that declare `_compile_conditions` are compiled
(most shipped models). The text encoder and VAE are not compiled by this
flag.
### Regional fullgraph compile (experimental)
`inference_torch_compile` is a stricter, kwargs-free variant that ports the
training-side regional compile of
[#1718](https://github.com/hao-ai-lab/FastVideo/pull/1718) to inference: the
loader wraps each `_compile_conditions` block in
`torch.compile(fullgraph=True)` with inductor
`options={"emulate_precision_casts": True}` right after the transformer
loads. Attention backends that cannot be traced end-to-end (VSA,
FLASH_ATTN on flash-attn 3, or the `FASTVIDEO_DISABLE_ATTENTION_COMPILE=1`
escape hatch) degrade the transformer to eager with one warning instead of
failing mid-denoise.
```python
generator = VideoGenerator.from_pretrained(
"MiniMaxAI/MiniMax-H3",
inference_torch_compile=True, # or FASTVIDEO_INFERENCE_TORCH_COMPILE=1
)
```
Do not combine it with `torch_compile_kwargs['mode']` (the loader injects
inductor options, and torch.compile forbids mode+options); it is
independent of `enable_torch_compile`, and when both are set the regional
compile wins for the DiT.
### What to expect
| Config | Effect |
+8
View File
@@ -85,6 +85,12 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--inference-torch-compile",
action="store_true",
help="regional fullgraph torch.compile of each DiT block after load. NOTE: this "
"script always runs the VSA-H3 backend, which is not fullgraph-traceable — the "
"loader logs one warning and keeps the transformer eager. The flag is exposed "
"here to exercise exactly that guard")
parser.add_argument("--repeats",
type=int,
default=1,
@@ -119,6 +125,8 @@ def main() -> None:
}
if args.vsa_sparsity > 0.0:
experimental["VSA_sparsity"] = args.vsa_sparsity
if args.inference_torch_compile:
experimental["inference_torch_compile"] = True
generator = VideoGenerator.from_config(
GeneratorConfig(
@@ -15,6 +15,7 @@ from fastvideo.api import (
OffloadConfig,
OutputConfig,
ParallelismConfig,
PipelineSelection,
SamplingConfig,
)
@@ -41,6 +42,12 @@ def parse_args() -> argparse.Namespace:
parser.add_argument("--compile-mode",
default=None,
help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs')
parser.add_argument("--inference-torch-compile",
action="store_true",
help="regional fullgraph torch.compile of each DiT block after load (the #1718 "
"training-port semantics: no kwargs; fullgraph + emulate_precision_casts injected). "
"First generation pays the inductor JIT (~1-2 min); use --repeats >= 2 and time "
"the last repeat. FASTVIDEO_INFERENCE_TORCH_COMPILE=1 is equivalent")
parser.add_argument("--repeats",
type=int,
default=1,
@@ -54,9 +61,16 @@ def main() -> None:
output_dir = Path(args.output)
output_dir.mkdir(parents=True, exist_ok=True)
# Boot-time run configuration folded into FastVideoArgs (the same
# experimental-dict route basic_fasth3.py uses for the VSA knobs).
experimental: dict[str, object] = {}
if args.inference_torch_compile:
experimental["inference_torch_compile"] = True
generator = VideoGenerator.from_config(
GeneratorConfig(
model_path=args.model_path,
pipeline=PipelineSelection(experimental=experimental),
engine=EngineConfig(
num_gpus=args.num_gpus,
use_fsdp_inference=args.num_gpus > 1,
+4 -4
View File
@@ -18,13 +18,13 @@ from fastvideo.layers.rotary_embedding import _apply_rotary_emb
def _attention_compile_disabled() -> bool:
"""Whether to keep attention ``forward`` out of the torch.compile graph.
Defaults to ``True`` (the historical behavior: attention runs eager via
``torch.compiler.disable``). Set ``FASTVIDEO_DISABLE_ATTENTION_COMPILE=0``
to let attention be traced/compiled into the surrounding graph.
Attention backends expose traceable custom-op boundaries, so compilation
is enabled by default. Set ``FASTVIDEO_DISABLE_ATTENTION_COMPILE=1`` for
an explicit eager escape hatch when debugging a backend.
"""
val = os.environ.get("FASTVIDEO_DISABLE_ATTENTION_COMPILE")
if val is None:
return True
return False
return val.strip().lower() not in ("0", "false", "no", "off", "")
+12
View File
@@ -21,6 +21,7 @@ if TYPE_CHECKING:
FASTVIDEO_TRACE_FUNCTION: int = 0
FASTVIDEO_ATTENTION_BACKEND: str | None = None
FASTVIDEO_FA4: bool = False
FASTVIDEO_INFERENCE_TORCH_COMPILE: bool = False
FASTVIDEO_MINIMAX_H3_FUSIONS: str = ""
FASTVIDEO_WORKER_MULTIPROC_METHOD: str = "spawn"
FASTVIDEO_TARGET_DEVICE: str = "cuda"
@@ -219,6 +220,17 @@ environment_variables: dict[str, Callable[[], Any]] = {
"FASTVIDEO_FA4":
lambda: os.getenv("FASTVIDEO_FA4", "0") != "0",
# If set (=1), enable regional (per-transformer-block) fullgraph
# torch.compile for the DiT at inference — the inference-side counterpart
# of the training regional-compile port of hao-ai-lab/FastVideo#1718.
# Equivalent to FastVideoArgs.inference_torch_compile=True (e.g. via
# PipelineSelection.experimental={"inference_torch_compile": True}). VSA
# and other non-fullgraph-traceable attention backends degrade to eager
# with one warning; see _regional_compile_unsupported_reason in
# fastvideo/models/loader/fsdp_load.py.
"FASTVIDEO_INFERENCE_TORCH_COMPILE":
lambda: os.getenv("FASTVIDEO_INFERENCE_TORCH_COMPILE", "0") != "0",
# Opt-in MiniMax-H3 inference-only Triton fusions adapted from the
# NVlabs/Sana Sol-Engine implementation. Accepts `all`, `1`, or a
# comma-separated subset of `modulate,qknorm_rope,swiglu`. An empty value
+28
View File
@@ -164,6 +164,18 @@ class FastVideoArgs:
torch_compile_kwargs_text_encoder: dict[str, Any] = field(default_factory=dict)
torch_compile_kwargs_vae: dict[str, Any] = field(default_factory=dict)
torch_compile_kwargs_audio_vae: dict[str, Any] = field(default_factory=dict)
# Regional (per-transformer-block) fullgraph torch.compile of the DiT at
# inference — the inference-side counterpart of the training regional
# compile ported from hao-ai-lab/FastVideo#1718. Applied by the loader
# right after the transformer loads, with fullgraph=True and inductor
# options {emulate_precision_casts: True} injected (no user kwargs
# needed). Attention backends that cannot be fullgraph-traced (VSA,
# FLASH_ATTN on flash-attn 3) degrade the transformer to eager with one
# warning. Opt-in via FASTVIDEO_INFERENCE_TORCH_COMPILE=1 (folded in
# __post_init__) or PipelineSelection.experimental
# {"inference_torch_compile": true}. Distinct from ``enable_torch_compile``,
# which keeps the pipeline-level compile semantics.
inference_torch_compile: bool = False
disable_autocast: bool = False
@@ -272,6 +284,13 @@ class FastVideoArgs:
self._apply_ltx2_vae_overrides()
self._resolve_refine_args()
self._apply_transformer_quant()
if not self.inference_torch_compile:
# Parse-once adapter (same pattern as attention_backend below): the
# environment variable is an input read once here, so the loader
# only ever consults the typed field.
import fastvideo.envs as envs
if envs.FASTVIDEO_INFERENCE_TORCH_COMPILE:
self.inference_torch_compile = True
if self.attention_backend is not None:
# Fail fast on typos instead of silently auto-selecting later.
from fastvideo.attention.selector import coerce_attn_backend
@@ -592,6 +611,15 @@ class FastVideoArgs:
help=
"JSON string of kwargs to pass to torch.compile. Example: '{\"backend\":\"inductor\",\"mode\":\"reduce-overhead\"}'",
)
parser.add_argument(
"--inference-torch-compile",
action=StoreBoolean,
default=FastVideoArgs.inference_torch_compile,
help="Regional fullgraph torch.compile of each DiT transformer block at inference "
"(port of the #1718 training-side regional compile). The loader injects fullgraph=True "
"and inductor options {emulate_precision_casts: true}; non-traceable attention backends "
"(VSA) degrade to eager with one warning. FASTVIDEO_INFERENCE_TORCH_COMPILE=1 is equivalent.",
)
parser.add_argument(
"--dit-cpu-offload",
@@ -1126,6 +1126,7 @@ class TransformerLoader(ComponentLoader):
training_mode=fastvideo_args.training_mode,
enable_torch_compile=fastvideo_args.enable_torch_compile,
torch_compile_kwargs=fastvideo_args.torch_compile_kwargs,
inference_regional_compile=fastvideo_args.inference_torch_compile,
)
total_params = sum(p.numel() for p in model.parameters())
+117
View File
@@ -136,6 +136,7 @@ def maybe_load_fsdp_model(
pin_cpu_memory: bool = True,
enable_torch_compile: bool = False,
torch_compile_kwargs: dict[str, Any] | None = None,
inference_regional_compile: bool = False,
) -> torch.nn.Module:
"""
Load the model with FSDP if is training, else load the model without FSDP.
@@ -235,9 +236,125 @@ def maybe_load_fsdp_model(
logger.info("Enabling torch.compile for FSDP training module with kwargs=%s", compile_kwargs)
model = torch.compile(model, **compile_kwargs)
logger.info("torch.compile enabled for %s", type(model).__name__)
elif inference_regional_compile and not training_mode:
# Inference-side counterpart of the #1718 training regional compile:
# per-block fullgraph compile right after the transformer loads, no
# user kwargs needed (fullgraph + emulate_precision_casts injected).
unsupported = _regional_compile_unsupported_reason(init_params)
if unsupported is not None:
logger.warning(
"inference_torch_compile requested but disabled: %s. "
"Inference continues in eager mode.", unsupported)
else:
prepare_for_compile = getattr(model, "prepare_for_compile", None)
if callable(prepare_for_compile):
logger.info("Running prepare_for_compile for %s", type(model).__name__)
prepare_for_compile()
_compile_model_regions(model, torch_compile_kwargs or {})
return model
def _regional_compile_unsupported_reason(init_params: dict[str, Any]) -> str | None:
"""Return why regional fullgraph compile cannot run, or None if it can.
FA3's grad-enabled attention path deliberately routes to the raw
autograd.Function at the cost of a dynamo graph break (see
flash_attn_default.py) — under regional ``fullgraph=True`` that break is
a hard RuntimeError at the first compiled forward. FA2 and FA4 route
through traceable custom ops and are compile-safe.
The VSA backends (Triton block-sparse kernels behind sequence-parallel
all-to-alls plus a host-synced metadata guard) are likewise not
fullgraph-traceable; a VSA-backed transformer (e.g. the FastH3 student)
falls back to eager while compile-safe dense loads still compile.
"""
try:
from fastvideo.attention.layer import _attention_compile_disabled
except Exception: # pragma: no cover - attention stack not importable
pass
else:
if _attention_compile_disabled():
# The escape hatch wraps attention forwards in
# torch.compiler.disable, which is a hard dynamo error inside a
# fullgraph region ("Skip inlining `torch.compiler.disable()`d
# function"). Degrade to eager instead, matching the hatch's
# debugging intent.
return ("FASTVIDEO_DISABLE_ATTENTION_COMPILE=1 keeps attention "
"forwards out of compiled graphs via torch.compiler."
"disable, which fullgraph regional compile cannot trace; "
"this model stays eager")
config = init_params.get("config")
resolved = getattr(config, "_resolved_attention_backend", None)
resolved_name = getattr(resolved, "name", "")
if resolved_name in ("VIDEO_SPARSE_ATTN", "VIDEO_SPARSE_ATTN_H3"):
return (f"attention backend resolved to {resolved_name}, whose Triton "
"kernels, sequence-parallel collectives, and sync metadata "
"guard graph-break (incompatible with fullgraph regional "
"compile); this model stays eager")
if resolved is None or resolved_name != "FLASH_ATTN":
return None
try:
from fastvideo.attention.utils.flash_attn_default import fa_version
except Exception: # pragma: no cover - flash-attn stack not importable
return None
if fa_version == "3":
return ("attention backend resolved to FLASH_ATTN with flash-attn 3, "
"whose grad-enabled path graph-breaks (incompatible with "
"fullgraph regional compile); use FA2, FA4 (FASTVIDEO_FA4=1), "
"or TORCH_SDPA for compiled runs")
return None
def _compile_model_regions(model: nn.Module, compile_kwargs: dict[str, Any]) -> int:
"""Compile repeated mathematical regions of a loaded model.
Only the selected module ``forward`` is replaced. This keeps activation
checkpoint wrappers structurally transparent while any module-level hooks
(FSDP pre/post, layerwise offload) execute outside the compiled region.
"""
compile_conditions = getattr(model, "_compile_conditions", None)
if not compile_conditions:
raise ValueError(f"{type(model).__name__} does not declare _compile_conditions")
if compile_kwargs.get("fullgraph", True) is not True:
raise ValueError("Regional compile requires fullgraph=True")
if "mode" in compile_kwargs:
# torch.compile forbids passing both `mode` and `options`, and
# regional compile always injects options (emulate_precision_casts)
# for bf16 numerics parity. Fail here with an actionable message
# instead of letting torch raise a mode/options conflict about an
# `options` key the user never wrote.
raise ValueError("Regional compile sets inductor options "
"(emulate_precision_casts) and cannot be combined "
"with torch_compile_kwargs['mode']. Remove 'mode' or "
"express its effect via torch_compile_kwargs['options'].")
kwargs = {**compile_kwargs, "fullgraph": True}
options = {"emulate_precision_casts": True}
options.update(kwargs.get("options") or {})
kwargs["options"] = options
compiled_count = 0
for name, submodule in list(model.named_modules()):
if not name:
continue
if any(condition(name, submodule) for condition in compile_conditions):
# Activation checkpoint wrappers are control-flow boundaries, not
# mathematical regions. Keep their saved-tensor/recompute logic
# eager and compile only the repeated block they own.
compile_target = getattr(submodule, "_checkpoint_wrapped_module", submodule)
compile_target.forward = torch.compile(compile_target.forward, **kwargs)
compiled_count += 1
if compiled_count == 0:
raise ValueError(f"No submodules in {type(model).__name__} matched _compile_conditions")
logger.info(
"Enabled regional torch.compile for %d submodules in %s with kwargs=%s",
compiled_count,
type(model).__name__,
kwargs,
)
return compiled_count
def shard_model(
model,
*,
@@ -203,6 +203,14 @@ class ComposedPipelineBase(ABC):
vae_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_vae or global_compile_kwargs)
audio_vae_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_audio_vae or global_compile_kwargs)
if compile_transformer and self.fastvideo_args.inference_torch_compile:
# The loader already applied the regional fullgraph
# compile to the DiT blocks (inference_torch_compile);
# wrapping the same forwards again here would stack
# compiled callables.
logger.info("inference_torch_compile already compiled the DiT regions in the "
"loader; skipping the pipeline-level DiT compile")
compile_transformer = False
if compile_transformer:
self._maybe_compile_pipeline_module(
module_name="transformer",
@@ -0,0 +1,117 @@
# SPDX-License-Identifier: Apache-2.0
"""Contract tests for the inference-side regional torch.compile port.
The loader applies a per-transformer-block fullgraph compile after the
transformer loads (``FastVideoArgs.inference_torch_compile``, env
``FASTVIDEO_INFERENCE_TORCH_COMPILE=1``). These tests pin the two pieces that
must not drift from the #1718 training-port semantics:
- ``_regional_compile_unsupported_reason``: VSA backends (and the attention
eager escape hatch) degrade to eager with a reason instead of hard-failing
fullgraph capture at the first denoising forward.
- ``_compile_model_regions``: exactly the ``_compile_conditions`` blocks are
compiled, fullgraph=True plus inductor ``emulate_precision_casts`` are
injected, and ``mode`` kwargs are rejected (torch.compile forbids
mode+options).
CPU-safe: torch.compile is monkeypatched, no CUDA needed.
"""
from types import SimpleNamespace
import pytest
import torch
from torch import nn
from fastvideo.models.loader import fsdp_load
from fastvideo.models.loader.fsdp_load import (
_compile_model_regions,
_regional_compile_unsupported_reason,
)
def _init_params_for(backend_name: str | None) -> dict:
resolved = None if backend_name is None else SimpleNamespace(name=backend_name)
return {"config": SimpleNamespace(_resolved_attention_backend=resolved)}
@pytest.mark.parametrize("backend_name", ["VIDEO_SPARSE_ATTN", "VIDEO_SPARSE_ATTN_H3"])
def test_vsa_backends_degrade_to_eager(backend_name, monkeypatch) -> None:
monkeypatch.delenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", raising=False)
reason = _regional_compile_unsupported_reason(_init_params_for(backend_name))
assert reason is not None
assert backend_name in reason
assert "eager" in reason
@pytest.mark.parametrize("backend_name", [None, "TORCH_SDPA"])
def test_dense_backends_allow_compile(backend_name, monkeypatch) -> None:
monkeypatch.delenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", raising=False)
assert _regional_compile_unsupported_reason(_init_params_for(backend_name)) is None
def test_attention_compile_escape_hatch_degrades_to_eager(monkeypatch) -> None:
monkeypatch.setenv("FASTVIDEO_DISABLE_ATTENTION_COMPILE", "1")
reason = _regional_compile_unsupported_reason(_init_params_for("TORCH_SDPA"))
assert reason is not None
assert "FASTVIDEO_DISABLE_ATTENTION_COMPILE" in reason
class _Block(nn.Module):
def __init__(self) -> None:
super().__init__()
self.linear = nn.Linear(4, 4)
def forward(self, x: torch.Tensor) -> torch.Tensor:
return self.linear(x)
class _Toy(nn.Module):
_compile_conditions = [lambda name, module: name.startswith("blocks.") and name.count(".") == 1]
def __init__(self) -> None:
super().__init__()
self.blocks = nn.ModuleList([_Block() for _ in range(3)])
self.proj_out = nn.Linear(4, 4)
def forward(self, x: torch.Tensor) -> torch.Tensor:
for block in self.blocks:
x = block(x)
return self.proj_out(x)
def test_compile_model_regions_injects_fullgraph_and_precision_casts(monkeypatch) -> None:
captured: list[dict] = []
def _fake_compile(fn, **kwargs):
captured.append(kwargs)
return fn
monkeypatch.setattr(fsdp_load.torch, "compile", _fake_compile)
model = _Toy()
count = _compile_model_regions(model, {})
# The three repeated blocks compile; proj_out and the root stay eager.
assert count == 3
assert len(captured) == 3
for kwargs in captured:
assert kwargs["fullgraph"] is True
assert kwargs["options"] == {"emulate_precision_casts": True}
def test_compile_model_regions_rejects_mode_kwargs() -> None:
with pytest.raises(ValueError, match="mode"):
_compile_model_regions(_Toy(), {"mode": "reduce-overhead"})
def test_compile_model_regions_requires_conditions_and_matches(monkeypatch) -> None:
monkeypatch.setattr(fsdp_load.torch, "compile", lambda fn, **kwargs: fn)
plain = nn.Linear(4, 4)
with pytest.raises(ValueError, match="_compile_conditions"):
_compile_model_regions(plain, {})
class _NoMatch(_Toy):
_compile_conditions = [lambda name, module: False]
with pytest.raises(ValueError, match="matched"):
_compile_model_regions(_NoMatch(), {})