Compare commits
2
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
9ba740e125 | ||
|
|
34346779d1 |
@@ -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 |
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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", "")
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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())
|
||||
|
||||
@@ -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(), {})
|
||||
Reference in New Issue
Block a user