diff --git a/docs/assets/cookbook-recipes.json b/docs/assets/cookbook-recipes.json index 9c8da1b49b..d146b05220 100644 --- a/docs/assets/cookbook-recipes.json +++ b/docs/assets/cookbook-recipes.json @@ -645,6 +645,29 @@ "evidence": "Source-backed", "limitations": ["The upstream checkpoint must be converted to Diffusers layout via scripts/checkpoint_conversion/convert_mmaudio_to_diffusers.py unless loaded from the FastVideo converted repo as done here."] }, + { + "id": "fasth3-spark-pair", + "family": "minimax_h3", + "stage": "inference", + "task": "Few-step text to video (with audio)", + "label": "FastH3 on two DGX Sparks (sequence parallel)", + "summary": "Run one FastH3 clip across two GB10s with Ray sequence parallel over QSFP RoCE. Sequential load and lazy module load stay on because SP replicates the DiT on each node.", + "model": "FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree", + "source": "examples/inference/basic/basic_fasth3_spark_pair.yaml", + "command": "source examples/inference/optimizations/spark_pair_env.sh && FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 FASTVIDEO_VAE_PARALLEL_DECODE=1 fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair.yaml", + "gpu_types": ["NVIDIA"], + "hardware": { + "platform": "cuda", + "gpu_count": 2, + "accelerator": "NVIDIA GB10 (DGX Spark pair)", + "evidence": "validated", + "evidence_url": "https://github.com/hao-ai-lab/FastVideo/pull/1803" + }, + "evidence": "Verified", + "expected_artifact": "MP4 under outputs/fasth3_spark_pair/", + "modes": ["T2VA", "2-Spark SP"], + "limitations": ["Requires a two-node Ray cluster on the QSFP interconnect. Height, width, frames, and steps in the YAML are examples. Edit them or pass CLI flags. See docs/getting_started/installation/spark_pair.md."] + }, { "id": "matrix-game-2", "family": "matrixgame", diff --git a/docs/cookbook/index.md b/docs/cookbook/index.md index 81dc4e4223..0dc666ad56 100644 --- a/docs/cookbook/index.md +++ b/docs/cookbook/index.md @@ -461,7 +461,10 @@ hide: Inference is the first complete stage. Distillation, fine-tuning, training, evaluation, optimization, and deployment will reuse the same family-first structure as their recipes land. Each family page shows - which stages are available and which are planned. + which stages are available and which are planned. Two DGX Sparks: bring up + the QSFP Ray cluster first + (pair two Sparks), + then pick the FastH3 two-Spark recipe.

diff --git a/docs/design/inference_schema_parity_inventory.yaml b/docs/design/inference_schema_parity_inventory.yaml index 38ac52ea1b..463977b4c8 100644 --- a/docs/design/inference_schema_parity_inventory.yaml +++ b/docs/design/inference_schema_parity_inventory.yaml @@ -30,6 +30,7 @@ surfaces: image_encoder_cpu_offload: generator.engine.offload.image_encoder vae_cpu_offload: generator.engine.offload.vae pin_cpu_memory: generator.engine.offload.pin_cpu_memory + lazy_module_load: generator.engine.offload.lazy_module_load enable_torch_compile: generator.engine.compile.enabled enable_torch_compile_text_encoder: generator.engine.compile.text_encoder_enabled enable_torch_compile_vae: generator.engine.compile.vae_enabled diff --git a/docs/getting_started/installation.md b/docs/getting_started/installation.md index 6d5c296393..50ee0a9226 100644 --- a/docs/getting_started/installation.md +++ b/docs/getting_started/installation.md @@ -4,8 +4,9 @@ FastVideo supports the following hardware platforms: - [NVIDIA CUDA](installation/gpu.md) -- [NVIDIA DGX Spark / GB10 (ARM64 + CUDA 13)](installation/spark.md) - ([performance & tuning](installation/spark_performance.md)) +- **NVIDIA DGX Spark / GB10 (ARM64 + CUDA 13)** — [install](installation/spark.md), + [performance](installation/spark_performance.md), + [pair two Sparks](installation/spark_pair.md) - [Apple silicon](installation/mps.md) ## Quick Installation diff --git a/docs/getting_started/installation/spark.md b/docs/getting_started/installation/spark.md index 326598826a..4a999e6423 100644 --- a/docs/getting_started/installation/spark.md +++ b/docs/getting_started/installation/spark.md @@ -144,6 +144,9 @@ for which models are practical on the GB10, what makes them faster, and what won't help on this hardware (and why) — so you don't spend a night tuning knobs that can't move here. +Two Sparks with QSFP cables: [Pair two NVIDIA DGX Sparks](spark_pair.md) for +one FastH3 clip across both GPUs (`sp_size=2` over Ray). + ## Development Environment Setup If you're planning to contribute to FastVideo please see the diff --git a/docs/getting_started/installation/spark_pair.md b/docs/getting_started/installation/spark_pair.md new file mode 100644 index 0000000000..f5bb92dde0 --- /dev/null +++ b/docs/getting_started/installation/spark_pair.md @@ -0,0 +1,209 @@ +# Pair two NVIDIA DGX Sparks + +One GB10 is 128 GB of unified LPDDR5X. FastH3 still fits on a single Spark with +[`lazy_module_load`](../../inference/offloading.md) (auto on GB10; sequential +load stands down when lazy owns deferral). Two boxes connected +by the QSFP ConnectX-7 cables can run **one clip faster** and can hold a +**longer clip** (up to the FastH3 15 s cap). + +This is FastVideo sequence parallel (`sp_size=2`) over Ray, not a third-party +xDiT vendor. Do not install xDiT for this path. + +## What two Sparks buy you + +| Goal | How | Use two Sparks? | +|---|---|---| +| Two independent videos at once | One process per box, `num_gpus=1` | Throughput only. Each clip still takes the 1-GPU time for that size. | +| One clip, faster | Ray + `sp_size=2` + parallel VAE | **Yes.** One 768×1344×124 recipe was 292 s vs 374 s on one GB10. | +| One clip, longer | Same, more frames | **Yes.** 345 frames (~14.4 s at 24 fps) finished in 587 s at 768×1344. | + +Sequence parallel **replicates** the DiT (~66 GiB per node). Lazy module load +is still required on each box. FSDP would shard weights; +it is untested on this fabric and is likely slower because every layer gathers +over ~21 GB/s RoCE. + +## Requirements + +- Two DGX Sparks with FastVideo [installed](spark.md) (CUDA 13, `aarch64`). +- The QSFP cables that ship with a dual-Spark kit, **ACTIVE** at 200 Gb/s: + `ibstat` should show the ConnectX-7 ports `LinkUp`. +- The same FastH3 snapshot on **both** NVMes. Copy the Hugging Face cache over + QSFP; do not download 100+ GB twice over Wi-Fi. +- Ray in the FastVideo venv (`uv pip install ray` if it is not already there). + +Each Spark has **one** GPU. `num_gpus=2` therefore means two nodes, which is +why the executor must be Ray (`mp` only works inside one process tree). + +## 1. Put IPv4 on the QSFP NICs + +The RoCE links often come up with no IPv4. Wi-Fi (`192.168.1.x`) is fine for +SSH and must stay the default route. NCCL and Ray must **not** use it. + +Pick a /24 that does not collide with your LAN. Example: + +| Node | QSFP IPv4 | Interface (typical) | +|---|---|---| +| Spark A | `192.168.23.1/24` | `enp1s0f1np1` | +| Spark B (Ray head) | `192.168.23.2/24` | `enp1s0f1np1` | + +Confirm names with `ibdev2netdev` and `ip -br link`. Then, as root, on each +box (NetworkManager likes to steal the NIC; unmanaged is enough for a session): + +```bash +sudo nmcli device set enp1s0f1np1 managed no +sudo ip addr replace 192.168.23.1/24 dev enp1s0f1np1 # .2 on the other box +sudo ip link set enp1s0f1np1 mtu 9000 up +``` + +These addresses do **not** survive reboot. Ping across the cable before +continuing: `ping -c 3 -I enp1s0f1np1 192.168.23.2`. + +A healthy fabric on this hardware looks like: + +- TCP iperf (jumbo 9000): ~40 Gb/s +- NCCL allreduce 1 GiB × 10: ~21 GB/s busbw (NVIDIA's dual-Spark figure is ~21.7) + +## 2. Start a two-node Ray cluster on the cable + +On **both** nodes, from the FastVideo repo, with the venv active: + +```bash +source examples/inference/optimizations/spark_pair_env.sh +``` + +That script pins NCCL and Gloo to the QSFP NIC/HCA, disables NVLink-style P2P +(there is none between boxes), and turns off Ray's memory monitor. The monitor +treats GB10 unified RSS during a 14-shard DiT load as a runaway and SIGTERMs +the worker around shard 11/14. Override `NCCL_SOCKET_IFNAME` / +`GLOO_SOCKET_IFNAME` if `ibdev2netdev` shows a different name. + +Cap Ray's object store. The default (~30% of 128 GB) leaves too little room +for the DiT: + +```bash +# Spark B — head +export FASTVIDEO_HOST_IP=192.168.23.2 +ray start --head --node-ip-address=192.168.23.2 --port=6379 --num-gpus=1 \ + --disable-usage-stats --object-store-memory=2147483648 --memory=4294967296 + +# Spark A — worker +export FASTVIDEO_HOST_IP=192.168.23.1 +ray start --address=192.168.23.2:6379 --node-ip-address=192.168.23.1 --num-gpus=1 \ + --disable-usage-stats --object-store-memory=2147483648 --memory=4294967296 +``` + +`FASTVIDEO_HOST_IP` **must** match `--node-ip-address`. If you omit it, Ray +advertises the Wi-Fi address, FastVideo builds a placement group for +`node:192.168.1.x`, and the QSFP workers never match. + +Check `ray status` on the head: `0.0/2.0 GPU` idle. + +## 3. Generate one FastH3 clip on both GPUs + +Run the driver on the **head**, same venv, same QSFP IP. + +`basic_fasth3.py` defaults target a four-GPU GB200 profile: 768×1344, `sm100a` +VSA, FA4, four GPUs. On Sparks you must override the kernel flags. Height, +width, frames, steps, seed, and prompt are yours. Change them. Legal +`num_frames` values are `17n+5`, capped at 345. + +GB10 has no FA4 / sm_100a VSA kernel, so `--vsa-kernel triton --no-fa4` stays +required on this box. `--execution-backend ray` is optional when `RAY_ADDRESS` +is already set. + +`--warmup --repeats 3` prints a median of three `generate()` calls after an +excluded warmup. Sequential load reloads Qwen for each later request, so that +protocol works. For a single cold process, pass `--no-warmup --repeats 1`. + +The command below is one example, not a required recipe: + +```bash +source examples/inference/optimizations/spark_pair_env.sh +export RAY_ADDRESS=192.168.23.2:6379 +export FASTVIDEO_HOST_IP=192.168.23.2 + +python examples/inference/basic/basic_fasth3.py \ + --model-path FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree \ + --num-gpus 2 --execution-backend ray \ + --vsa-kernel triton --no-fa4 \ + --warmup --repeats 3 --parallel-vae \ + --height 768 --width 1344 --num-frames 124 --steps 5 \ + --seed 2026 \ + --prompt "A wide cinematic shot of an alpine meadow at sunrise, pale pink mountain peaks above a blue valley filled with thin morning mist." \ + --output outputs/fasth3_spark_pair +``` + +Config-first equivalent. Edit the YAML the same way, `request.sampling` is not +locked: + +```bash +FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 \ +FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 \ +FASTVIDEO_VAE_PARALLEL_DECODE=1 FASTVIDEO_STAGE_LOGGING=1 \ +fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair.yaml +``` + +Stop the cluster when you are done: `ray stop` on both nodes. + +## FastH3 frame counts + +H3 is 24 fps. Legal `num_frames` values are `17n+5`. The pipeline rejects +clips longer than **15 s**. The longest legal length is **345 frames** +(14.375 s). 360 frames aligns to 362 and fails the duration check. + +## Measured on two GB10s (2026-08-31) + +These rows are full H3 VAE decode, Triton VSA, sequential + lazy load (auto on +GB10), parallel VAE. They are not a required size. Denoise times include +deferred DiT load (~35 s on the first generate). + +Cold process, `--no-warmup --repeats 1`, alpine prompt, 768×1344, 5 sigma +points (4 DiT forwards): + +| Run | GPUs | Frames | E2E | Denoise | VAE decode | +|---|---:|---:|---:|---:|---:| +| One Spark | 1 | 124 | 374–393 s | 180–188 s | 151–156 s | +| Two Sparks, SP=2 | 2 | 124 | **292 s** | **122 s** | **102 s** | +| Two Sparks, SP=2 | 2 | 345 | **587 s** | **351 s** | **173 s** | + +Warmup excluded, `--warmup --repeats 3` median, 512×896, 5 sigma points, full +VAE, same 4-step schedule: + +| Run | GPUs | Frames | Median E2E | Median denoise | +|---|---:|---:|---:|---:| +| One Spark | 1 | 124 | **251.4 s** | 94.2 s | +| Two Sparks, SP=2 | 2 | 124 | **215.2 s** | 72.4 s | + +Those medians used `--height` / `--width` / `--num-frames` as CLI flags. Swap +them. Native 480p on this model is 480×832, 124 frames. The 15 s cap is 345 +frames. + +The first VAE decode still pays `torch.compile`. Later `generate()` calls in +the same workers are cheaper. GB10 regional DiT compile stays off because the +sm_100a VSA kernel is not on this chip, so denoise is slower than a GB200 +`sm100a` run at the same geometry. + +## Troubleshooting + +| Symptom | Fix | +|---|---| +| Placement group waits forever / `node:192.168.1.x` | Set `FASTVIDEO_HOST_IP` to the QSFP address on **every** `ray start` **and** on the driver. | +| `RayDistributedExecutor` TypeError / abstract `set_log_queue` | Use a FastVideo build that implements those methods on the Ray executor (this page). | +| Worker SIGTERM during DiT shard 11/14 | `RAY_memory_monitor_refresh_ms=0` **before** `ray start`. Do not leave Ray's default 30% object store. | +| NCCL hangs or uses Wi-Fi | `source spark_pair_env.sh`. Confirm `NCCL_SOCKET_IFNAME` is the QSFP NIC. | +| Gloo `connectFullMesh` / `remote=[127.0.0.1]` | Two 1-GPU nodes must not use loopback as the Gloo store. Source `spark_pair_env.sh` so `GLOO_SOCKET_IFNAME` is the QSFP NIC on **each** box. FastVideo no longer copies that NIC name from the driver onto workers. | +| Second `generate()` crashes `NoneType.parameters` | Sequential load used to drop the text encoder without reloading it. This branch reloads Qwen for later requests so `--warmup --repeats N` works. | +| OOM / `earlyoom` prefers Python | Lazy module load must stay on (do not pass `--no-lazy-module-load`). Peak GPU during 345-frame denoise is ~90 GiB/node. | +| `num_gpus=2` on one Spark | Each Spark has one GPU. Use Ray across two nodes, or `num_gpus=1` on one box. | + +## What we are not claiming + +- **Throughput of many clips.** Two independent 1-GPU jobs still win if you + want two videos, not one faster video. +- **xDiT PipeFusion / CFG-parallel.** FastH3 is 4-step and has no CFG. +- **FSDP or tensor parallel as a speedup** on this 21 GB/s link. +- **Persistent networking.** The example IPs are session `ip addr replace`. + +More GPUs are legal while `num_attention_heads` (56 on FastH3) is divisible by +`sp_size`. Four Sparks would need a four-node fabric that this bring-up did +not exercise. diff --git a/docs/getting_started/installation/spark_performance.md b/docs/getting_started/installation/spark_performance.md index 490df05cb3..49c4c91460 100644 --- a/docs/getting_started/installation/spark_performance.md +++ b/docs/getting_started/installation/spark_performance.md @@ -161,21 +161,26 @@ is power-cycled. To avoid it: on: "CPU" offload uses the same unified RAM. Multi-GPU FSDP sharding remains available because it partitions weights without parking them in a separate host pool. -- **MiniMax H3 / FastH3** still needs sequential loading on one GB10. The Qwen3-VL +- **MiniMax H3 / FastH3** still needs deferred loading on one GB10. The Qwen3-VL conditioner is tens of gigabytes of BF16. If the DiT and VAEs load while that encoder is still resident, the process is a typical `earlyoom` kill (Python is - preferred). `h3_sequential_load` defaults to auto and turns this split on for - unified-memory devices. Do not pass `--no-h3-sequential-load` here. Force - `--h3-sequential-load` only if auto-detect misses the device. The CUDA pipeline - encodes first, releases the encoder, then loads DiT and VAEs onto the - accelerator (`to_cpu` follows `cpu_offload`, which is off here). See - [Offloading](../../inference/offloading.md). + preferred). On unified memory, `lazy_module_load` auto-enables and owns that + split (encoder, then DiT, then VAE; DiT can drop before decode). Sequential + load is the H3-only fallback when lazy is off; do not pass + `--no-lazy-module-load` here. Geometry scalars come from checkpoint + `config.json`, not live weights. See [Offloading](../../inference/offloading.md). - **FastH3 TAEH3** (`--video-decode-backend taeh3`) is an opt-in preview decoder. T2VA never materializes the 9.7 GiB video VAE (DiT still loads after Qwen via sequential start). On this box, alpine 768×1344×124 decoded in **2.4 s** versus **68 s** for the full VAE, and one T2VA generation finished in **224 s** end-to-end. Reconstruction is approximate, not lossless. FL2VA/Ref2VA still need the full VAE to encode references. +- **Two Sparks, one clip.** Sequence parallel (`sp_size=2`) over the QSFP RoCE + link ran one 768×1344×124 FastH3 recipe in **292 s** vs **374–393 s** on + one GB10, and a 345-frame (~14.4 s) clip in **587 s**. Other heights, widths, + and frame counts are valid. Weights stay replicated, so lazy module load + (auto on GB10) is still required on each box. Bring-up and knobs: + [Pair two NVIDIA DGX Sparks](spark_pair.md). ## Gotchas specific to the GB10 @@ -193,10 +198,11 @@ A few things that surprise people on this box (beyond the memory notes above): build recent enough to include its `transformers`-compatibility handling before running it. - **MiniMax H3 worker init can look healthy and still die on the first generate** - if sequential load is off (`--no-h3-sequential-load`, or auto-off on a - misclassified device) and encoder, VAE, and DiT load together. Confirm the log - contains `Released MiniMax-H3 text encoder after conditioning` before - `Loading MiniMax-H3 denoise modules`. + if deferred loading is off (`--no-lazy-module-load` and sequential also off) + and encoder, VAE, and DiT load together. On GB10 the log should show + `lazy_module_load owns deferral` (or, if lazy is off, sequential + `Released MiniMax-H3 text encoder after conditioning` before + `Loading MiniMax-H3 denoise modules`). ## Reproduce these numbers diff --git a/docs/inference/configuration.md b/docs/inference/configuration.md index 2fdfe17fd1..2c7c9bceb3 100644 --- a/docs/inference/configuration.md +++ b/docs/inference/configuration.md @@ -12,6 +12,22 @@ generator = VideoGenerator.from_pretrained( ) ``` +One node uses the multiprocessing executor (`execution_backend: mp`, the +default). Two machines — for example two DGX Sparks, one GPU each — need Ray: + +```yaml +generator: + engine: + num_gpus: 2 + execution_backend: ray + parallelism: + sp_size: 2 +``` + +Set `RAY_ADDRESS` and `FASTVIDEO_HOST_IP` to the interconnect IPs, not Wi-Fi. +The FastH3 example selects Ray automatically when `RAY_ADDRESS` is set. Full +bring-up: [Pair two NVIDIA DGX Sparks](../getting_started/installation/spark_pair.md). + ## Customizing Generation - `PipelineConfig`: Initialization time parameters diff --git a/docs/inference/offloading.md b/docs/inference/offloading.md index 08ad0f6aad..9045812719 100644 --- a/docs/inference/offloading.md +++ b/docs/inference/offloading.md @@ -12,6 +12,7 @@ text_encoder_cpu_offload: bool = True image_encoder_cpu_offload: bool = True vae_cpu_offload: bool = True pin_cpu_memory: bool = True +lazy_module_load: bool | None = None ``` On unified-memory accelerators such as NVIDIA GB10 and Apple silicon, FastVideo @@ -21,18 +22,25 @@ pool there, so offload adds transfers and duplicate residency instead of freeing memory. CUDA FSDP sharding remains enabled when requested; MPS continues to disable FSDP. `pin_cpu_memory` is not an offload mode and is left unchanged. -MiniMax H3 CUDA inference can use a second lever that does not copy weights to a -host pool: load the Qwen3-VL text encoder, run conditioning, then release that -encoder before loading the DiT and video/audio VAEs. `h3_sequential_load` -defaults to auto (`None`): on for unified-memory devices such as GB10, off on -discrete GPUs. Pass `--h3-sequential-load` to force it, or -`--no-h3-sequential-load` to keep the encoder resident for later `generate()` -calls on the same worker. Sequential load currently cannot re-encode a new prompt -on that worker; start a new generator until prompt-cache reload exists. The MLX -FastH3 runtime always uses this phase order. When host offload is off, DiT -safetensors are read onto the accelerator instead of CPU-then-copy. -Input-preparation geometry (spatial ratio, latent channels, audio sample rate) -comes from the VAE arch configs until those weights load. +MiniMax H3 CUDA inference can use two levers that do not copy weights to a host +pool. `lazy_module_load` is the general path: each opted-in component loads on +first use and is freed after its last stage, so a later `generate()` reloads +from disk in-process and the DiT can drop before VAE decode. On GB10 it +auto-enables and owns deferral. `h3_sequential_load` is the H3-only fallback +when lazy is off: load Qwen3-VL, run conditioning, release that encoder, then +load the DiT and VAEs. When both would arm, sequential stands down so VAE +`torch.compile` can attach to the lazy proxy. Input preparation and unpatchify +read geometry from checkpoint `config.json` (VAE spatial ratio / latent +channels, DiT patch size) so those stages do not materialize weights just to +read two integers. The MLX FastH3 runtime always uses this phase order. When +host offload is off, DiT safetensors are read onto the accelerator instead of +CPU-then-copy. Both flags default to auto (`None`) and turn on for +unified-memory devices such as GB10; lazy then disables sequential. Pass +`--no-lazy-module-load` to keep every component resident (sequential may still +auto-arm). Two-node Spark +jobs still need this split: sequence parallel replicates the DiT on each GB10 +(~66 GiB of weights plus activations). See +[Pair two NVIDIA DGX Sparks](../getting_started/installation/spark_pair.md). ## Behavior Explanation @@ -93,9 +101,12 @@ because the encoder has been released. #### Usage Recommendation -Leave the default on Spark / DGX Spark. Force `--h3-sequential-load` only when -you need the split on a discrete GPU. Use `--no-h3-sequential-load` when you -need more than one prompt per worker and have enough memory to keep the encoder. +Leave the default on Spark / DGX Spark when `lazy_module_load` is off. When +both would arm (the GB10 auto case), lazy owns deferral and sequential stands +down so VAE `torch.compile` can attach to the lazy proxy. Force +`--h3-sequential-load` only when you need the split on a discrete GPU without +lazy load. Use `--no-h3-sequential-load` when you need more than one prompt per +worker and have enough memory to keep the encoder. ### `text_encoder_cpu_offload` @@ -121,6 +132,52 @@ These options introduce performance overhead due to PCIe data transfer. We recommend enabling these options when OOM happens. +### `lazy_module_load` + +Every option above moves weights between host and device. This one changes +whether they are in memory at all. + +By default a pipeline loads every component before the first stage runs, so +peak memory is the sum of all of them even though no two are needed at the same +moment. With `lazy_module_load` enabled, each heavy component loads on first use +and is freed once the last stage that needs it has returned, so peak memory +becomes the largest overlapping set instead of the sum. MiniMax-H3 T2VA is +`max(text encoder, DiT, VAE)` rather than `text encoder + DiT + VAE`, because +the DiT is not held through VAE decode. + +#### Performance Impact + +A freed component is read from disk again on the next generation, so a +multi-prompt run pays one reload per component per request. For a large text +encoder that is tens of seconds. If pipeline-level `torch.compile` is enabled, +the compile setup is reapplied after each reload; PyTorch can reuse its graph +and kernel caches when the component structure and input shapes are unchanged. + +#### Usage Recommendation + +Enable this when a model does not fit at load time, which the CPU offload +options above cannot help with because they act after loading. It is +particularly relevant on unified-memory devices, where host and device draw on +the same pool and moving weights to the host frees nothing. FastVideo +auto-enables it there (`lazy_module_load=None`). Leave it off when the model +already fits, or pass `--no-lazy-module-load` to keep components resident for +later `generate()` calls. + +This option applies to inference only. Training keeps every component resident +and logs a warning if the flag is set. + +Deferral is opt-in per pipeline. Releasing a component and loading it again is +only safe when nothing outside the loader has changed it, and two common habits +break that without raising: mutating a component after load, as LongCat does +when it enables block-sparse attention, and reading a component's attributes +while stages are built, as the shared denoising stage does to pick an attention +backend. A pipeline therefore lists the components it has checked in +`_lazy_module_names`, which is empty in the base class. MiniMax-H3 opts in. On +a pipeline that has not, the flag is a no-op: hooks are not installed and no +warning is logged. Sequential MiniMax-H3 (`h3_sequential_load`) reloads the +text encoder for a later `generate()` on the same worker; you do not need to +start a new generator. + ## General Recommendations ### Single GPU Inference @@ -131,6 +188,12 @@ We recommend enabling `dit_layerwise_offload`. If OOM happens, also enable `imag We recommend enabling `use_fsdp_inference` and disabling both `dit_layerwise_offload` and `dit_cpu_offload`. If OOM happens, consider enabling `text_encoder_cpu_offload`, `image_encoder_cpu_offload`, and `vae_cpu_offload`. If OOM still happens, consider enabling `dit_cpu_offload`. +### When the Model Does Not Fit at Load Time + +The offload options only help once loading has finished. If the run dies while +components are still being placed, or if the machine has unified memory so +there is no separate host pool to offload into, enable `lazy_module_load`. + ## Examples ### Single GPU with Layerwise Offloading diff --git a/docs/inference/optimizations.md b/docs/inference/optimizations.md index 753fef2835..98e3920d11 100644 --- a/docs/inference/optimizations.md +++ b/docs/inference/optimizations.md @@ -7,7 +7,8 @@ This page describes the various options for speeding up generation times in Fast Several options on this page behave differently on the GB10's unified-memory hardware — some give little or nothing there. See [DGX Spark: Performance & Tuning](../getting_started/installation/spark_performance.md) - for what actually helps on that platform and why. + for what actually helps on that platform and why. Two Sparks, one clip: + [Pair two NVIDIA DGX Sparks](../getting_started/installation/spark_pair.md). ## Table of Contents diff --git a/docs/inference/support_matrix.md b/docs/inference/support_matrix.md index ab6f1227f6..6c58ebcb1e 100644 --- a/docs/inference/support_matrix.md +++ b/docs/inference/support_matrix.md @@ -221,6 +221,8 @@ Per the installation guides: [GPU install guide](../getting_started/installation/gpu.md). - **NVIDIA DGX Spark (GB10, aarch64)** — CUDA 13, from-source kernel build; see the [DGX Spark install guide](../getting_started/installation/spark.md). + Two Sparks over QSFP use Ray sequence parallel; see + [Pair two NVIDIA DGX Sparks](../getting_started/installation/spark_pair.md). - **Apple silicon** — macOS 14 or newer; FastMetal-QAD via the MLX runtime. See the [Apple Silicon guide](../getting_started/installation/mps.md). The older [`basic_mps.py`](https://github.com/hao-ai-lab/FastVideo/blob/main/examples/inference/basic/basic_mps.py) diff --git a/examples/inference/basic/basic_fasth3.py b/examples/inference/basic/basic_fasth3.py index d2ecfd125c..02f1e918b7 100644 --- a/examples/inference/basic/basic_fasth3.py +++ b/examples/inference/basic/basic_fasth3.py @@ -48,6 +48,14 @@ def build_parser(description: str | None = None) -> argparse.ArgumentParser: # License review completes. A local snapshot can be passed here instead. parser.add_argument("--prompt", required=True) parser.add_argument("--output", default="outputs/fasth3") + parser.add_argument("--lazy-module-load", + action=argparse.BooleanOptionalAction, + default=None, + help="load each heavy component on first use and free it after the last stage that " + "needs it, so peak memory is the largest overlapping set instead of the sum of every " + "component. Omit for auto (on for unified-memory devices such as GB10; off on discrete " + "GPUs). Costs a reload per generation; pass --no-lazy-module-load to keep every " + "component resident") parser.add_argument("--profile", choices=("all", "strict"), default="all", @@ -69,6 +77,13 @@ def build_parser(description: str | None = None) -> argparse.ArgumentParser: default=True, help="run one excluded request before timing") parser.add_argument("--num-gpus", type=int, default=4) + parser.add_argument( + "--execution-backend", + choices=("mp", "ray"), + default=None, + help="mp for one node; ray for a Ray cluster (two DGX Sparks). " + "Default: ray when RAY_ADDRESS is set, otherwise mp", + ) parser.add_argument("--vsa-sparsity", type=float, default=0.9, @@ -231,6 +246,12 @@ def validate_profile_dependencies(args: argparse.Namespace) -> None: "`cd fastvideo-kernel && ./build.sh`), or pass --vsa-kernel triton.") +def _execution_backend(args: argparse.Namespace) -> str: + if args.execution_backend is not None: + return args.execution_backend + return "ray" if os.environ.get("RAY_ADDRESS") else "mp" + + def build_generator_config(args: argparse.Namespace) -> GeneratorConfig: use_vsa = _uses_vsa(args) experimental: dict[str, object] = { @@ -261,6 +282,7 @@ def build_generator_config(args: argparse.Namespace) -> GeneratorConfig: ), engine=EngineConfig( num_gpus=args.num_gpus, + execution_backend=_execution_backend(args), use_fsdp_inference=args.num_gpus > 1 and not args.replicated_dit, parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus), offload=OffloadConfig( @@ -269,6 +291,7 @@ def build_generator_config(args: argparse.Namespace) -> GeneratorConfig: text_encoder=True, vae=True, pin_cpu_memory=args.pin_cpu_memory, + lazy_module_load=args.lazy_module_load, ), compile=CompileConfig( enabled=args.torch_compile, @@ -331,6 +354,7 @@ def run(args: argparse.Namespace) -> list[float]: f"Denoising contract override: {args.steps} sigma points = {args.steps - 1} DiT forwards") print("Profile environment: " + " ".join(f"{key}={value if value is not None else ''}" for key, value in environment.items())) + print(f"Execution backend: {_execution_backend(args)}") generator = VideoGenerator.from_config(build_generator_config(args)) measured_wall_times: list[float] = [] diff --git a/examples/inference/basic/basic_fasth3_spark_pair.yaml b/examples/inference/basic/basic_fasth3_spark_pair.yaml new file mode 100644 index 0000000000..b34d36ad29 --- /dev/null +++ b/examples/inference/basic/basic_fasth3_spark_pair.yaml @@ -0,0 +1,59 @@ +# FastH3 on two DGX Sparks (one GPU each) over QSFP RoCE. +# Bring up the Ray cluster first: docs/getting_started/installation/spark_pair.md +# +# request.sampling below is an example, not a required recipe. Change height, +# width, num_frames, num_inference_steps, seed, and prompt. Legal H3 frame +# counts are 17n+5, max 345 (15 s). +# +# source examples/inference/optimizations/spark_pair_env.sh +# export RAY_ADDRESS=:6379 +# export FASTVIDEO_HOST_IP= +# FASTVIDEO_VSA_SM100A=0 FASTVIDEO_FA4=0 FASTVIDEO_ATTENTION_BACKEND=VIDEO_SPARSE_ATTN_H3 \ +# FASTVIDEO_VAE_PARALLEL_DECODE=1 FASTVIDEO_STAGE_LOGGING=1 \ +# fastvideo generate --config examples/inference/basic/basic_fasth3_spark_pair.yaml +generator: + model_path: FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree + engine: + num_gpus: 2 + execution_backend: ray + use_fsdp_inference: false + parallelism: + tp_size: 1 + sp_size: 2 + offload: + dit: false + dit_layerwise: false + text_encoder: true + vae: true + pin_cpu_memory: true + lazy_module_load: true + compile: + enabled: false + vae_enabled: true + pipeline: + experimental: + attention_backend: VIDEO_SPARSE_ATTN_H3 + VSA_sparsity: 0.9 + VSA_tile_size: 64 + inference_torch_compile: true + vae_parallel_decode: true + vae_parallel_decode_strategy: gather + h3_sequential_load: true +request: + prompt: >- + A wide cinematic shot of an alpine meadow at sunrise, pale pink mountain + peaks above a blue valley filled with thin morning mist. + negative_prompt: "" + sampling: + seed: 2026 + height: 768 + width: 1344 + num_frames: 124 + fps: 24 + num_inference_steps: 5 + guidance_scale: 1.0 + batch_cfg: false + output: + output_path: outputs/fasth3_spark_pair/ + save_video: true + return_frames: false diff --git a/examples/inference/basic/basic_minimax_h3_t2v.py b/examples/inference/basic/basic_minimax_h3_t2v.py index 3e4344ae92..d4224f1162 100644 --- a/examples/inference/basic/basic_minimax_h3_t2v.py +++ b/examples/inference/basic/basic_minimax_h3_t2v.py @@ -4,6 +4,7 @@ from __future__ import annotations import argparse +import os from pathlib import Path from fastvideo import VideoGenerator @@ -38,7 +39,19 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--steps", type=int, default=50) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--num-gpus", type=int, default=4) + parser.add_argument( + "--execution-backend", + choices=("mp", "ray"), + default=None, + help="mp for one node; ray for a Ray cluster (two DGX Sparks). " + "Default: ray when RAY_ADDRESS is set, otherwise mp", + ) parser.add_argument("--torch-compile", action="store_true", help="torch.compile the DiT transformer path") + parser.add_argument("--compile-vae", + action=argparse.BooleanOptionalAction, + default=True, + help="compile the video VAE decoder independently of the DiT (on by default; " + "the Spark lazy-load path needs this registered before first materialize)") parser.add_argument("--compile-mode", default=None, help='torch.compile mode, e.g. "reduce-overhead" for CUDA graphs') @@ -48,6 +61,14 @@ def parse_args() -> argparse.Namespace: "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("--lazy-module-load", + action=argparse.BooleanOptionalAction, + default=None, + help="load each heavy component on first use and free it after the last stage that " + "needs it, so peak memory is the largest overlapping set instead of the sum of every " + "component. Omit for auto (on for unified-memory devices such as GB10; off on discrete " + "GPUs). Costs a reload per generation; pass --no-lazy-module-load to keep every " + "component resident") parser.add_argument("--repeats", type=int, default=1, @@ -67,12 +88,14 @@ def main() -> None: if args.inference_torch_compile: experimental["inference_torch_compile"] = True + execution_backend = args.execution_backend or ("ray" if os.environ.get("RAY_ADDRESS") else "mp") generator = VideoGenerator.from_config( GeneratorConfig( model_path=args.model_path, pipeline=PipelineSelection(experimental=experimental), engine=EngineConfig( num_gpus=args.num_gpus, + execution_backend=execution_backend, use_fsdp_inference=args.num_gpus > 1, parallelism=ParallelismConfig(tp_size=1, sp_size=args.num_gpus), offload=OffloadConfig( @@ -81,10 +104,12 @@ def main() -> None: text_encoder=True, vae=True, pin_cpu_memory=False, + lazy_module_load=args.lazy_module_load, ), compile=CompileConfig( enabled=args.torch_compile, mode=args.compile_mode, + vae_enabled=args.compile_vae, ), ), )) diff --git a/examples/inference/optimizations/spark_pair_env.sh b/examples/inference/optimizations/spark_pair_env.sh new file mode 100644 index 0000000000..c021faab5f --- /dev/null +++ b/examples/inference/optimizations/spark_pair_env.sh @@ -0,0 +1,21 @@ +# Source on every DGX Spark before `ray start` and before the FastH3 driver. +# QSFP ConnectX-7 interface names match the GB10 dual-Spark bring-up +# (enp1s0f1np1 / rocep1s0f1). Override NCCL_SOCKET_IFNAME / GLOO_SOCKET_IFNAME / +# NCCL_IB_HCA if `ibdev2netdev` shows different names. +# +# source examples/inference/optimizations/spark_pair_env.sh +# export FASTVIDEO_HOST_IP= +# +# See docs/getting_started/installation/spark_pair.md + +export NCCL_SOCKET_IFNAME="${NCCL_SOCKET_IFNAME:-enp1s0f1np1}" +export GLOO_SOCKET_IFNAME="${GLOO_SOCKET_IFNAME:-enp1s0f1np1}" +export NCCL_IB_HCA="${NCCL_IB_HCA:-rocep1s0f1,roceP2p1s0f1}" +# GB10 has no NVLink between boxes. Intra-node C2C P2P fights the QSFP path. +export NCCL_P2P_DISABLE="${NCCL_P2P_DISABLE:-1}" +export NCCL_CUMEM_ENABLE="${NCCL_CUMEM_ENABLE:-0}" +export NCCL_NVLS_ENABLE="${NCCL_NVLS_ENABLE:-0}" +# Ray's default memory monitor treats GB10 unified RSS during DiT load as a +# runaway and SIGTERMs the worker around shard 11/14. +export RAY_memory_monitor_refresh_ms="${RAY_memory_monitor_refresh_ms:-0}" +export RAY_memory_usage_threshold="${RAY_memory_usage_threshold:-1.0}" diff --git a/fastvideo/api/compat.py b/fastvideo/api/compat.py index 9e91d71298..07d4762cbf 100644 --- a/fastvideo/api/compat.py +++ b/fastvideo/api/compat.py @@ -126,6 +126,8 @@ def legacy_from_pretrained_to_config( offload["vae"] = value elif key == "pin_cpu_memory": offload["pin_cpu_memory"] = value + elif key == "lazy_module_load": + offload["lazy_module_load"] = value elif key == "enable_torch_compile": compile_config["enabled"] = value elif key == "enable_torch_compile_text_encoder": @@ -253,6 +255,7 @@ def generator_config_to_fastvideo_args(config: GeneratorConfig | Mapping[str, An "image_encoder_cpu_offload": engine.offload.image_encoder, "vae_cpu_offload": engine.offload.vae, "pin_cpu_memory": engine.offload.pin_cpu_memory, + "lazy_module_load": engine.offload.lazy_module_load, "enable_torch_compile": engine.compile.enabled, "torch_compile_kwargs": _compile_config_to_torch_kwargs(engine.compile), "enable_stage_verification": engine.enable_stage_verification, diff --git a/fastvideo/api/schema.py b/fastvideo/api/schema.py index 7b75cbc290..77224d0368 100644 --- a/fastvideo/api/schema.py +++ b/fastvideo/api/schema.py @@ -30,6 +30,12 @@ class OffloadConfig: image_encoder: bool = True vae: bool = True pin_cpu_memory: bool = True + # Not a CPU offload: loads each heavy component on first use and frees it + # after the last stage that needs it, so peak memory is the largest + # overlapping set rather than the sum. Grouped here because it is the same + # decision the offload knobs answer, which is how much of the model has to + # be resident at once. ``None`` auto-enables on unified-memory devices. + lazy_module_load: bool | None = None @dataclass diff --git a/fastvideo/attention/backends/video_sparse_attn_h3.py b/fastvideo/attention/backends/video_sparse_attn_h3.py index 38f0daa045..dcfd54f3a9 100644 --- a/fastvideo/attention/backends/video_sparse_attn_h3.py +++ b/fastvideo/attention/backends/video_sparse_attn_h3.py @@ -442,12 +442,20 @@ def __init__( self.prefix = prefix self.layer_idx = layer_idx_from_prefix(prefix, default=-1) self.head_size = head_size + # Generic torch.compile must not specialize the shared VSA forward on + # the Python ``layer_idx`` value of each of H3's 50 blocks. This + # tensor is prepared after weights load and drives only the compiled + # dense-layer decision; it does not opt the module into sm_100a. + self._compile_layer_idx: torch.Tensor | None = None # None means the regional-compile preparation hook has not run. The # eager path deliberately ignores this cache and preserves its # request-time env/probe/fallback behavior; only Dynamo capture reads # the prepared, static route. self._regional_compile_sm100a_enabled: bool | None = None - self._regional_compile_layer_idx: torch.Tensor | None = None + + def prepare_for_compile(self, device: torch.device) -> None: + """Tensorize per-layer state shared by every torch.compile route.""" + self._compile_layer_idx = torch.tensor(self.layer_idx, device=device, dtype=torch.int64) def prepare_for_regional_compile(self, device: torch.device) -> str | None: """Resolve the inference-only sm_100a route before fullgraph capture. @@ -459,6 +467,8 @@ def prepare_for_regional_compile(self, device: torch.device) -> str | None: the loaded model's device now, then let ``forward`` specialize on the resulting plain bool while Dynamo is compiling. """ + if self._compile_layer_idx is None: + self.prepare_for_compile(device) requested = os.environ.get(VSA_SM100A_ENV, "0") == "1" enabled = False reason = None if requested else f"{VSA_SM100A_ENV}=1 is required for compile-safe VSA-H3 attention" @@ -485,10 +495,6 @@ def prepare_for_regional_compile(self, device: torch.device) -> str | None: enabled = reason is None self._regional_compile_sm100a_enabled = enabled - # Keep this marker unset when preparation fails. Generic/training - # torch.compile must retain the established Triton attention route. - self._regional_compile_layer_idx = (torch.tensor(self.layer_idx, device=device, dtype=torch.int64) - if enabled else None) if enabled: route = ("native fastvideo-kernel mask entry" if callable( getattr(_sm100a, "block_sparse_attn_sm100a_from_mask", None)) else @@ -514,7 +520,7 @@ def tile(self, x: torch.Tensor, attn_metadata: MiniMaxH3VSAMetadata) -> torch.Te n_tiles = attn_metadata.variable_block_sizes.numel() grad_mode = torch.is_grad_enabled() and x.requires_grad compiling = torch.compiler.is_compiling() - regional_compiling = compiling and self._regional_compile_layer_idx is not None + regional_compiling = compiling and self._regional_compile_sm100a_enabled is True if regional_compiling: sm100a_requested = bool(self._regional_compile_sm100a_enabled) elif compiling: @@ -560,7 +566,7 @@ def forward( # type: ignore[override] attn_metadata: MiniMaxH3VSAMetadata, ) -> torch.Tensor: compiling = torch.compiler.is_compiling() - regional_compiling = compiling and self._regional_compile_layer_idx is not None + regional_compiling = compiling and self._regional_compile_sm100a_enabled is True tile_elems = attn_metadata.tile_elems if regional_compiling and tile_elems != 64: @@ -603,13 +609,14 @@ def forward( # type: ignore[override] logical_gate = gate_compress[:, :logical_seq_len] if gate_compress is not None else None # Probe-guided per-layer opt-out: diffuse layers run dense (all-True - # mask) while the rest keep the configured sparsity. During regional - # capture, keep the layer decision tensor-valued so the 50 block - # instances reuse one graph instead of specializing on layer_idx. + # mask) while the rest keep the configured sparsity. During any + # prepared capture, keep the layer decision tensor-valued so the 50 + # block instances reuse one graph instead of specializing on the + # Python layer_idx attribute. force_dense = None - if regional_compiling: - assert self._regional_compile_layer_idx is not None - force_dense = (attn_metadata.dense_layers_tensor == self._regional_compile_layer_idx).any() + compile_layer_idx = self._compile_layer_idx if compiling else None + if compile_layer_idx is not None: + force_dense = (attn_metadata.dense_layers_tensor == compile_layer_idx).any() layer_sparsity = attn_metadata.VSA_sparsity else: layer_sparsity = 0.0 if self.layer_idx in attn_metadata.dense_layers else attn_metadata.VSA_sparsity diff --git a/fastvideo/fastvideo_args.py b/fastvideo/fastvideo_args.py index dd6dbb1810..a648a02538 100644 --- a/fastvideo/fastvideo_args.py +++ b/fastvideo/fastvideo_args.py @@ -173,6 +173,17 @@ class FastVideoArgs: taeh3_checkpoint: str | None = None taeh3_chunk_size: int = 5 + # Load each heavy component on first use and free it once the last stage + # that holds it has run, instead of keeping every component resident from + # load time to shutdown. Peak memory becomes the largest overlapping set + # rather than the sum of all components. ``None`` (auto) turns this on for + # unified-memory devices (GB10 / Spark) after the worker binds its device, + # and leaves it off on discrete GPUs. Explicit True / False overrides the + # probe. A released component is re-read from disk on the next generation, + # so this trades per-request latency for headroom. Inference only; training + # keeps every component resident. + lazy_module_load: bool | None = None + # Sequence-parallel MiniMax-H3 VAE (opt-in, default off). With SP > 1 the # video VAE's temporal chunks (decode) and clips (reference encode) are # round-robined across the sequence-parallel ranks and reassembled @@ -721,6 +732,15 @@ def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser: action=StoreBoolean, help="Use CPU offload for VAE. Enable if run out of memory.", ) + parser.add_argument( + "--lazy-module-load", + action=argparse.BooleanOptionalAction, + default=None, + help="Load each heavy component on first use and free it after the last stage that needs it, " + "so peak memory is the largest overlapping set of components instead of their sum. " + "Omit for auto (on for unified-memory devices such as GB10; off on discrete GPUs). " + "Pass --no-lazy-module-load to keep every component resident.", + ) parser.add_argument( "--pin-cpu-memory", action=StoreBoolean, @@ -989,6 +1009,20 @@ def _resolve_device_offload_conflicts(self) -> None: def finalize_device_offload_policy(self, device_id: int = 0) -> bool: """Apply device-local memory policy, then resolve incompatible modes.""" has_unified_memory = self.disable_offload_on_unified_memory(device_id) + if self.lazy_module_load is None: + self.lazy_module_load = bool(has_unified_memory) and not self.training_mode + if self.lazy_module_load: + from fastvideo.platforms import current_platform + + try: + device_name = current_platform.get_device_name(device_id) + except Exception: + device_name = current_platform.device_name + logger.info( + "Enabling lazy_module_load: %s has unified memory, so encoder, DiT, and VAEs cannot stay " + "resident together. Pass --no-lazy-module-load to keep every component loaded.", + device_name, + ) self._resolve_device_offload_conflicts() return has_unified_memory diff --git a/fastvideo/models/dits/minimax_h3.py b/fastvideo/models/dits/minimax_h3.py index ec168cec18..3dff637fe4 100644 --- a/fastvideo/models/dits/minimax_h3.py +++ b/fastvideo/models/dits/minimax_h3.py @@ -733,31 +733,49 @@ def __init__(self, config: MiniMaxH3Config, hf_config: dict[str, Any]) -> None: ) self.__post_init__() + @staticmethod + def _compile_setup_device(attention: MiniMaxH3Attention) -> torch.device: + """Return the loaded device even when FP8 replaced the query weight.""" + query_state = next(attention.to_q.parameters(), None) + if query_state is None: + query_state = next(attention.to_q.buffers(), None) + if query_state is None: + raise RuntimeError("MiniMax H3 to_q has no materialized parameter or buffer for compile setup.") + return query_state.device + def prepare_for_compile(self) -> None: """Pipeline hook, called once right before torch.compile wraps the blocks. - Resolve each loaded VSA compression gate eagerly. Generic and training - compile retain their established attention dispatch; only the - inference loader's separate ``prepare_for_regional_compile`` hook may - preselect the inference-only sm_100a path. + Resolve each loaded VSA compression gate eagerly and tensorize its + layer identity so repeated blocks share one Dynamo graph. Generic and + training compile retain their established attention dispatch; only + the inference loader's separate ``prepare_for_regional_compile`` hook + may preselect the inference-only sm_100a path. The inference-only Triton fusions expose fake-backed custom operators, so Dynamo can keep them active as opaque nodes inside each fullgraph block instead of tracing into their launcher implementation. """ gate_states: list[bool] = [] + prepared_vsa_impls = 0 for block in self.transformer_blocks: attention = block.attn if attention.to_gate_compress is not None: attention._resolve_gate_compress_for_compile() assert attention._gate_compress_active is not None gate_states.append(attention._gate_compress_active) + prepare_vsa = getattr(attention.distributed_attention.attn_impl, "prepare_for_compile", None) + if callable(prepare_vsa): + prepare_vsa(self._compile_setup_device(attention)) + prepared_vsa_impls += 1 if gate_states: logger.info( "Resolved MiniMax H3 VSA compression gates before torch.compile: %d active, %d inactive", sum(gate_states), len(gate_states) - sum(gate_states), ) + if prepared_vsa_impls: + logger.info("Prepared %d MiniMax H3 VSA layer indices for torch.compile", prepared_vsa_impls) if self.enabled_fusions: logger.info( "MiniMax H3 inference fusions remain active under torch.compile through custom-op boundaries: %s", @@ -774,14 +792,7 @@ def prepare_for_regional_compile(self) -> str | None: prepare_vsa = getattr(attention.distributed_attention.attn_impl, "prepare_for_regional_compile", None) if not callable(prepare_vsa): continue - # Post-load FP8 conversion may replace to_q.weight with packed - # buffers. Either representation identifies the local device. - query_state = next(attention.to_q.parameters(), None) - if query_state is None: - query_state = next(attention.to_q.buffers(), None) - if query_state is None: - raise RuntimeError("MiniMax H3 to_q has no materialized parameter or buffer for compile setup.") - unsupported = prepare_vsa(query_state.device) + unsupported = prepare_vsa(self._compile_setup_device(attention)) if unsupported: unsupported_reasons.add(str(unsupported)) prepared_vsa_impls += 1 diff --git a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py index 02d618497c..691561c40a 100644 --- a/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py +++ b/fastvideo/pipelines/basic/minimax_h3/minimax_h3_pipeline.py @@ -11,6 +11,7 @@ import torch from fastvideo.configs.models.vaes.minimax_h3_audio import MiniMaxH3AudioVAEArchConfig +from fastvideo.configs.models.vaes.minimax_h3_video import MiniMaxH3VideoVAEArchConfig from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig from fastvideo.fastvideo_args import FastVideoArgs from fastvideo.logger import init_logger @@ -35,39 +36,59 @@ _DENOISE_MODULE_NAMES = ("vae", "audio_vae", "transformer") +@dataclass(frozen=True) +class _H3VideoGeometry: + spatial_compression_ratio: int + latent_channels: int + + +@dataclass(frozen=True) +class _H3AudioGeometry: + sampling_rate: int + + +def _default_video_geometry() -> _H3VideoGeometry: + arch = MiniMaxH3VideoVAEArchConfig() + return _H3VideoGeometry( + spatial_compression_ratio=int(arch.spatial_compression_ratio), + latent_channels=int(arch.latent_channels), + ) + + +def _default_audio_geometry() -> _H3AudioGeometry: + return _H3AudioGeometry(sampling_rate=int(MiniMaxH3AudioVAEArchConfig().sampling_rate)) + + def _apply_h3_checkpoint_arch_configs(model_path: str, fastvideo_args: FastVideoArgs, extra_config_module_map: dict[str, str]) -> None: """Overlay checkpoint config.json onto pipeline configs without loading weights.""" root = Path(model_path) - vae_dir = root / "vae" + vae_dir = root / extra_config_module_map.get("vae", "vae") if (vae_dir / "config.json").is_file(): fastvideo_args.pipeline_config.vae_config.update_model_arch(get_diffusers_config(str(vae_dir))) + audio_vae_dir = root / extra_config_module_map.get("audio_vae", "audio_vae") + audio_vae_config = getattr(fastvideo_args.pipeline_config, "audio_vae_config", None) + if audio_vae_config is not None and (audio_vae_dir / "config.json").is_file(): + audio_vae_config.update_model_arch(get_diffusers_config(str(audio_vae_dir))) transformer_dir = root / extra_config_module_map.get("transformer", "transformer") if (transformer_dir / "config.json").is_file(): fastvideo_args.pipeline_config.dit_config.update_model_arch(get_diffusers_config(str(transformer_dir))) - dit_arch = getattr(fastvideo_args.pipeline_config.dit_config, "arch_config", None) + dit_config = fastvideo_args.pipeline_config.dit_config vae_arch = getattr(fastvideo_args.pipeline_config.vae_config, "arch_config", None) - logger.info( - "MiniMax-H3 geometry from config: patch_size=%s spatial_compression_ratio=%s latent_channels=%s", - getattr(dit_arch, "patch_size", None), - getattr(vae_arch, "spatial_compression_ratio", None), - getattr(vae_arch, "latent_channels", None), - ) + patch_size = getattr(dit_config, "patch_size", None) + if patch_size is not None and vae_arch is not None: + logger.info( + "MiniMax-H3 geometry from config: patch_size=%s spatial_compression_ratio=%s latent_channels=%s", + tuple(patch_size), + int(getattr(vae_arch, "spatial_compression_ratio", 0)), + int(getattr(vae_arch, "latent_channels", 0)), + ) def _use_taeh3_t2va(fastvideo_args: FastVideoArgs | None, *, ref2va: bool) -> bool: return (not ref2va) and getattr(fastvideo_args, "video_decode_backend", "h3-vae") == "taeh3" -@dataclass(frozen=True) -class _H3AudioGeometry: - sampling_rate: int - - -def _default_audio_geometry() -> _H3AudioGeometry: - return _H3AudioGeometry(sampling_rate=int(MiniMaxH3AudioVAEArchConfig().sampling_rate)) - - class MiniMaxH3BasePipeline(LoRAPipeline, ComposedPipelineBase): """Shared loading and target-generation path for MiniMax H3. @@ -104,6 +125,11 @@ class MiniMaxH3BasePipeline(LoRAPipeline, ComposedPipelineBase): "scheduler", "audio_scheduler", ] + # Deferral is safe here: geometry scalars come from checkpoint config.json + # (applied in initialize_pipeline without loading weights), no stage + # constructor reads a deferred component, and initialize_pipeline only + # inspects the schedulers, which are never deferred. + _lazy_module_names = ("text_encoder", "transformer", "vae", "audio_vae") def __init__(self, *args: Any, **kwargs: Any) -> None: self._ref2va = getattr(self, "_ref2va_default", False) @@ -127,6 +153,16 @@ def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None: def _defer_denoise_modules(self, fastvideo_args: FastVideoArgs) -> bool: if not fastvideo_args.inference_mode or bool(getattr(fastvideo_args, "training_mode", False)): return False + # Both mechanisms defer the same four modules and both decide when to + # free them. Running them together strips DiT/VAEs from the first load + # (sequential) while the base wraps the encoder in a proxy (lazy), so + # post_init's VAE compile transform has nothing to attach to. Lazy is + # the more general owner — including auto-on for unified memory — so it + # wins whenever it is on. Sequential remains the H3-only fallback when + # lazy is off. + if bool(getattr(fastvideo_args, "lazy_module_load", False)): + logger.info("MiniMax-H3 sequential module load off: lazy_module_load owns deferral") + return False requested = fastvideo_args.h3_sequential_load if requested is True: return True @@ -191,6 +227,7 @@ def _load_denoise_modules(self, fastvideo_args: FastVideoArgs) -> None: loaded = super().load_modules(fastvideo_args, loaded_modules=self.modules) for name, module in loaded.items(): self.add_module(name, module) + self._apply_inference_compile(tuple(name for name in loaded if name in _DENOISE_MODULE_NAMES)) finally: self._required_config_modules = saved @@ -207,23 +244,70 @@ def _release_text_encoder(self) -> None: if torch.cuda.is_available(): torch.cuda.empty_cache() + def _ensure_text_encoder(self, fastvideo_args: FastVideoArgs) -> None: + """Reload Qwen3-VL after `_release_text_encoder` so a later request can encode.""" + encoder = self.get_module("text_encoder") + stage = self._stage_name_mapping.get("conditioning_stage") + if encoder is not None: + if stage is not None and getattr(stage, "conditioner", None) is None: + stage.conditioner = encoder + return + saved = list(self.required_config_modules) + self._required_config_modules = ["text_encoder"] + try: + logger.info("Reloading MiniMax-H3 text encoder for a subsequent request") + loaded = super().load_modules(fastvideo_args, loaded_modules=self.modules) + for name, module in loaded.items(): + self.add_module(name, module) + self._apply_inference_compile(("text_encoder", )) + finally: + self._required_config_modules = saved + if stage is not None: + stage.conditioner = self.get_module("text_encoder") + + def _run_condition_then_denoise(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch: + for name in ("input_preparation_stage", "conditioning_stage"): + batch = self._stage_name_mapping[name](batch, fastvideo_args) + self._release_text_encoder() + self._load_denoise_modules(fastvideo_args) + if not self._denoise_stages_ready: + self._add_denoise_stages(ref2va=self._ref2va) + for name in ( + "latent_preparation_stage", + "denoising_stage", + "video_decoding_stage", + "audio_decoding_stage", + ): + batch = self._stage_name_mapping[name](batch, fastvideo_args) + return batch + + def _input_video_geometry(self, fastvideo_args: FastVideoArgs) -> Any: + """Read canvas scalars from checkpoint JSON, not a live VAE proxy.""" + arch = getattr(getattr(fastvideo_args.pipeline_config, "vae_config", None), "arch_config", None) + if arch is not None: + return arch + return _default_video_geometry() + def _input_vae(self) -> Any: live = self.get_module("vae") if live is not None: return live - return self.fastvideo_args.pipeline_config.vae_config.arch_config + return self._input_video_geometry(self.fastvideo_args) - def _input_audio_vae(self, *, ref2va: bool) -> Any | None: + def _input_audio_vae(self, fastvideo_args: FastVideoArgs, *, ref2va: bool) -> Any | None: if not ref2va: return None - return self.get_module("audio_vae") or _default_audio_geometry() + arch = getattr(getattr(fastvideo_args.pipeline_config, "audio_vae_config", None), "arch_config", None) + if arch is not None: + return arch + return _default_audio_geometry() - def _add_condition_stages(self, *, ref2va: bool) -> None: + def _add_condition_stages(self, fastvideo_args: FastVideoArgs, *, ref2va: bool) -> None: self.add_stage( "input_preparation_stage", MiniMaxH3InputPreparationStage( - vae=self._input_vae(), - audio_vae=self._input_audio_vae(ref2va=ref2va), + vae=self._input_video_geometry(fastvideo_args), + audio_vae=self._input_audio_vae(fastvideo_args, ref2va=ref2va), ref2va=ref2va, ), ) @@ -270,9 +354,9 @@ def _add_denoise_stages(self, *, ref2va: bool) -> None: self.add_stage("audio_decoding_stage", MiniMaxH3AudioDecodingStage(audio_vae=audio_vae)) self._denoise_stages_ready = True - def _add_stages(self, *, ref2va: bool) -> None: + def _add_stages(self, fastvideo_args: FastVideoArgs, *, ref2va: bool) -> None: self._ref2va = ref2va - self._add_condition_stages(ref2va=ref2va) + self._add_condition_stages(fastvideo_args, ref2va=ref2va) if self._denoise_modules_loaded(): self._add_denoise_stages(ref2va=ref2va) @@ -280,31 +364,30 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward if not self.post_init_called: self.post_init() - if self._denoise_stages_ready: - return super().forward(batch, fastvideo_args) - - logger.info("Running MiniMax-H3 condition stages before loading DiT/VAE weights") - for stage in self.stages: - batch = stage(batch, fastvideo_args) - self._release_text_encoder() - self._load_denoise_modules(fastvideo_args) - self._add_denoise_stages(ref2va=self._ref2va) - for name in ( - "latent_preparation_stage", - "denoising_stage", - "video_decoding_stage", - "audio_decoding_stage", - ): - batch = self._stage_name_mapping[name](batch, fastvideo_args) - return batch + # Sequential encode-then-release is the H3-only fallback. Lazy and the + # fully-resident discrete-GPU path both keep a complete stage list and + # must use the base forward so abort cleanup and text_encoder_cpu_offload + # still apply. Releasing Qwen on every request was re-reading it from disk + # when neither deferral flag was on. + if self._defer_denoise_modules(fastvideo_args): + try: + self._ensure_text_encoder(fastvideo_args) + if self._denoise_stages_ready: + logger.info("Running MiniMax-H3 condition stages before denoise (subsequent request)") + else: + logger.info("Running MiniMax-H3 condition stages before loading DiT/VAE weights") + return self._run_condition_then_denoise(batch, fastvideo_args) + except BaseException: + self._release_all_lazy_modules() + raise + return super().forward(batch, fastvideo_args) class MiniMaxH3Pipeline(MiniMaxH3BasePipeline): """One-request joint video/stereo-audio pipeline for T2VA and FL2VA.""" def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None: - del fastvideo_args - self._add_stages(ref2va=False) + self._add_stages(fastvideo_args, ref2va=False) class MiniMaxH3RefPipeline(MiniMaxH3BasePipeline): @@ -314,8 +397,7 @@ class MiniMaxH3RefPipeline(MiniMaxH3BasePipeline): _ref2va_default = True def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None: - del fastvideo_args - self._add_stages(ref2va=True) + self._add_stages(fastvideo_args, ref2va=True) class MiniMaxH3ModularPipeline(MiniMaxH3Pipeline): diff --git a/fastvideo/pipelines/basic/minimax_h3/packing.py b/fastvideo/pipelines/basic/minimax_h3/packing.py index 1c67c49a35..db5d0538ed 100644 --- a/fastvideo/pipelines/basic/minimax_h3/packing.py +++ b/fastvideo/pipelines/basic/minimax_h3/packing.py @@ -38,19 +38,36 @@ MINIMAX_H3_KEYFRAME_NOISE_AUG = 0.999 MINIMAX_H3_KEYFRAME_ENCODE_SEED = 42 +_PATCH_SIZE_CACHE: dict[int, tuple[int, int, int]] = {} + def h3_dit_patch_size(fastvideo_args: Any) -> tuple[int, int, int]: """Read DiT patch size from pipeline config, not live transformer weights.""" dit_config = getattr(getattr(fastvideo_args, "pipeline_config", None), "dit_config", None) + cached = _PATCH_SIZE_CACHE.get(id(dit_config)) if dit_config is not None else None + if cached is not None: + return cached patch_size = getattr(dit_config, "patch_size", None) if patch_size is None: raise ValueError("MiniMax-H3 requires pipeline_config.dit_config.patch_size.") - values = tuple(int(axis) for axis in patch_size) - if len(values) != 3 or min(values) <= 0: + axes = tuple(int(axis) for axis in patch_size) + if len(axes) != 3 or min(axes) <= 0: raise ValueError(f"MiniMax-H3 patch_size must be three positive ints, got {patch_size!r}.") + values = (axes[0], axes[1], axes[2]) + if dit_config is not None: + _PATCH_SIZE_CACHE[id(dit_config)] = values return values +def h3_latent_channels(model_config: Any, name: str) -> int: + """Read VAE latent width from arch config, not a live VAE proxy.""" + arch = getattr(model_config, "arch_config", None) + value = getattr(arch, "latent_channels", None) + if value is None: + raise ValueError(f"MiniMax-H3 requires {name}.arch_config.latent_channels") + return int(value) + + MINIMAX_H3_ROPE_FRAME_RESCALE = 5.0 / 3.0 MINIMAX_H3_ROPE_FRAMES_PER_LATENT = (1, 4, 4, 4, 4) _ROPE_SPATIAL_SCALE = 32 diff --git a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_latent_preparation.py b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_latent_preparation.py index 2d3b90c6a7..60ffb25541 100644 --- a/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_latent_preparation.py +++ b/fastvideo/pipelines/basic/minimax_h3/stages/minimax_h3_latent_preparation.py @@ -22,6 +22,7 @@ build_packed_sequence, build_ref2va_packed_sequence, h3_dit_patch_size, + h3_latent_channels, keyframe_condition_noise, patchify_video_latents, ) @@ -59,6 +60,17 @@ def _sample_visual_posterior(posterior: Any) -> torch.Tensor: return posterior.sample(generator=generator) +def _video_latent_channels(fastvideo_args: FastVideoArgs) -> int: + return h3_latent_channels(fastvideo_args.pipeline_config.vae_config, "vae_config") + + +def _audio_latent_channels(fastvideo_args: FastVideoArgs) -> int: + return h3_latent_channels( + getattr(fastvideo_args.pipeline_config, "audio_vae_config", None), + "audio_vae_config", + ) + + class MiniMaxH3LatentPreparationStage(PipelineStage): """Encode fixed conditions, build the row layout, then draw target noise.""" @@ -196,7 +208,7 @@ def _encode_fl2va_conditions( noise = keyframe_condition_noise( shapes, h3_dit_patch_size(fastvideo_args), - self.vae.latent_channels, + _video_latent_channels(fastvideo_args), generator=batch.generator, device=device, ) @@ -244,7 +256,7 @@ def _encode_ref2va_conditions( noise = keyframe_condition_noise( shapes, h3_dit_patch_size(fastvideo_args), - self.vae.latent_channels, + _video_latent_channels(fastvideo_args), generator=batch.generator, device=device, ) @@ -320,10 +332,11 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward h3_dit_patch_size(fastvideo_args)) num_audio_latents = layout.num_audio_latents - expected_audio_shape = (MINIMAX_H3_AUDIO_CHANNELS, self.audio_vae.latent_channels, num_audio_latents) + audio_channels = _audio_latent_channels(fastvideo_args) + expected_audio_shape = (MINIMAX_H3_AUDIO_CHANNELS, audio_channels, num_audio_latents) if audio_noise is None: audio_rows = randn_tensor( - (num_audio_latents * MINIMAX_H3_AUDIO_CHANNELS, self.audio_vae.latent_channels), + (num_audio_latents * MINIMAX_H3_AUDIO_CHANNELS, audio_channels), generator=batch.generator, device=device, dtype=torch.float32, @@ -332,9 +345,7 @@ def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> Forward if tuple(audio_noise.shape) != expected_audio_shape: raise ValueError(f"MiniMax-H3 injected audio latents must have shape {expected_audio_shape}, " f"got {tuple(audio_noise.shape)}.") - audio_rows = audio_noise.to(device=device, - dtype=torch.float32).permute(0, 2, - 1).reshape(-1, self.audio_vae.latent_channels) + audio_rows = audio_noise.to(device=device, dtype=torch.float32).permute(0, 2, 1).reshape(-1, audio_channels) if condition_video is not None: video_rows = torch.cat((condition_video.to(device), video_rows)) diff --git a/fastvideo/pipelines/composed_pipeline_base.py b/fastvideo/pipelines/composed_pipeline_base.py index ec1acb963c..183a34af47 100644 --- a/fastvideo/pipelines/composed_pipeline_base.py +++ b/fastvideo/pipelines/composed_pipeline_base.py @@ -8,6 +8,8 @@ import argparse import os from abc import ABC, abstractmethod +from collections.abc import Iterator +from functools import partial from typing import Any, cast import torch @@ -24,6 +26,7 @@ from fastvideo.logger import init_logger from fastvideo.profiler import get_or_create_profiler from fastvideo.models.loader.component_loader import PipelineComponentLoader +from fastvideo.pipelines.lazy_module import LazyModule, is_lazy_module from fastvideo.pipelines.pipeline_batch_info import ForwardBatch from fastvideo.pipelines.stages import PipelineStage import fastvideo.envs as envs @@ -32,6 +35,40 @@ logger = init_logger(__name__) +def _iter_held_objects(stage: PipelineStage) -> Iterator[Any]: + """Yield everything a stage holds, walking into nested stages. + + A stage can compose others rather than hold a component directly: + ``Cosmos25AutoDenoisingStage`` keeps the transformer inside its ``_t2w`` + and ``_v2w`` children. A scan that stopped at the outer stage would call + that transformer unreferenced and never free it, so the flag would quietly + deliver less than it promises on those pipelines. + """ + stack: list[Any] = [stage] + visited: set[int] = set() + while stack: + obj = stack.pop() + # is_lazy_module first: isinstance() on a proxy forwards __class__ and + # would load every deferred component just to work out where to free it. + if is_lazy_module(obj): + yield obj + continue + if isinstance(obj, PipelineStage | list | tuple | dict): + if id(obj) in visited: + continue + visited.add(id(obj)) + if isinstance(obj, PipelineStage): + for key, value in vars(obj).items(): + if key == "_lazy_modules_to_release": + # Installed by this schedule, not a real use. + continue + stack.append(value) + elif isinstance(obj, list | tuple): + stack.extend(obj) + elif isinstance(obj, dict): + stack.extend(obj.values()) + + class ComposedPipelineBase(ABC): """ Base class for pipelines composed of multiple stages. @@ -51,6 +88,27 @@ class ComposedPipelineBase(ABC): trainable_transformer_names: list[str] = ["transformer"] trainable_transformer_modules: dict[str, torch.nn.Module] = {} post_init_called: bool = False + # Set once the deferred-release schedule has been derived from the stage + # list, so a stage added afterwards can rebuild it instead of running + # against a plan that predates it. + _lazy_release_hooks_installed: bool = False + # Components this pipeline allows ``lazy_module_load`` to defer and free. + # Empty by default: deferral is opt-in per pipeline, because releasing a + # component and loading it again is only safe when nothing outside the + # loader has changed it. Two habits break that and neither raises: + # + # * mutating a component after load. ``LongCatPipeline.initialize_pipeline`` + # turns on block-sparse attention and writes parameters into every + # transformer block. That runs once, so a re-materialized component + # silently comes back with the feature off. + # * reading a component's attributes while building stages. The shared + # ``DenoisingStage.__init__`` derives the attention backend from + # ``transformer.hidden_size``, which materializes the DiT before the + # first request and defeats the deferral it was meant to gain. + # + # A pipeline opts in by listing the components it has checked. Names match + # the diffusers manifest. + _lazy_module_names: tuple[str, ...] = () @classmethod def get_hf_download_component_dirs(cls) -> tuple[str, ...] | None: @@ -144,29 +202,27 @@ def _compile_with_conditions( compiled_count += 1 return compiled_count - def _maybe_compile_pipeline_module( - self, + @staticmethod + def _compile_pipeline_module_instance( module_name: str, + module: torch.nn.Module, fsdp_module_cls: type | None, compile_kwargs: dict[str, Any], - ) -> None: - if module_name not in self.modules: - return - - module = self.modules[module_name] + ) -> Any: + """Apply pipeline-level compile setup to one loaded component.""" if fsdp_module_cls is not None and isinstance(module, fsdp_module_cls): logger.info( "%s is already FSDP-wrapped; skipping torch.compile in pipeline", module_name.capitalize(), ) - return + return module prepare_for_compile = getattr(module, "prepare_for_compile", None) if callable(prepare_for_compile): logger.info("Running prepare_for_compile for %s", module_name) prepare_for_compile() - compiled_count = self._compile_with_conditions(module, compile_kwargs) + compiled_count = ComposedPipelineBase._compile_with_conditions(module, compile_kwargs) if compiled_count > 0: logger.info( "Enabled torch.compile for %d submodules in %s via _compile_conditions with kwargs=%s", @@ -174,108 +230,167 @@ def _maybe_compile_pipeline_module( module_name, compile_kwargs, ) - return + return module # Backward-compatible fallback: compile full module if no condition matched. logger.info("Enabling torch.compile for %s with kwargs=%s", module_name, compile_kwargs) - self.modules[module_name] = torch.compile(module, **compile_kwargs) + return torch.compile(module, **compile_kwargs) - def post_init(self) -> None: - assert self.fastvideo_args is not None, "fastvideo_args must be set" - if self.post_init_called: + def _maybe_compile_pipeline_module( + self, + module_name: str, + fsdp_module_cls: type | None, + compile_kwargs: dict[str, Any], + ) -> None: + if module_name not in self.modules: return - self.post_init_called = True + + entry = self.modules[module_name] + if is_lazy_module(entry): + # Compilation is part of materialization, not a one-time mutation + # of the first loaded instance. The proxy remains in the module + # map, so whole-module and conditional compile both survive every + # release/reload cycle without making initialization eager. + entry.set_materialize_transform( + partial( + ComposedPipelineBase._compile_pipeline_module_instance, + module_name, + fsdp_module_cls=fsdp_module_cls, + compile_kwargs=dict(compile_kwargs), + )) + logger.info("Configured torch.compile for every materialization of deferred %s", module_name) + return + + self.modules[module_name] = self._compile_pipeline_module_instance( + module_name, + entry, + fsdp_module_cls, + compile_kwargs, + ) + + def _apply_inference_compile(self, module_names: tuple[str, ...] | None = None) -> None: + """Attach pipeline-level compile to modules that are present now. + + Sequential MiniMax-H3 loads DiT/VAEs after ``post_init``, so this is + also called once those modules appear. Lazy proxies register a + materialize transform and can be configured at ``post_init``. + """ + if self.fastvideo_args is None: + return + compile_requested = any(( + self.fastvideo_args.enable_torch_compile, + self.fastvideo_args.enable_torch_compile_text_encoder, + self.fastvideo_args.enable_torch_compile_vae, + self.fastvideo_args.enable_torch_compile_audio_vae, + )) + if self.fastvideo_args.training_mode and compile_requested: + logger.info("Torch Compile enabled via FSDP loader for training; skipping additional pipeline compile") if self.fastvideo_args.training_mode: - assert isinstance(self.fastvideo_args, TrainingArgs) - self.training_args = self.fastvideo_args - assert self.training_args is not None - self.initialize_training_pipeline(self.training_args) - if self.training_args.log_validation: - self.initialize_validation_pipeline(self.training_args) + return - self.initialize_pipeline(self.fastvideo_args) compile_transformer = self.fastvideo_args.enable_torch_compile - compile_text_encoder = (self.fastvideo_args.enable_torch_compile_text_encoder) + compile_text_encoder = self.fastvideo_args.enable_torch_compile_text_encoder compile_vae = self.fastvideo_args.enable_torch_compile_vae compile_audio_vae = self.fastvideo_args.enable_torch_compile_audio_vae - if (compile_transformer or compile_text_encoder or compile_vae or compile_audio_vae): - if self.fastvideo_args.training_mode: - logger.info("Torch Compile enabled via FSDP loader for training; skipping additional pipeline compile") - else: - fsdp_module_cls = None - try: - from torch.distributed.fsdp import FSDPModule # type: ignore - fsdp_module_cls = FSDPModule - except Exception: # pragma: no cover - FSDP not always available - fsdp_module_cls = None - - global_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs or {}) - dit_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_dit or global_compile_kwargs) - text_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_text_encoder or global_compile_kwargs) - 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", - fsdp_module_cls=fsdp_module_cls, - compile_kwargs=dit_compile_kwargs, - ) - self._maybe_compile_pipeline_module( - module_name="transformer_refine", - fsdp_module_cls=fsdp_module_cls, - compile_kwargs=dit_compile_kwargs, - ) + if not (compile_transformer or compile_text_encoder or compile_vae or compile_audio_vae): + return + + wanted = None if module_names is None else set(module_names) + + def _want(name: str) -> bool: + return wanted is None or name in wanted + + fsdp_module_cls = None + try: + from torch.distributed.fsdp import FSDPModule # type: ignore + fsdp_module_cls = FSDPModule + except Exception: # pragma: no cover - FSDP not always available + fsdp_module_cls = None + + global_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs or {}) + dit_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_dit or global_compile_kwargs) + text_compile_kwargs = (self.fastvideo_args.torch_compile_kwargs_text_encoder or global_compile_kwargs) + 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: + 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 and any(_want(name) for name in ("transformer", "transformer_refine", "transformer_2")): + for name in ("transformer", "transformer_refine", "transformer_2"): + if _want(name): self._maybe_compile_pipeline_module( - module_name="transformer_2", + module_name=name, fsdp_module_cls=fsdp_module_cls, compile_kwargs=dit_compile_kwargs, ) - logger.info("Torch Compile enabled for DiT") + if any(name in self.modules for name in ("transformer", "transformer_refine", "transformer_2")): + logger.info("Torch Compile enabled for DiT") - if compile_text_encoder: - self._maybe_compile_pipeline_module( - module_name="text_encoder", - fsdp_module_cls=fsdp_module_cls, - compile_kwargs=text_compile_kwargs, - ) + if compile_text_encoder and any(_want(name) for name in ("text_encoder", "text_encoder_2")): + for name in ("text_encoder", "text_encoder_2"): + if _want(name): self._maybe_compile_pipeline_module( - module_name="text_encoder_2", + module_name=name, fsdp_module_cls=fsdp_module_cls, compile_kwargs=text_compile_kwargs, ) - logger.info("Torch Compile enabled for text encoder") - - if compile_vae: - self._maybe_compile_pipeline_module( - module_name="vae", - fsdp_module_cls=fsdp_module_cls, - compile_kwargs=vae_compile_kwargs, - ) - logger.info("Torch Compile enabled for VAE") + if any(name in self.modules for name in ("text_encoder", "text_encoder_2")): + logger.info("Torch Compile enabled for text encoder") + + if compile_vae and _want("vae"): + self._maybe_compile_pipeline_module( + module_name="vae", + fsdp_module_cls=fsdp_module_cls, + compile_kwargs=vae_compile_kwargs, + ) + if "vae" in self.modules: + logger.info("Torch Compile enabled for VAE") + + if compile_audio_vae and _want("audio_vae"): + self._maybe_compile_pipeline_module( + module_name="audio_vae", + fsdp_module_cls=fsdp_module_cls, + compile_kwargs=audio_vae_compile_kwargs, + ) + if "audio_vae" in self.modules: + logger.info("Torch Compile enabled for audio VAE") - if compile_audio_vae: - self._maybe_compile_pipeline_module( - module_name="audio_vae", - fsdp_module_cls=fsdp_module_cls, - compile_kwargs=audio_vae_compile_kwargs, - ) - logger.info("Torch Compile enabled for audio VAE") + def post_init(self) -> None: + assert self.fastvideo_args is not None, "fastvideo_args must be set" + if self.post_init_called: + return + self.post_init_called = True + if self.fastvideo_args.training_mode: + assert isinstance(self.fastvideo_args, TrainingArgs) + self.training_args = self.fastvideo_args + assert self.training_args is not None + self.initialize_training_pipeline(self.training_args) + if self.training_args.log_validation: + self.initialize_validation_pipeline(self.training_args) - self._trace_mgr = attach_activation_trace(self.modules.get("transformer")) + self.initialize_pipeline(self.fastvideo_args) + self._apply_inference_compile() + + trace_target = self.modules.get("transformer") + if is_lazy_module(trace_target): + # The hook manager keeps a strong reference to every module it + # wraps, so attaching here would materialize the DiT before the + # first request and pin that instance past any release. + if envs.FASTVIDEO_TRACE_ACTIVATIONS: + logger.warning("Activation trace is not attached to a deferred transformer; " + "turn off lazy_module_load to trace it") + trace_target = None + self._trace_mgr = attach_activation_trace(trace_target) if not self.fastvideo_args.training_mode: logger.info("Creating pipeline stages...") self.create_pipeline_stages(self.fastvideo_args) + if self._lazy_module_load_enabled(self.fastvideo_args) and self._lazy_module_names: + self._install_lazy_release_hooks() + # Warmup NCCL communicators for sequence parallelism to avoid # slow first forward pass due to lazy initialization warmup_sequence_parallel_communication() @@ -336,7 +451,13 @@ def get_module(self, module_name: str, default_value: Any = None) -> Any: return self.modules[module_name] def add_module(self, module_name: str, module: Any): + previous = self.modules.get(module_name) self.modules[module_name] = module + # The release schedule keys proxies by identity. Replacing a deferred + # module (or swapping a proxy for a freshly loaded instance) leaves + # stages holding the old object unless the schedule is rebuilt. + if self._lazy_release_hooks_installed and (is_lazy_module(previous) or is_lazy_module(module)): + self._install_lazy_release_hooks() def _load_config(self, model_path: str) -> dict[str, Any]: revision = getattr(self.fastvideo_args, "revision", None) @@ -486,13 +607,23 @@ def load_modules(self, load_module_name = module_name component_model_path = os.path.join(self.model_path, load_module_name) - module = PipelineComponentLoader.load_module( - module_name=load_module_name, - component_model_path=component_model_path, - transformers_or_diffusers=transformers_or_diffusers, - fastvideo_args=fastvideo_args, - ) - logger.info("Loaded module %s from %s", module_name, component_model_path) + + def load_component(load_module_name: str = load_module_name, + component_model_path: str = component_model_path, + transformers_or_diffusers: str = transformers_or_diffusers) -> Any: + return PipelineComponentLoader.load_module( + module_name=load_module_name, + component_model_path=component_model_path, + transformers_or_diffusers=transformers_or_diffusers, + fastvideo_args=fastvideo_args, + ) + + if self._lazy_module_load_enabled(fastvideo_args) and module_name in self._lazy_module_names: + module = LazyModule(module_name, load_component) + logger.info("Deferred module %s from %s", module_name, component_model_path) + else: + module = load_component() + logger.info("Loaded module %s from %s", module_name, component_model_path) if module_name in modules: logger.warning("Overwriting module %s", module_name) @@ -507,6 +638,82 @@ def load_modules(self, return modules + @staticmethod + def _lazy_module_load_enabled(fastvideo_args: FastVideoArgs) -> bool: + """Deferred loading is inference only; training needs every component.""" + if not fastvideo_args.lazy_module_load: + return False + if fastvideo_args.training_mode: + logger.warning("lazy_module_load is not supported in training mode; loading all modules eagerly") + return False + return True + + def _build_lazy_release_schedule(self) -> dict[int, list[str]]: + """Map each stage index to the deferred modules it is the last user of. + + Derived from what the stages actually hold rather than declared per + pipeline, so a stage added later cannot have its module freed out from + under it. A module no stage references is never released, which is the + safe direction: it stays loaded rather than disappearing mid-run. + """ + lazy_names_by_id = {id(module): name for name, module in self.modules.items() if is_lazy_module(module)} + if not lazy_names_by_id: + return {} + + last_use: dict[str, int] = {} + for index, stage in enumerate(self._stages): + for held in _iter_held_objects(stage): + name = lazy_names_by_id.get(id(held)) + if name is not None: + last_use[name] = index + + schedule: dict[int, list[str]] = {} + for name, index in sorted(last_use.items()): + schedule.setdefault(index, []).append(name) + + unreferenced = sorted(set(lazy_names_by_id.values()) - set(last_use)) + if unreferenced: + logger.info("Deferred modules held by no stage, so never released: %s", unreferenced) + return schedule + + def _install_lazy_release_hooks(self) -> None: + """Tell each stage which deferred modules to free once it returns.""" + if not self._lazy_module_names: + # Unified-memory auto-enable turns the flag on for every pipeline. + # Only opted-in families (currently MiniMax-H3) should log about it. + self._lazy_release_hooks_installed = True + return + schedule = self._build_lazy_release_schedule() + for index, stage in enumerate(self._stages): + stage._lazy_modules_to_release = tuple(self.modules[name] for name in schedule.get(index, ())) + + if not schedule: + # Deferring without releasing still lowers the load-time peak, but + # it is not what the flag promises, so say so rather than let a + # no-op look like a win. + logger.warning( + "lazy_module_load is on but no deferred module is held by a stage, so nothing will be " + "freed mid-run. Pipeline %s may load its modules eagerly or hold them outside its stages.", + type(self).__name__) + self._lazy_release_hooks_installed = True + return + + for index, names in sorted(schedule.items()): + logger.info("Deferred modules to free after stage %d (%s): %s", index, + getattr(self._stages[index], "_pipeline_stage_name", "?"), names) + self._lazy_release_hooks_installed = True + + def _release_all_lazy_modules(self) -> None: + """Free every deferred component that is currently materialized.""" + for module_name, module in self.modules.items(): + if not is_lazy_module(module): + continue + try: + module.release() + except Exception: + # Never let cleanup replace the exception being propagated. + logger.exception("Failed to release deferred module %s", module_name) + def add_stage(self, stage_name: str, stage: PipelineStage): assert self.modules is not None, "No modules are registered" # Preserve the pipeline-unique stage key for structured metrics. @@ -517,6 +724,17 @@ def add_stage(self, stage_name: str, stage: PipelineStage): self._stage_name_mapping[stage_name] = stage setattr(self, stage_name, stage) + if self._lazy_release_hooks_installed: + # The schedule maps each deferred module to its last holder. A + # stage appended afterwards may hold a module an earlier stage has + # already been told to free, which would hand it a released + # component mid-run. Rebuild rather than trust the stale plan. + # H3 sequential load adds denoise stages on the first request; that + # is the designed path, so do not log it as a warning. + logger.debug("Stage %s was added after the deferred-release schedule was built; rebuilding the schedule", + stage_name) + self._install_lazy_release_hooks() + # TODO(will): don't hardcode no_grad @torch.no_grad() def forward( @@ -539,8 +757,17 @@ def forward( # Execute each stage logger.info("Running pipeline stages: %s", self._stage_name_mapping.keys()) # logger.info("Batch: %s", batch) - for stage in self.stages: - batch = stage(batch, fastvideo_args) + try: + for stage in self.stages: + batch = stage(batch, fastvideo_args) + except BaseException: + # A stage's own hook frees only what that stage was the last user + # of. When the run aborts earlier, everything already materialized + # stays for the life of the generator, and the retry a + # memory-constrained caller is most likely to attempt starts from a + # worse position than the request that just failed. + self._release_all_lazy_modules() + raise # Return the output return batch diff --git a/fastvideo/pipelines/lazy_module.py b/fastvideo/pipelines/lazy_module.py new file mode 100644 index 0000000000..c6f6e28b2e --- /dev/null +++ b/fastvideo/pipelines/lazy_module.py @@ -0,0 +1,205 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Deferred loading and release of heavy pipeline modules. + +A pipeline normally materializes every component before the first stage runs, +so peak memory is the sum of all components even though no two of them are +needed at the same moment. On a unified-memory device that sum is charged +against the same pool the activations come from, and a model whose components +individually fit can still fail to load. + +``LazyModule`` turns that sum into a maximum. It stands in for a component, +loads it on first use, and drops it once the last stage holding it has run. +The pipeline decides when to release; this module only owns the proxying and +the load/free mechanics. +""" + +from __future__ import annotations + +import functools +import gc +import inspect +from collections.abc import Callable +from typing import Any, TypeGuard + +import torch + +from fastvideo.logger import init_logger + +logger = init_logger(__name__) + + +def _cuda_allocated_gib() -> float | None: + if not torch.cuda.is_available(): + return None + return torch.cuda.memory_allocated() / 1024**3 + + +class LazyModule: + """A stand-in for a pipeline component that loads on first use. + + Every attribute access, call, and ``isinstance`` check forwards to the real + component, materializing it if needed. ``release`` drops the reference and + frees the allocator cache; a later access re-runs the loader, so releasing + early is a latency cost, never a correctness one. + """ + + __slots__ = ("_lazy_name", "_lazy_loader", "_lazy_materialize_transform", "_lazy_module", "_lazy_release_callbacks") + + def __init__(self, name: str, loader: Callable[[], Any]) -> None: + object.__setattr__(self, "_lazy_name", name) + object.__setattr__(self, "_lazy_loader", loader) + object.__setattr__(self, "_lazy_materialize_transform", None) + object.__setattr__(self, "_lazy_module", None) + object.__setattr__(self, "_lazy_release_callbacks", []) + + @property + def lazy_name(self) -> str: + return object.__getattribute__(self, "_lazy_name") + + @property + def is_materialized(self) -> bool: + return object.__getattribute__(self, "_lazy_module") is not None + + def materialize(self) -> Any: + """Return the real component, loading it if this is the first use.""" + module = object.__getattribute__(self, "_lazy_module") + if module is not None: + return module + + name = object.__getattribute__(self, "_lazy_name") + loader = object.__getattribute__(self, "_lazy_loader") + logger.info("Loading deferred module %s", name) + module = loader() + if module is None: + raise ValueError(f"Deferred loader for module {name} returned None") + + transform = object.__getattribute__(self, "_lazy_materialize_transform") + if transform is not None: + module = transform(module) + if module is None: + raise ValueError(f"Materialize transform for module {name} returned None") + object.__setattr__(self, "_lazy_module", module) + + allocated = _cuda_allocated_gib() + if allocated is not None: + logger.info("Loaded deferred module %s, cuda allocated now %.2f GiB", name, allocated) + return module + + def set_materialize_transform(self, transform: Callable[[Any], Any]) -> None: + """Apply ``transform`` to this and every future loaded instance. + + Registering a transform does not itself load a deferred component. If + something has already materialized the component, transform that + instance immediately so current and future instances have the same + setup. A transform may return a wrapper, as ``torch.compile`` does. + """ + current_transform = object.__getattribute__(self, "_lazy_materialize_transform") + module = object.__getattribute__(self, "_lazy_module") + if current_transform is not None: + inner = current_transform + + def chained(loaded: Any) -> Any: + return transform(inner(loaded)) + + stored: Callable[[Any], Any] = chained + # The resident instance already ran ``inner``; only apply the new outer. + immediate = transform + else: + stored = transform + immediate = transform + + transformed = immediate(module) if module is not None else None + if module is not None and transformed is None: + raise ValueError(f"Materialize transform for module {self.lazy_name} returned None") + + object.__setattr__(self, "_lazy_materialize_transform", stored) + if module is not None: + object.__setattr__(self, "_lazy_module", transformed) + + def add_release_callback(self, callback: Callable[[], None]) -> None: + """Run ``callback`` each time the real component is dropped.""" + object.__getattribute__(self, "_lazy_release_callbacks").append(callback) + + def release(self) -> bool: + """Drop the real component. Returns True if something was released.""" + module = object.__getattribute__(self, "_lazy_module") + if module is None: + return False + + name = object.__getattribute__(self, "_lazy_name") + before = _cuda_allocated_gib() + object.__setattr__(self, "_lazy_module", None) + del module + for callback in list(object.__getattribute__(self, "_lazy_release_callbacks")): + callback() + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + + after = _cuda_allocated_gib() + if before is not None and after is not None: + logger.info("Released deferred module %s, cuda allocated %.2f -> %.2f GiB, freed %.2f GiB", name, before, + after, before - after) + else: + logger.info("Released deferred module %s", name) + return True + + # ------------------------------------------------------------------ + # Proxying + # ------------------------------------------------------------------ + + def __getattr__(self, item: str) -> Any: + # __slots__ and the methods above are found by normal lookup, so + # reaching here means the attribute belongs to the real component. + attr = getattr(self.materialize(), item) + if inspect.ismethod(attr) or inspect.isbuiltin(attr): + return self._preserve_identity(attr) + return attr + + def _preserve_identity(self, method: Any) -> Any: + """Return the proxy, not the component, from self-returning methods. + + ``nn.Module.to`` and its relatives return ``self``, and callers write + ``self.vae = self.vae.to(device)`` all over the stages. Handing back + the real component there would quietly replace the proxy with a strong + reference the pipeline cannot release, and the run would look normal + while freeing nothing. + """ + + @functools.wraps(method) + def wrapper(*args: Any, **kwargs: Any) -> Any: + result = method(*args, **kwargs) + if result is object.__getattribute__(self, "_lazy_module"): + return self + return result + + return wrapper + + def __setattr__(self, item: str, value: Any) -> None: + setattr(self.materialize(), item, value) + + def __delattr__(self, item: str) -> None: + delattr(self.materialize(), item) + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + return self.materialize()(*args, **kwargs) + + @property # type: ignore[misc] + def __class__(self) -> type: # type: ignore[override] + # isinstance() consults __class__ when the exact type does not match, + # so forwarding it keeps `isinstance(module, FSDPModule)` and friends + # honest. The cost is that an isinstance check materializes; a wrong + # answer would be worse, because callers branch on it silently. + return type(self.materialize()) + + def __repr__(self) -> str: + # Deliberately does not materialize: logging a pipeline must not + # trigger a multi-gigabyte load. + name = object.__getattribute__(self, "_lazy_name") + state = "materialized" if object.__getattribute__(self, "_lazy_module") is not None else "deferred" + return f"" + + +def is_lazy_module(obj: Any) -> TypeGuard[LazyModule]: + """Type test that does not materialize, unlike ``isinstance``.""" + return type(obj) is LazyModule diff --git a/fastvideo/pipelines/lora_pipeline.py b/fastvideo/pipelines/lora_pipeline.py index b291d6e535..3dbeae6451 100644 --- a/fastvideo/pipelines/lora_pipeline.py +++ b/fastvideo/pipelines/lora_pipeline.py @@ -26,6 +26,7 @@ from fastvideo.models.loader.lora_patch import DenseLoRAPatch, normalize_lora_key from fastvideo.models.loader.utils import get_param_names_mapping from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase +from fastvideo.pipelines.lazy_module import is_lazy_module from fastvideo.utils import maybe_download_lora logger = init_logger(__name__) @@ -127,6 +128,7 @@ def __init__(self, *args, **kwargs) -> None: # Adapter tensors and wrapped model layers belong to this pipeline's module # instances. Sharing either cache across two generators can apply one model's # adapter to another model's layers. + self.trainable_transformer_modules = {} self.lora_adapters = defaultdict(dict) self.lora_adapter_paths = {} self.lora_layers = {} @@ -154,11 +156,6 @@ def __init__(self, *args, **kwargs) -> None: self.trainable_transformer_modules.keys(), ) - for ( - transformer_name, - transformer_module, - ) in self.trainable_transformer_modules.items(): - self.exclude_lora_layers[transformer_name] = (transformer_module.config.arch_config.exclude_lora_layers) # Only override the pipeline class's own default when the caller actually set # one. Assigning unconditionally erases per-model defaults, and a model that # declares one usually does so because wrapping every linear breaks its forward. @@ -208,15 +205,16 @@ def __init__(self, *args, **kwargs) -> None: # Inference elif not self.training_mode and self.lora_path is not None: self.convert_to_lora_layers() - self._setting_constructor_adapter = True - try: - self.set_lora_adapter( - self.lora_nickname, # type: ignore - self.lora_path, - strength=self.lora_strength, - ) # type: ignore - finally: - self._setting_constructor_adapter = False + if not any(is_lazy_module(module) for module in self.trainable_transformer_modules.values()): + self._setting_constructor_adapter = True + try: + self.set_lora_adapter( + self.lora_nickname, # type: ignore + self.lora_path, + strength=self.lora_strength, + ) # type: ignore + finally: + self._setting_constructor_adapter = False def is_target_layer(self, module_name: str) -> bool: if self.lora_target_modules is None: @@ -254,6 +252,91 @@ def set_lora_grads(lora_layers: LoRAModelLayers, device_mesh: DeviceMesh): else: raise ValueError(f"Transformer {transformer_name} should be trainable but not found in lora_layers") + def _exclude_lora_layers_for(self, transformer_name: str, transformer_module: Any) -> list[str]: + excluded = self.exclude_lora_layers.get(transformer_name) + if excluded is not None: + return excluded + # Prefer the pipeline config so a LazyModule is not materialized just to + # read a list of layer name fragments. + dit_config = getattr(getattr(self.fastvideo_args, "pipeline_config", None), "dit_config", None) + arch = getattr(dit_config, "arch_config", None) + if arch is not None and hasattr(arch, "exclude_lora_layers"): + excluded = list(arch.exclude_lora_layers) + elif is_lazy_module(transformer_module): + excluded = [] + else: + excluded = list(transformer_module.config.arch_config.exclude_lora_layers) + self.exclude_lora_layers[transformer_name] = excluded + return excluded + + def _apply_constructor_adapter(self) -> None: + if self.lora_path is None: + return + self.cur_adapter_name = "" + self.cur_adapter_path = "" + self._setting_constructor_adapter = True + try: + self.set_lora_adapter( + self.lora_nickname, + self.lora_path, + strength=self.lora_strength, + ) + finally: + self._setting_constructor_adapter = False + + def _convert_one_transformer(self, transformer_name: str, transformer_module: nn.Module) -> None: + excluded_lora_layers = self._exclude_lora_layers_for(transformer_name, transformer_module) + # Fresh instance after a lazy rematerialize must not keep the previous + # block mapping — those modules pin the released DiT and never get freed. + block_list = [] + for name, submodule in transformer_module.named_children(): + if isinstance(submodule, nn.ModuleList): + block_list = [(f"{name}.{i}", m) for i, m in enumerate(submodule)] + break + self.lora_layers[transformer_name] = LoRAModelLayers(block_list) + logger.info("Converting %s to LoRA Transformer", transformer_name) + converted_count = 0 + for block_name, block_modules in _named_module_by_prefix( + transformer_module, + list(self.lora_layers[transformer_name].block_mapping), + ): + if block_name is not None and (not self.fastvideo_args.training_mode + and self.fastvideo_args.dit_layerwise_offload): + scope_ctx = _get_hook_ctx(self.lora_layers[transformer_name].block_mapping[block_name]) + else: + scope_ctx = nullcontext() + with scope_ctx: + for name, layer in block_modules: + if not self.is_target_layer(name): + continue + + excluded = False + for exclude_layer in excluded_lora_layers: + if exclude_layer in name: + excluded = True + break + if excluded: + continue + + layer = get_lora_layer( + layer, + lora_rank=self.lora_rank, + lora_alpha=self.lora_alpha, + training_mode=self.training_mode, + ) + if layer is not None: + block_name_split = name.split(".", 2) + if len(block_name_split) > 2: + block_name = (block_name_split[0] + "." + block_name_split[1]) + else: + block_name = None + if (block_name not in self.lora_layers[transformer_name].block_mapping): + block_name = None + self.lora_layers[transformer_name].add_lora_layer(block_name, name, layer) + replace_submodule(transformer_module, name, layer) + converted_count += 1 + logger.info("Converted %d layers to LoRA layers", converted_count) + def convert_to_lora_layers(self) -> None: """ Unified method to convert the transformer to a LoRA transformer. @@ -265,59 +348,22 @@ def convert_to_lora_layers(self) -> None: transformer_name, transformer_module, ) in self.trainable_transformer_modules.items(): - converted_count = 0 - # init bookkeeping structures - if transformer_name not in self.lora_layers: - # get block list - block_list = [] - for name, submodule in transformer_module.named_children(): - if isinstance(submodule, nn.ModuleList): - block_list = [(f"{name}.{i}", m) for i, m in enumerate(submodule)] - break - self.lora_layers[transformer_name] = LoRAModelLayers(block_list) - logger.info("Converting %s to LoRA Transformer", transformer_name) - # scan every module and convert to LoRA layer if applicable - - for block_name, block_modules in _named_module_by_prefix( - transformer_module, - list(self.lora_layers[transformer_name].block_mapping), - ): - if block_name is not None and (not self.fastvideo_args.training_mode - and self.fastvideo_args.dit_layerwise_offload): - scope_ctx = _get_hook_ctx(self.lora_layers[transformer_name].block_mapping[block_name]) - else: - scope_ctx = nullcontext() - with scope_ctx: - for name, layer in block_modules: - if not self.is_target_layer(name): - continue - - excluded = False - for exclude_layer in self.exclude_lora_layers[transformer_name]: - if exclude_layer in name: - excluded = True - break - if excluded: - continue - - layer = get_lora_layer( - layer, - lora_rank=self.lora_rank, - lora_alpha=self.lora_alpha, - training_mode=self.training_mode, - ) - if layer is not None: - block_name_split = name.split(".", 2) - if len(block_name_split) > 2: - block_name = (block_name_split[0] + "." + block_name_split[1]) - else: - block_name = None - if (block_name not in self.lora_layers[transformer_name].block_mapping): - block_name = None - self.lora_layers[transformer_name].add_lora_layer(block_name, name, layer) - replace_submodule(transformer_module, name, layer) - converted_count += 1 - logger.info("Converted %d layers to LoRA layers", converted_count) + if is_lazy_module(transformer_module): + + def _drop_lora_refs(*, _name: str = transformer_name) -> None: + self.lora_layers.pop(_name, None) + self.cur_adapter_name = "" + self.cur_adapter_path = "" + + def _lora_after_load(module: nn.Module, *, _name: str = transformer_name) -> nn.Module: + self._convert_one_transformer(_name, module) + self._apply_constructor_adapter() + return module + + transformer_module.add_release_callback(_drop_lora_refs) + transformer_module.set_materialize_transform(_lora_after_load) + continue + self._convert_one_transformer(transformer_name, transformer_module) def set_lora_adapter(self, lora_nickname: str, diff --git a/fastvideo/pipelines/stages/base.py b/fastvideo/pipelines/stages/base.py index 1c879c3a7c..0f990b256b 100644 --- a/fastvideo/pipelines/stages/base.py +++ b/fastvideo/pipelines/stages/base.py @@ -15,6 +15,7 @@ import fastvideo.envs as envs from fastvideo.fastvideo_args import FastVideoArgs from fastvideo.logger import init_logger +from fastvideo.pipelines.lazy_module import LazyModule from fastvideo.pipelines.pipeline_batch_info import ForwardBatch from fastvideo.pipelines.stages.validators import VerificationResult @@ -35,6 +36,12 @@ class PipelineStage(ABC): for a specific part of the process, such as prompt encoding, latent preparation, etc. """ performance_component_metric: str | None = None + # Deferred modules this stage is the last user of, installed by the + # pipeline under ``lazy_module_load``. Released once __call__ returns. + # Living here rather than in the pipeline's stage loop means a pipeline + # that overrides forward still frees, since __call__ is the one entry + # point subclasses are told not to override. + _lazy_modules_to_release: tuple[LazyModule, ...] = () def verify_input(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> VerificationResult: """ @@ -143,7 +150,34 @@ def __call__( logger.error("Input verification failed for %s: %s", stage_name, str(e)) raise - # Execute the actual stage logic + # Execute the actual stage logic, then optional output verification. + # One BaseException net: KeyboardInterrupt inside verify_output must + # still free this stage's deferred modules (OOM is already an Exception). + try: + result = self._execute(batch, fastvideo_args, stage_key, stage_class_name, stage_name) + if enable_verification: + try: + output_result = self.verify_output(result, fastvideo_args) + self._run_verification(output_result, stage_name, "output") + except Exception as e: + logger.error("Output verification failed for %s: %s", stage_name, str(e)) + raise + except BaseException: + self._release_deferred_modules(stage_name) + raise + + self._release_deferred_modules(stage_name) + return result + + def _execute( + self, + batch: ForwardBatch, + fastvideo_args: FastVideoArgs, + stage_key: str, + stage_class_name: str, + stage_name: str, + ) -> ForwardBatch: + """Run forward, with the optional timing and logging wrapper.""" if envs.FASTVIDEO_STAGE_LOGGING: logger.info("[%s] Starting execution", stage_name) torch.cuda.synchronize() @@ -169,17 +203,24 @@ def __call__( # Direct execution (current behavior) result = self.forward(batch, fastvideo_args) - if enable_verification: - # Post-execution output verification - try: - output_result = self.verify_output(result, fastvideo_args) - self._run_verification(output_result, stage_name, "output") - except Exception as e: - logger.error("Output verification failed for %s: %s", stage_name, str(e)) - raise - return result + def _release_deferred_modules(self, stage_name: str) -> None: + """Free the deferred components this stage is the last user of. + + Called on the way out whether or not the stage succeeded. A stage that + raises after materializing a multi-gigabyte component would otherwise + keep it for the life of the generator, and the retry that a + memory-constrained caller is most likely to attempt would start from a + worse position than the request that just failed. + """ + for lazy_module in self._lazy_modules_to_release: + try: + lazy_module.release() + except Exception: + # Never let cleanup replace the exception being propagated. + logger.exception("Failed to release deferred module after %s", stage_name) + @abstractmethod def forward( self, diff --git a/fastvideo/tests/api/test_parser.py b/fastvideo/tests/api/test_parser.py index e5370cc7bf..9b0f2104ee 100644 --- a/fastvideo/tests/api/test_parser.py +++ b/fastvideo/tests/api/test_parser.py @@ -114,6 +114,7 @@ def test_load_run_config_supports_yaml_roundtrip(tmp_path) -> None: "image_encoder": True, "vae": True, "pin_cpu_memory": True, + "lazy_module_load": None, }, "compile": { "enabled": False, diff --git a/fastvideo/tests/attention/test_vsa_h3_sm100a_route.py b/fastvideo/tests/attention/test_vsa_h3_sm100a_route.py index 904c3ca3ed..41cf6ebf62 100644 --- a/fastvideo/tests/attention/test_vsa_h3_sm100a_route.py +++ b/fastvideo/tests/attention/test_vsa_h3_sm100a_route.py @@ -142,8 +142,8 @@ def test_prepare_for_regional_compile_resolves_supported_route(monkeypatch): assert probe_q.dtype == torch.bfloat16 assert probe_vbs.dtype == torch.int32 assert probe_vbs.tolist() == [64, 64] - assert impl._regional_compile_layer_idx is not None - assert impl._regional_compile_layer_idx.item() == -1 + assert impl._compile_layer_idx is not None + assert impl._compile_layer_idx.item() == -1 def test_prepare_for_regional_compile_env_off_skips_probe(monkeypatch): @@ -156,8 +156,8 @@ def test_prepare_for_regional_compile_env_off_skips_probe(monkeypatch): assert unsupported is not None assert VSA_SM100A_ENV in unsupported + assert impl._compile_layer_idx is not None assert impl._regional_compile_sm100a_enabled is False - assert impl._regional_compile_layer_idx is None assert fake_sm.support_calls == [] @@ -177,8 +177,8 @@ def is_supported(self, q, variable_block_sizes): unsupported = impl.prepare_for_regional_compile(torch.device("cpu")) assert unsupported is not None + assert impl._compile_layer_idx is not None assert impl._regional_compile_sm100a_enabled is False - assert impl._regional_compile_layer_idx is None assert len(warnings) == 1 assert "compatibility route" in warnings[0] @@ -308,6 +308,61 @@ def compile_safe_from_mask(q, k, v, block_map, variable_block_sizes): torch._dynamo.reset() +def test_generic_compile_reuses_graph_across_layer_indices_and_stays_on_triton(monkeypatch): + """Pipeline compile must share one graph without selecting sm_100a.""" + fake_sm = _FakeSm100a(supported=True) + monkeypatch.setattr(vsa_h3, "_sm100a", fake_sm) + monkeypatch.setenv(VSA_SM100A_ENV, "1") + monkeypatch.setattr(vsa_h3, "probe_enabled", lambda: None) + meta = _build_meta(sparsity=0.5, dense_layers=(0, 17), prefix_segments=(64, 64)) + q, k, v = _tiled_qkv(meta) + + def fake_triton(q, k, v, block_map, variable_block_sizes): + del k, v, variable_block_sizes + return q + block_map.all().to(q.dtype), None + + def fail_sm100a(*args, **kwargs): + raise AssertionError("generic torch.compile unexpectedly selected sm_100a") + + monkeypatch.setattr(vsa_h3, "block_sparse_attn_64_bhsd", fake_triton) + monkeypatch.setattr(fake_sm, "is_supported", fail_sm100a) + monkeypatch.setattr(fake_sm, "block_sparse_attn_sm100a", fail_sm100a) + monkeypatch.setattr(fake_sm, "block_sparse_attn_sm100a_from_mask", fail_sm100a) + monkeypatch.setattr(vsa_h3, "_sm100a_unavailable_reason", fail_sm100a) + + implementations = [] + for layer_idx in range(20): + impl = MiniMaxH3VSAImpl( + num_heads=_HEADS, + head_size=_DIM, + causal=False, + softmax_scale=_DIM**-0.5, + prefix=f"transformer_blocks.{layer_idx}.attn", + ) + impl.prepare_for_compile(torch.device("cpu")) + implementations.append(impl) + + compiled_graphs = [] + + def recording_backend(graph_module, _example_inputs): + compiled_graphs.append(graph_module) + return graph_module.forward + + torch._dynamo.reset() + try: + compiled = [torch.compile(impl.forward, backend=recording_backend, fullgraph=True) + for impl in implementations] + with torch.inference_mode(): + for layer_idx, run in enumerate(compiled): + actual = run(q, k, v, None, meta) + expected_delta = 1.0 if layer_idx in meta.dense_layers else 0.0 + torch.testing.assert_close(actual, q + expected_delta, atol=0, rtol=0) + finally: + torch._dynamo.reset() + + assert len(compiled_graphs) == 1 + + def test_default_off_routes_triton(routed, monkeypatch): fake_sm, fake_triton, run, _ = routed monkeypatch.delenv(VSA_SM100A_ENV, raising=False) diff --git a/fastvideo/tests/inference/test_basic_fasth3_profile.py b/fastvideo/tests/inference/test_basic_fasth3_profile.py index af9b28fae4..779d101e20 100644 --- a/fastvideo/tests/inference/test_basic_fasth3_profile.py +++ b/fastvideo/tests/inference/test_basic_fasth3_profile.py @@ -84,6 +84,18 @@ def test_default_all_profile_matches_fastest_contract(tmp_path): assert request.sampling.guidance_scale == 1.0 assert request.sampling.batch_cfg is False assert request.output.output_path == str(tmp_path / "result.mp4") + assert config.engine.offload.lazy_module_load is None + + +def test_lazy_module_load_is_tri_state(): + config = fasth3.build_generator_config(_args("--num-gpus", "1")) + assert config.engine.offload.lazy_module_load is None + + enabled = fasth3.build_generator_config(_args("--lazy-module-load")) + assert enabled.engine.offload.lazy_module_load is True + + disabled = fasth3.build_generator_config(_args("--no-lazy-module-load")) + assert disabled.engine.offload.lazy_module_load is False @pytest.mark.parametrize("num_frames", (124, 243, 345)) diff --git a/fastvideo/tests/inference/test_inference_regional_compile.py b/fastvideo/tests/inference/test_inference_regional_compile.py index ffb5d0b4bd..e3aa224e07 100644 --- a/fastvideo/tests/inference/test_inference_regional_compile.py +++ b/fastvideo/tests/inference/test_inference_regional_compile.py @@ -106,11 +106,15 @@ def test_h3_vsa_probe_degrades_regional_compile_to_eager(monkeypatch) -> None: class _RegionalPrepareProbe: def __init__(self, unsupported: str | None = None) -> None: - self.devices: list[torch.device] = [] + self.compile_devices: list[torch.device] = [] + self.regional_devices: list[torch.device] = [] self.unsupported = unsupported + def prepare_for_compile(self, device: torch.device) -> None: + self.compile_devices.append(device) + def prepare_for_regional_compile(self, device: torch.device) -> str | None: - self.devices.append(device) + self.regional_devices.append(device) return self.unsupported @@ -157,7 +161,8 @@ def test_minimax_h3_prepare_for_compile_resolves_loaded_vsa_gates() -> None: assert [block.attn._gate_compress_active for block in model.transformer_blocks] == [False, True] for block in model.transformer_blocks: impl = block.attn.distributed_attention.attn_impl - assert impl.devices == [] + assert impl.compile_devices == [next(block.parameters()).device] + assert impl.regional_devices == [] def test_training_compile_prepare_does_not_probe_inference_kernel() -> None: @@ -168,7 +173,9 @@ def test_training_compile_prepare_does_not_probe_inference_kernel() -> None: assert reason is None attention = model.transformer_blocks[0].attn assert attention._gate_compress_active is True - assert attention.distributed_attention.attn_impl.devices == [] + impl = attention.distributed_attention.attn_impl + assert impl.compile_devices == [next(model.parameters()).device] + assert impl.regional_devices == [] def test_regional_compile_prepare_prefers_specialized_hook() -> None: @@ -179,7 +186,8 @@ def test_regional_compile_prepare_prefers_specialized_hook() -> None: assert reason is None impl = model.transformer_blocks[0].attn.distributed_attention.attn_impl - assert impl.devices == [expected_device] + assert impl.compile_devices == [expected_device] + assert impl.regional_devices == [expected_device] def test_minimax_h3_prepare_for_regional_compile_does_not_require_quantized_q_weight() -> None: @@ -193,7 +201,8 @@ def test_minimax_h3_prepare_for_regional_compile_does_not_require_quantized_q_we assert reason is None impl = attention.distributed_attention.attn_impl - assert impl.devices == [expected_device] + assert impl.compile_devices == [expected_device] + assert impl.regional_devices == [expected_device] def test_minimax_h3_prepare_for_regional_compile_propagates_backend_rejection() -> None: diff --git a/fastvideo/tests/platforms/test_unified_memory_offload.py b/fastvideo/tests/platforms/test_unified_memory_offload.py index aec74b0960..cc35c30019 100644 --- a/fastvideo/tests/platforms/test_unified_memory_offload.py +++ b/fastvideo/tests/platforms/test_unified_memory_offload.py @@ -104,6 +104,7 @@ def test_discrete_device_finalization_retains_layerwise_precedence(monkeypatch) assert args.text_encoder_cpu_offload is True assert args.image_encoder_cpu_offload is True assert args.vae_cpu_offload is True + assert args.lazy_module_load is False def test_workers_classify_their_own_device(monkeypatch) -> None: @@ -189,3 +190,18 @@ def unsupported_name(device_id): assert args.disable_offload_on_unified_memory() is True assert args.text_encoder_cpu_offload is False + + +def test_unified_device_auto_enables_lazy_module_load(as_unified_cuda) -> None: + args = FastVideoArgs(model_path="unused/for-this-test") + + assert args.lazy_module_load is None + assert args.finalize_device_offload_policy(device_id=6) is True + assert args.lazy_module_load is True + + +def test_explicit_false_lazy_module_load_stays_off_on_unified(as_unified_cuda) -> None: + args = FastVideoArgs(model_path="unused/for-this-test", lazy_module_load=False) + + args.finalize_device_offload_policy(device_id=6) + assert args.lazy_module_load is False diff --git a/fastvideo/tests/stages/test_lazy_module_load.py b/fastvideo/tests/stages/test_lazy_module_load.py new file mode 100644 index 0000000000..57bf8658e7 --- /dev/null +++ b/fastvideo/tests/stages/test_lazy_module_load.py @@ -0,0 +1,933 @@ +# SPDX-License-Identifier: Apache-2.0 +"""Deferred pipeline-module loading and release. + +CPU only. Exercises the proxy contract and the release schedule the pipeline +derives from what its stages hold; no model weights are touched. +""" + +import dataclasses + +import torch +from types import SimpleNamespace + +import pytest + +from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase +from fastvideo.pipelines.lazy_module import LazyModule, is_lazy_module +from fastvideo.pipelines.stages.base import PipelineStage + + +class _Component: + + def __init__(self, tag: str) -> None: + self.tag = tag + + def __call__(self, value: int) -> int: + return value * 2 + + +def _counting_loader(tag: str = "c"): + calls = [] + + def loader(): + calls.append(tag) + return _Component(tag) + + return loader, calls + + +def test_deferred_until_first_use(): + loader, calls = _counting_loader() + module = LazyModule("transformer", loader) + + assert calls == [] + assert not module.is_materialized + assert "deferred" in repr(module) + + assert module.tag == "c" + assert calls == ["c"] + assert module.is_materialized + + +def test_repr_does_not_materialize(): + loader, calls = _counting_loader() + module = LazyModule("transformer", loader) + + repr(module) + f"{module!r}" + + assert calls == [] + + +def test_loads_exactly_once_across_many_accesses(): + loader, calls = _counting_loader() + module = LazyModule("vae", loader) + + module.tag + module.tag + module(3) + + assert calls == ["c"] + + +def test_call_forwards_to_component(): + loader, _ = _counting_loader() + module = LazyModule("vae", loader) + + assert module(21) == 42 + + +def test_setattr_and_delattr_forward_to_component(): + loader, _ = _counting_loader() + module = LazyModule("vae", loader) + + module.tag = "changed" + assert module.materialize().tag == "changed" + + del module.tag + assert not hasattr(module.materialize(), "tag") + + +def test_self_returning_methods_hand_back_the_proxy(): + # Stages write `self.vae = self.vae.to(device)`. Returning the component + # would swap the proxy out and leave nothing releasable, with no error. + import torch + + module = LazyModule("vae", lambda: torch.nn.Linear(2, 2)) + + assert module.to("cpu") is module + assert module.eval() is module + assert module.float() is module + assert module.requires_grad_(False) is module + + +def test_non_self_returning_methods_pass_their_result_through(): + import torch + + module = LazyModule("vae", lambda: torch.nn.Linear(2, 2)) + + assert isinstance(module.state_dict(), dict) + assert module.extra_repr() == "in_features=2, out_features=2, bias=True" + + +def test_callable_submodule_attribute_is_not_wrapped(): + # Only bound methods get the identity wrapper. A callable submodule must + # come back as itself so attribute chains and further calls keep working. + import torch + + inner = torch.nn.Linear(2, 2) + module = LazyModule("vae", lambda: torch.nn.Sequential(inner)) + + assert module.__getattr__("0") is inner + + +def test_isinstance_reports_the_real_class(): + # Callers branch on isinstance (FSDPModule, nn.Module). A proxy that + # answered False would take the wrong branch silently. + loader, _ = _counting_loader() + module = LazyModule("transformer", loader) + + assert isinstance(module, _Component) + assert is_lazy_module(module) + + +def test_is_lazy_module_does_not_materialize(): + loader, calls = _counting_loader() + module = LazyModule("transformer", loader) + + assert is_lazy_module(module) + assert calls == [] + assert not is_lazy_module(_Component("plain")) + + +def test_release_then_reload_is_correct_not_broken(): + loader, calls = _counting_loader() + module = LazyModule("text_encoder", loader) + + first = module.materialize() + assert module.release() is True + assert not module.is_materialized + + second = module.materialize() + assert calls == ["c", "c"] + assert second is not first + assert second.tag == "c" + + +def test_materialize_transform_applies_to_every_loaded_instance_without_loading_eagerly(): + loader, calls = _counting_loader() + transformed = [] + module = LazyModule("vae", loader) + + def transform(component): + transformed.append(component) + component.tag = f"compiled-{component.tag}" + return component + + module.set_materialize_transform(transform) + assert calls == [] + + first = module.materialize() + assert first.tag == "compiled-c" + assert module.release() is True + + second = module.materialize() + assert second.tag == "compiled-c" + assert second is not first + assert calls == ["c", "c"] + assert transformed == [first, second] + + +def test_materialize_transforms_compose_in_registration_order(): + loader, calls = _counting_loader() + module = LazyModule("vae", loader) + tags = [] + + def inner(component): + tags.append("inner") + component.tag = f"inner-{component.tag}" + return component + + def outer(component): + tags.append("outer") + component.tag = f"outer-{component.tag}" + return component + + module.set_materialize_transform(inner) + module.set_materialize_transform(outer) + first = module.materialize() + assert first.tag == "outer-inner-c" + assert tags == ["inner", "outer"] + assert calls == ["c"] + + +def test_release_without_materializing_is_a_noop(): + loader, calls = _counting_loader() + module = LazyModule("text_encoder", loader) + + assert module.release() is False + assert module.release() is False + assert calls == [] + + +def test_loader_returning_none_raises_instead_of_proxying_none(): + module = LazyModule("transformer", lambda: None) + + with pytest.raises(ValueError, match="returned None"): + module.materialize() + + +# ---------------------------------------------------------------------- +# Release schedule +# ---------------------------------------------------------------------- + + +class _EchoStage(PipelineStage): + + def __init__(self, **held): + for name, value in held.items(): + setattr(self, name, value) + + def forward(self, batch, fastvideo_args): + return batch + + +class _FakePipeline(ComposedPipelineBase): + """Just enough pipeline to exercise the schedule; no weights, no loading.""" + + def __init__(self, modules, stages): # deliberately does not call super() + self.modules = modules + self._stages = stages + self._lazy_module_names = tuple(name for name, module in modules.items() if is_lazy_module(module)) + + def create_pipeline_stages(self, fastvideo_args): + raise NotImplementedError + + +def _schedule(modules, stages): + return _FakePipeline(modules, stages)._build_lazy_release_schedule() + + +def _lazy(name): + return LazyModule(name, lambda: _Component(name)) + + +def test_pipeline_compile_is_reapplied_after_lazy_release(monkeypatch): + loads = [] + compile_calls = [] + + class _CompileAwareComponent(torch.nn.Module): + _compile_conditions = (lambda name, module: name == "block", ) + + def __init__(self): + super().__init__() + self.block = torch.nn.Linear(2, 2) + self.prepare_calls = 0 + + def prepare_for_compile(self): + self.prepare_calls += 1 + + def load_component(): + component = _CompileAwareComponent() + loads.append(component) + return component + + def fake_compile(target, **kwargs): + compile_calls.append((target, kwargs)) + return target + + monkeypatch.setattr(torch, "compile", fake_compile) + lazy = LazyModule("vae", load_component) + pipeline = _FakePipeline({"vae": lazy}, []) + + pipeline._maybe_compile_pipeline_module("vae", None, {"mode": "reduce-overhead"}) + assert loads == [] + + first = lazy.materialize() + assert first.prepare_calls == 1 + assert lazy.release() is True + + second = lazy.materialize() + assert second is not first + assert second.prepare_calls == 1 + assert loads == [first, second] + assert len(compile_calls) == 2 + assert [kwargs for _, kwargs in compile_calls] == [ + {"mode": "reduce-overhead"}, + {"mode": "reduce-overhead"}, + ] + + +def test_whole_module_compile_keeps_lazy_proxy_and_recompiles_after_release(monkeypatch): + loads = [] + compile_calls = [] + + class _WholeComponent(torch.nn.Module): + + def forward(self, value): + return value + + class _Compiled: + + def __init__(self, original): + self.original = original + + def load_component(): + component = _WholeComponent() + loads.append(component) + return component + + def fake_compile(target, **kwargs): + compile_calls.append((target, kwargs)) + return _Compiled(target) + + monkeypatch.setattr(torch, "compile", fake_compile) + lazy = LazyModule("text_encoder", load_component) + pipeline = _FakePipeline({"text_encoder": lazy}, []) + + pipeline._maybe_compile_pipeline_module("text_encoder", None, {"dynamic": True}) + assert pipeline.modules["text_encoder"] is lazy + assert loads == [] + + first = lazy.materialize() + assert first.original is loads[0] + assert lazy.release() is True + + second = lazy.materialize() + assert second.original is loads[1] + assert second is not first + assert len(compile_calls) == 2 + + +def test_schedule_releases_after_the_last_stage_that_holds_a_module(): + text_encoder = _lazy("text_encoder") + transformer = _lazy("transformer") + vae = _lazy("vae") + modules = {"text_encoder": text_encoder, "transformer": transformer, "vae": vae, "scheduler": object()} + + stages = [ + _EchoStage(vae=vae), # 0 input prep + _EchoStage(conditioner=text_encoder), # 1 conditioning + _EchoStage(transformer=transformer), # 2 denoising + _EchoStage(vae=vae, transformer=transformer), # 3 decoding + ] + + assert _schedule(modules, stages) == {1: ["text_encoder"], 3: ["transformer", "vae"]} + + +def test_building_the_schedule_does_not_materialize_anything(): + # isinstance() on a proxy forwards __class__, so a careless scan of stage + # attributes would load every deferred module before the run starts and + # silently undo the whole point of deferring. + loaded = [] + + def tracked(name): + return LazyModule(name, lambda: loaded.append(name) or _Component(name)) + + text_encoder, transformer = tracked("text_encoder"), tracked("transformer") + stages = [ + _EchoStage(conditioner=text_encoder, flags=[1, 2], opts={"a": 1}, ref2va=False), + _EchoStage(transformer=transformer), + ] + + _schedule({"text_encoder": text_encoder, "transformer": transformer}, stages) + + assert loaded == [] + + +def test_schedule_ignores_eager_modules(): + transformer = _lazy("transformer") + scheduler = object() + modules = {"transformer": transformer, "scheduler": scheduler} + stages = [_EchoStage(transformer=transformer, scheduler=scheduler)] + + assert _schedule(modules, stages) == {0: ["transformer"]} + + +def test_schedule_finds_modules_held_inside_containers(): + text_encoder = _lazy("text_encoder") + vae = _lazy("vae") + modules = {"text_encoder": text_encoder, "vae": vae} + stages = [ + _EchoStage(text_encoders=[text_encoder]), + _EchoStage(by_name={"vae": vae}), + ] + + assert _schedule(modules, stages) == {0: ["text_encoder"], 1: ["vae"]} + + +def test_unreferenced_module_is_never_released(): + # Safe direction: a module no stage holds stays loaded rather than + # disappearing under a caller the schedule cannot see. + orphan = _lazy("image_encoder") + + assert _schedule({"image_encoder": orphan}, [_EchoStage(other=1)]) == {} + + +def test_schedule_is_empty_without_lazy_modules(): + assert _schedule({"vae": object()}, [_EchoStage(vae=object())]) == {} + + +# ---------------------------------------------------------------------- +# Enablement +# ---------------------------------------------------------------------- + + +@pytest.mark.parametrize(("lazy", "training", "expected"), [ + (False, False, False), + (True, False, True), + (True, True, False), + (False, True, False), + (None, False, False), +]) +def test_training_mode_never_defers(lazy, training, expected): + args = SimpleNamespace(lazy_module_load=lazy, training_mode=training) + + assert ComposedPipelineBase._lazy_module_load_enabled(args) is expected + + +def test_flag_defaults_to_auto(): + from fastvideo.fastvideo_args import FastVideoArgs + + fields = {f.name: f for f in dataclasses.fields(FastVideoArgs)} + assert fields["lazy_module_load"].default is None + + +# ---------------------------------------------------------------------- +# Release hooks on the stages +# ---------------------------------------------------------------------- + + +def test_hooks_land_on_the_last_stage_that_holds_each_module(): + text_encoder, transformer = _lazy("text_encoder"), _lazy("transformer") + stages = [_EchoStage(conditioner=text_encoder), _EchoStage(transformer=transformer, extra=text_encoder)] + pipeline = _FakePipeline({"text_encoder": text_encoder, "transformer": transformer}, stages) + + pipeline._install_lazy_release_hooks() + + assert stages[0]._lazy_modules_to_release == () + assert set(stages[1]._lazy_modules_to_release) == {text_encoder, transformer} + + +def test_installing_hooks_twice_does_not_shift_the_schedule(): + # The installed tuple is itself a container of proxies; a rebuild that + # counted it as a use would keep pushing every release to the last stage. + text_encoder = _lazy("text_encoder") + stages = [_EchoStage(conditioner=text_encoder), _EchoStage(other=1)] + pipeline = _FakePipeline({"text_encoder": text_encoder}, stages) + + pipeline._install_lazy_release_hooks() + pipeline._install_lazy_release_hooks() + + assert stages[0]._lazy_modules_to_release == (text_encoder, ) + assert stages[1]._lazy_modules_to_release == () + + +def test_stage_call_releases_its_modules(): + loader, calls = _counting_loader() + text_encoder = LazyModule("text_encoder", loader) + stages = [_EchoStage(conditioner=text_encoder), _EchoStage(other=1)] + pipeline = _FakePipeline({"text_encoder": text_encoder}, stages) + pipeline._install_lazy_release_hooks() + + text_encoder.tag # the stage would use it + assert text_encoder.is_materialized + + batch = object() + args = SimpleNamespace(enable_stage_verification=False) + assert stages[0](batch, args) is batch + + assert not text_encoder.is_materialized + assert calls == ["c"] + + +def test_stage_without_hooks_releases_nothing(): + loader, _ = _counting_loader() + module = LazyModule("vae", loader) + stage = _EchoStage(vae=module) + module.tag + + stage(object(), SimpleNamespace(enable_stage_verification=False)) + + assert module.is_materialized + + +def test_pipeline_warns_when_no_stage_holds_a_deferred_module(caplog): + # A silent no-op here would look exactly like a working run, so the flag + # has to say when it cannot do anything. + orphan = _lazy("image_encoder") + pipeline = _FakePipeline({"image_encoder": orphan}, [_EchoStage(other=1)]) + + with caplog.at_level("WARNING"): + pipeline._install_lazy_release_hooks() + + assert "nothing will be freed" in caplog.text + + +def test_empty_opt_in_list_is_silent(caplog): + pipeline = _FakePipeline({}, [_EchoStage(other=1)]) + with caplog.at_level("WARNING"): + pipeline._install_lazy_release_hooks() + assert caplog.text == "" + + +def test_a_stage_that_rebinds_through_to_can_still_be_released(): + # The end-to-end shape of the identity rule: a stage does the + # `self.vae = self.vae.to(device)` dance, the pipeline still releases. + import torch + + vae = LazyModule("vae", lambda: torch.nn.Linear(2, 2)) + stage = _EchoStage(vae=vae) + pipeline = _FakePipeline({"vae": vae}, [stage]) + pipeline._install_lazy_release_hooks() + + stage.vae = stage.vae.to("cpu") + assert stage.vae is vae + assert vae.is_materialized + + stage(object(), SimpleNamespace(enable_stage_verification=False)) + + assert not vae.is_materialized + + +def test_a_stage_added_after_the_schedule_rebuilds_it(caplog): + # The schedule is derived from the stage list. A stage appended afterwards + # could hold a module an earlier stage was already told to free, which + # would hand it a released component mid-run. + vae = _lazy("vae") + first = _EchoStage(vae=vae) + pipeline = _FakePipeline({"vae": vae}, []) + pipeline._stage_name_mapping = {} + pipeline.add_stage("first", first) + pipeline._install_lazy_release_hooks() + + assert first._lazy_modules_to_release == (vae, ) + + later = _EchoStage(vae=vae) + with caplog.at_level("DEBUG"): + pipeline.add_stage("later", later) + + assert "rebuilding the schedule" in caplog.text + assert first._lazy_modules_to_release == () + assert later._lazy_modules_to_release == (vae, ) + + +class _CompositeStage(PipelineStage): + """Mirrors Cosmos25AutoDenoisingStage: the component lives in a child.""" + + def __init__(self, **held): + self._child = _EchoStage(**held) + + def forward(self, batch, fastvideo_args): + return self._child.forward(batch, fastvideo_args) + + +def test_schedule_walks_into_nested_stages(): + # A stage can compose others rather than hold the component itself. Left + # unwalked, the component reads as unreferenced and is never freed. + transformer = _lazy("transformer") + stages = [_EchoStage(other=1), _CompositeStage(transformer=transformer)] + + assert _schedule({"transformer": transformer}, stages) == {1: ["transformer"]} + + +def test_nested_walk_survives_a_cycle(): + vae = _lazy("vae") + outer = _EchoStage(vae=vae) + inner = _EchoStage(back=outer) + outer.inner = inner + + assert _schedule({"vae": vae}, [outer]) == {0: ["vae"]} + + +def test_a_raising_stage_still_releases_its_modules(): + # The retry a memory-constrained caller attempts must not start from a + # worse position than the request that just failed. + class _Boom(_EchoStage): + + def forward(self, batch, fastvideo_args): + raise RuntimeError("out of activation memory") + + vae = LazyModule("vae", lambda: torch.nn.Linear(2, 2)) + stage = _Boom(vae=vae) + pipeline = _FakePipeline({"vae": vae}, [stage]) + pipeline._install_lazy_release_hooks() + vae.materialize() + + with pytest.raises(RuntimeError, match="out of activation memory"): + stage(object(), SimpleNamespace(enable_stage_verification=False)) + + assert not vae.is_materialized + + +def test_a_failing_release_does_not_mask_the_original_error(): + class _Boom(_EchoStage): + + def forward(self, batch, fastvideo_args): + raise RuntimeError("original") + + class _BadRelease(LazyModule): + + def release(self): + raise ValueError("cleanup blew up") + + stage = _Boom(vae=None) + stage._lazy_modules_to_release = (_BadRelease("vae", lambda: object()), ) + + with pytest.raises(RuntimeError, match="original"): + stage(object(), SimpleNamespace(enable_stage_verification=False)) + + +def test_deferral_is_opt_in_per_pipeline(): + # A pipeline that has not been checked must get no deferral at all, + # because releasing and reloading is only safe when nothing outside the + # loader mutates the component or reads it while stages are built. + from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import MiniMaxH3BasePipeline + + assert ComposedPipelineBase._lazy_module_names == () + assert set(MiniMaxH3BasePipeline._lazy_module_names) == {"text_encoder", "transformer", "vae", "audio_vae"} + + +def test_an_aborted_run_releases_everything_already_materialized(): + # A stage frees only what it is the last user of. When the run aborts + # earlier, the rest would stay for the life of the generator. + vae = LazyModule("vae", lambda: torch.nn.Linear(2, 2)) + transformer = LazyModule("transformer", lambda: torch.nn.Linear(2, 2)) + + class _Boom(_EchoStage): + + def forward(self, batch, fastvideo_args): + raise RuntimeError("out of activation memory") + + early = _EchoStage(vae=vae) + boom = _Boom(transformer=transformer) + late = _EchoStage(vae=vae) + pipeline = _FakePipeline({"vae": vae, "transformer": transformer}, [early, boom, late]) + pipeline._install_lazy_release_hooks() + vae.materialize() + transformer.materialize() + + assert vae.is_materialized and transformer.is_materialized + + pipeline._release_all_lazy_modules() + + assert not vae.is_materialized + assert not transformer.is_materialized + + +# ---------------------------------------------------------------------- +# Production wiring +# +# The tests above build stages and pipelines by hand. These two run the real +# code paths where the two defects review found would live: a `load_modules` +# that never reaches the deferral, and a stage constructor that reads a +# component's attributes and materializes it before the first request. +# ---------------------------------------------------------------------- + + +class _StubLoader: + """Stands in for PipelineComponentLoader and counts what it is asked for.""" + + def __init__(self): + self.loaded: list[str] = [] + + def load_module(self, *, module_name, component_model_path, transformers_or_diffusers, fastvideo_args): + self.loaded.append(module_name) + return _Component(module_name) + + +def _run_real_load_modules(monkeypatch, lazy_names, manifest_modules): + from fastvideo.pipelines import composed_pipeline_base as cpb + + stub = _StubLoader() + monkeypatch.setattr(cpb.PipelineComponentLoader, "load_module", stub.load_module) + + class _Pipeline(ComposedPipelineBase): + _required_config_modules = list(manifest_modules) + _lazy_module_names = lazy_names + + def __init__(self): # deliberately does not call super() + self.model_path = "/nowhere" + self.fastvideo_args = None + + def _load_config(self, model_path): + index = {"_class_name": "X", "_diffusers_version": "0"} + index.update({name: ["diffusers", "Cls", {}] for name in manifest_modules}) + return index + + def create_pipeline_stages(self, fastvideo_args): + raise NotImplementedError + + args = SimpleNamespace(lazy_module_load=True, training_mode=False, revision=None) + modules = _Pipeline().load_modules(args) + return modules, stub.loaded + + +def test_real_load_modules_defers_only_the_opted_in_components(monkeypatch): + modules, loaded = _run_real_load_modules(monkeypatch, ("transformer", "vae"), ["transformer", "vae", "scheduler"]) + + assert is_lazy_module(modules["transformer"]) + assert is_lazy_module(modules["vae"]) + assert not is_lazy_module(modules["scheduler"]) + # The loader is asked only for what stays eager. + assert loaded == ["scheduler"] + + +def test_real_load_modules_defers_nothing_when_the_pipeline_opts_out(monkeypatch): + # The base class ships an empty list, so an unchecked pipeline must load + # everything eagerly even with the flag on. + modules, loaded = _run_real_load_modules(monkeypatch, (), ["transformer", "vae", "scheduler"]) + + assert not any(is_lazy_module(m) for m in modules.values()) + assert sorted(loaded) == ["scheduler", "transformer", "vae"] + + +def test_building_the_real_h3_stages_materializes_nothing(): + # `DenoisingStage.__init__` in the shared stage set reads + # `transformer.hidden_size` to pick an attention backend, which would pull + # the DiT in during post_init. H3's stages must not acquire that habit. + from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig + from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import MiniMaxH3Pipeline + from fastvideo.pipelines.composed_pipeline_base import _iter_held_objects + + loaded: list[str] = [] + + def tracked(name): + return LazyModule(name, lambda: loaded.append(name) or _Component(name)) + + pipeline = MiniMaxH3Pipeline.__new__(MiniMaxH3Pipeline) + pipeline._stages = [] + pipeline._stage_name_mapping = {} + pipeline.modules = { + "text_encoder": tracked("text_encoder"), + "transformer": tracked("transformer"), + "vae": tracked("vae"), + "audio_vae": tracked("audio_vae"), + "tokenizer": object(), + "processor": object(), + "scheduler": object(), + "audio_scheduler": object(), + } + args = SimpleNamespace(pipeline_config=MiniMaxH3PipelineConfig()) + + pipeline._add_stages(args, ref2va=False) + + assert loaded == [], f"building stages materialized {loaded}" + assert len(pipeline._stages) == 6 + input_held = {id(obj) for obj in _iter_held_objects(pipeline._stage_name_mapping["input_preparation_stage"])} + latent_held = {id(obj) for obj in _iter_held_objects(pipeline._stage_name_mapping["latent_preparation_stage"])} + decode_held = {id(obj) for obj in _iter_held_objects(pipeline._stage_name_mapping["video_decoding_stage"])} + denoise_held = {id(obj) for obj in _iter_held_objects(pipeline._stage_name_mapping["denoising_stage"])} + assert id(pipeline.modules["vae"]) not in input_held + assert id(pipeline.modules["transformer"]) not in input_held + assert id(pipeline.modules["transformer"]) not in latent_held + assert id(pipeline.modules["transformer"]) not in decode_held + assert id(pipeline.modules["transformer"]) in denoise_held + assert id(pipeline.modules["vae"]) in decode_held + + +def test_h3_lazy_release_drops_dit_before_vae_decode(): + from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig + from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import MiniMaxH3Pipeline + + pipeline = MiniMaxH3Pipeline.__new__(MiniMaxH3Pipeline) + pipeline._stages = [] + pipeline._stage_name_mapping = {} + pipeline.modules = { + "text_encoder": LazyModule("text_encoder", lambda: _Component("text_encoder")), + "transformer": LazyModule("transformer", lambda: _Component("transformer")), + "vae": LazyModule("vae", lambda: _Component("vae")), + "audio_vae": LazyModule("audio_vae", lambda: _Component("audio_vae")), + "tokenizer": object(), + "processor": object(), + "scheduler": object(), + "audio_scheduler": object(), + } + args = SimpleNamespace(pipeline_config=MiniMaxH3PipelineConfig()) + pipeline._add_stages(args, ref2va=False) + schedule = pipeline._build_lazy_release_schedule() + names = {pipeline._stages[index]._pipeline_stage_name: modules for index, modules in schedule.items()} + assert names["conditioning_stage"] == ["text_encoder"] + assert names["denoising_stage"] == ["transformer"] + assert "transformer" not in names.get("video_decoding_stage", []) + assert "vae" in names["video_decoding_stage"] + + +def test_h3_checkpoint_json_updates_dit_patch_size_without_weights(tmp_path): + from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig + from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import _apply_h3_checkpoint_arch_configs + + transformer_dir = tmp_path / "transformer" + transformer_dir.mkdir() + (transformer_dir / "config.json").write_text('{"patch_size": [1, 1, 1]}') + args = SimpleNamespace(pipeline_config=MiniMaxH3PipelineConfig()) + assert tuple(args.pipeline_config.dit_config.patch_size) == (1, 2, 2) + _apply_h3_checkpoint_arch_configs(str(tmp_path), args, {}) + assert tuple(args.pipeline_config.dit_config.patch_size) == (1, 1, 1) + + +def test_h3_checkpoint_json_updates_audio_sampling_rate_without_weights(tmp_path): + from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig + from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import _apply_h3_checkpoint_arch_configs + + audio_dir = tmp_path / "audio_vae" + audio_dir.mkdir() + (audio_dir / "config.json").write_text('{"sampling_rate": 16000, "latent_channels": 16}') + args = SimpleNamespace(pipeline_config=MiniMaxH3PipelineConfig()) + assert int(args.pipeline_config.audio_vae_config.arch_config.sampling_rate) == 32000 + _apply_h3_checkpoint_arch_configs(str(tmp_path), args, {}) + assert int(args.pipeline_config.audio_vae_config.arch_config.sampling_rate) == 16000 + assert int(args.pipeline_config.audio_vae_config.arch_config.latent_channels) == 16 + + +class _LoRAConfigComponent(torch.nn.Module): + + def __init__(self, excluded_layers): + super().__init__() + self.config = SimpleNamespace( + arch_config=SimpleNamespace(exclude_lora_layers=excluded_layers), + ) + self.blocks = torch.nn.ModuleList([torch.nn.Linear(2, 2)]) + + +def _build_stub_lora_pipeline(monkeypatch, transformer, excluded_layers=None): + from fastvideo.pipelines import lora_pipeline as lora_module + + args = SimpleNamespace( + lora_target_modules=None, + lora_path=None, + lora_nickname="default", + lora_strength=1.0, + training_mode=False, + lora_training=False, + dit_layerwise_offload=False, + pipeline_config=SimpleNamespace( + dit_config=SimpleNamespace( + arch_config=SimpleNamespace(exclude_lora_layers=list(excluded_layers or [])), + ), + ), + ) + + def initialize_base(pipeline, *unused_args, **unused_kwargs): + pipeline.fastvideo_args = args + pipeline.modules = {"transformer": transformer} + + monkeypatch.setattr(ComposedPipelineBase, "__init__", initialize_base) + monkeypatch.setattr(lora_module, "get_local_torch_device", lambda: torch.device("cpu")) + + class _Pipeline(lora_module.LoRAPipeline): + + def create_pipeline_stages(self, fastvideo_args): + raise NotImplementedError + + return _Pipeline("unused", args) + + +def test_no_lora_setup_keeps_the_transformer_deferred(monkeypatch): + loaded = [] + transformer = LazyModule( + "transformer", + lambda: loaded.append("transformer") or _LoRAConfigComponent(["proj_out"]), + ) + + pipeline = _build_stub_lora_pipeline(monkeypatch, transformer) + + assert loaded == [] + assert not transformer.is_materialized + assert pipeline.exclude_lora_layers == {} + assert pipeline.trainable_transformer_modules == {"transformer": transformer} + + +def test_lora_conversion_does_not_materialize_a_deferred_dit(monkeypatch): + loaded = [] + transformer = LazyModule( + "transformer", + lambda: loaded.append("transformer") or _LoRAConfigComponent(["proj_out"]), + ) + pipeline = _build_stub_lora_pipeline(monkeypatch, transformer, excluded_layers=["proj_out"]) + + pipeline.convert_to_lora_layers() + + assert loaded == [] + assert not transformer.is_materialized + assert pipeline.exclude_lora_layers == {} + + transformer.materialize() + assert loaded == ["transformer"] + assert transformer.is_materialized + assert pipeline.exclude_lora_layers == {"transformer": ["proj_out"]} + + +def test_lora_release_drops_block_mapping_so_the_dit_can_free(monkeypatch): + import gc + import weakref + + holder = {} + + def load(): + component = _LoRAConfigComponent([]) + holder["component"] = component + return component + + transformer = LazyModule("transformer", load) + pipeline = _build_stub_lora_pipeline(monkeypatch, transformer) + pipeline.convert_to_lora_layers() + transformer.materialize() + ref = weakref.ref(holder["component"]) + del holder["component"] + + assert transformer.release() is True + gc.collect() + assert ref() is None + assert pipeline.lora_layers == {} + + +def test_lora_transformer_bookkeeping_is_per_pipeline(monkeypatch): + first_transformer = LazyModule("transformer", lambda: _LoRAConfigComponent([])) + first = _build_stub_lora_pipeline(monkeypatch, first_transformer) + second_transformer = LazyModule("transformer", lambda: _LoRAConfigComponent([])) + second = _build_stub_lora_pipeline(monkeypatch, second_transformer) + + assert first.trainable_transformer_modules == {"transformer": first_transformer} + assert second.trainable_transformer_modules == {"transformer": second_transformer} + assert first.trainable_transformer_modules is not second.trainable_transformer_modules diff --git a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py index 8f31973789..02955b73fe 100644 --- a/fastvideo/tests/stages/test_minimax_h3_sequential_start.py +++ b/fastvideo/tests/stages/test_minimax_h3_sequential_start.py @@ -116,6 +116,13 @@ def fake_add_denoise(*, ref2va: bool) -> None: assert pipeline.get_module("transformer") is not None assert pipeline._denoise_stages_ready is True + second = pipeline.forward(ForwardBatch(data_type="video", prompt="second clip"), args) + assert second is not None + assert len(loads) == 3 + assert loads[2] == ["text_encoder"] + assert pipeline.get_module("text_encoder") is None + assert condition_stage.conditioner is None + def test_injected_denoise_weights_skip_the_deferred_split(monkeypatch) -> None: events: list = [] @@ -167,7 +174,7 @@ def fake_load(self, fastvideo_args, loaded_modules=None): monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) monkeypatch.setattr("fastvideo.platforms.current_platform.has_unified_memory", lambda device_id: True) - args = FastVideoArgs(model_path="unused/for-this-test") + args = FastVideoArgs(model_path="unused/for-this-test", lazy_module_load=False) MiniMaxH3Pipeline("unused/for-this-test", args) assert loads @@ -175,6 +182,25 @@ def fake_load(self, fastvideo_args, loaded_modules=None): assert all(name not in loads[0] for name in _DENOISE_MODULE_NAMES) +def test_lazy_module_load_owns_deferral_when_both_would_arm(monkeypatch) -> None: + events: list = [] + _patch_pipeline_construction(monkeypatch, events) + loads: list[list[str]] = [] + + def fake_load(self, fastvideo_args, loaded_modules=None): + del fastvideo_args, loaded_modules + loads.append(list(self.required_config_modules)) + return {name: _stub_module(name) for name in self.required_config_modules} + + monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) + monkeypatch.setattr("fastvideo.platforms.current_platform.has_unified_memory", lambda device_id: True) + args = FastVideoArgs(model_path="unused/for-this-test", h3_sequential_load=True) + MiniMaxH3Pipeline("unused/for-this-test", args) + + assert loads == [list(MiniMaxH3Pipeline._required_config_modules)] + assert all(name in loads[0] for name in _DENOISE_MODULE_NAMES) + + def test_auto_loads_together_without_unified_memory(monkeypatch) -> None: events: list = [] _patch_pipeline_construction(monkeypatch, events) @@ -248,3 +274,52 @@ def fake_add_denoise(*, ref2va: bool) -> None: assert "vae" not in loads[1] assert pipeline.get_module("vae") is None assert pipeline.get_module("transformer") is not None + + +def test_generic_pipeline_config_does_not_crash_geometry_overlay(monkeypatch) -> None: + events: list = [] + _patch_pipeline_construction(monkeypatch, events) + + def fake_load(self, fastvideo_args, loaded_modules=None): + del fastvideo_args, loaded_modules + return {name: _stub_module(name) for name in self.required_config_modules} + + monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) + args = FastVideoArgs(model_path="unused/for-this-test", h3_sequential_load=True) + pipeline = MiniMaxH3Pipeline("unused/for-this-test", args) + pipeline.post_init() + assert pipeline.get_module("text_encoder") is not None + + +def test_resident_path_does_not_reread_encoder_on_later_request(monkeypatch) -> None: + events: list = [] + _patch_pipeline_construction(monkeypatch, events) + loads: list[list[str]] = [] + + def fake_load(self, fastvideo_args, loaded_modules=None): + del fastvideo_args + requested = list(self.required_config_modules) + loads.append(requested) + modules = dict(loaded_modules or {}) + for name in requested: + modules.setdefault(name, _stub_module(name)) + return modules + + monkeypatch.setattr(ComposedPipelineBase, "load_modules", fake_load) + args = FastVideoArgs( + model_path="unused/for-this-test", + enable_stage_verification=False, + h3_sequential_load=False, + lazy_module_load=False, + ) + pipeline = MiniMaxH3Pipeline("unused/for-this-test", args) + pipeline.post_init() + passthrough = lambda batch, _args: batch + for stage in pipeline._stages: + monkeypatch.setattr(stage, "forward", passthrough) + + first = pipeline.forward(ForwardBatch(data_type="video", prompt="one"), args) + second = pipeline.forward(ForwardBatch(data_type="video", prompt="two"), args) + assert first is not None and second is not None + assert len(loads) == 1 + assert pipeline.get_module("text_encoder") is not None diff --git a/fastvideo/tests/worker/test_ray_distributed_executor.py b/fastvideo/tests/worker/test_ray_distributed_executor.py new file mode 100644 index 0000000000..0bf8b4d1e3 --- /dev/null +++ b/fastvideo/tests/worker/test_ray_distributed_executor.py @@ -0,0 +1,38 @@ +# SPDX-License-Identifier: Apache-2.0 +from inspect import signature + +from fastvideo.worker.executor import Executor +from fastvideo.worker.ray_distributed_executor import ( + RayDistributedExecutor, + should_use_gloo_loopback, +) + + +def test_ray_executor_implements_executor_abc() -> None: + remaining = getattr(RayDistributedExecutor, "__abstractmethods__", frozenset()) + assert remaining == frozenset(), remaining + + +def test_gloo_loopback_follows_worker_ips_not_node_count() -> None: + assert should_use_gloo_loopback(["192.168.23.2"]) is True + assert should_use_gloo_loopback(["192.168.23.2", "192.168.23.2"]) is True + assert should_use_gloo_loopback(["192.168.23.2", "192.168.23.1"]) is False + + +def test_ray_does_not_copy_per_node_nic_env_vars() -> None: + nic = RayDistributedExecutor.WORKER_LOCAL_NIC_ENV_VARS + assert "NCCL_SOCKET_IFNAME" in nic + assert "NCCL_IB_HCA" in nic + assert "GLOO_SOCKET_IFNAME" in nic + copied = RayDistributedExecutor.ADDITIONAL_ENV_VARS + assert not (nic & copied) + + +def test_ray_log_queue_stays_on_the_driver() -> None: + """multiprocessing.Queue cannot be pickled onto a remote Ray worker.""" + executor = RayDistributedExecutor.__new__(RayDistributedExecutor) + executor.set_log_queue(object()) + assert executor._log_queue is not None + executor.clear_log_queue() + assert executor._log_queue is None + assert "log_queue" in signature(Executor.set_log_queue).parameters diff --git a/fastvideo/worker/ray_distributed_executor.py b/fastvideo/worker/ray_distributed_executor.py index e276bb90c8..2cf2462afa 100644 --- a/fastvideo/worker/ray_distributed_executor.py +++ b/fastvideo/worker/ray_distributed_executor.py @@ -3,6 +3,7 @@ import asyncio from collections import defaultdict +from queue import Queue import os import cloudpickle @@ -11,7 +12,7 @@ from typing import Any, TYPE_CHECKING from collections.abc import Callable -from fastvideo.utils import get_ip, get_distributed_init_method, get_open_port +from fastvideo.utils import get_ip, get_distributed_init_method, get_open_port, get_loopback_ip from fastvideo.fastvideo_args import FastVideoArgs from fastvideo.pipelines.pipeline_batch_info import ForwardBatch from fastvideo.worker.executor import Executor @@ -35,6 +36,18 @@ logger = init_logger(__name__) +def should_use_gloo_loopback(worker_ips: list[str]) -> bool: + """Loopback is only safe when every worker shares one host IP. + + Single-node Ray (one or many GPUs on the same box) can dial the Gloo store + on loopback. Two Sparks already have distinct worker IPs, so this returns + False and Gloo stays on the fabric address. Per-node NIC names are a + separate issue: do not copy ``NCCL_SOCKET_IFNAME`` / ``GLOO_SOCKET_IFNAME`` + from the driver onto those workers. + """ + return len(set(worker_ips)) <= 1 + + @dataclass class RayWorkerMetaData: """ @@ -60,8 +73,27 @@ class RayDistributedExecutor(Executor): "CUDA_VISIBLE_DEVICES", } - # These non-vLLM env vars are copied from the driver to workers - ADDITIONAL_ENV_VARS = {"HF_TOKEN", "HUGGING_FACE_HUB_TOKEN"} + # Per-node fabric names. spark_pair_env.sh / ibdev2netdev can differ across + # boxes; pushing the driver's value overwrites the export set before ray start. + WORKER_LOCAL_NIC_ENV_VARS = { + "NCCL_SOCKET_IFNAME", + "NCCL_IB_HCA", + "GLOO_SOCKET_IFNAME", + } + + # These non-vLLM env vars are copied from the driver to workers. + # NCCL_* knobs present on the driver are added dynamically in + # ``_env_vars_to_copy_from_driver``, except the per-node NIC trio above. + ADDITIONAL_ENV_VARS = { + "HF_TOKEN", + "HUGGING_FACE_HUB_TOKEN", + "NCCL_IB_DISABLE", + "NCCL_P2P_DISABLE", + "NCCL_CUMEM_ENABLE", + "NCCL_NVLS_ENABLE", + "NCCL_DEBUG", + "NCCL_DEBUG_SUBSYS", + } def _init_executor(self) -> None: initialize_ray_cluster(self.fastvideo_args) @@ -196,9 +228,10 @@ def sort_by_driver_then_worker_ip(item: RayWorkerMetaData): } for (node_id, _) in worker_node_and_gpu_ids] # Environment variables to copy from driver to workers + extra_nccl = {k for k in os.environ if k.startswith("NCCL_") and k not in self.WORKER_LOCAL_NIC_ENV_VARS} env_vars_to_copy = get_env_vars_to_copy( - exclude_vars=self.WORKER_SPECIFIC_ENV_VARS, - additional_vars=set(current_platform.additional_env_vars).union(self.ADDITIONAL_ENV_VARS), + exclude_vars=self.WORKER_SPECIFIC_ENV_VARS | self.WORKER_LOCAL_NIC_ENV_VARS, + additional_vars=set(current_platform.additional_env_vars).union(self.ADDITIONAL_ENV_VARS).union(extra_nccl), destination="workers", ) @@ -213,16 +246,8 @@ def sort_by_driver_then_worker_ip(item: RayWorkerMetaData): self._run_ray_workers("update_environment_variables", self._get_env_vars_to_be_updated()) - if len(node_gpus) == 1: - # in single node case, we don't need to get the IP address. - # the loopback address is sufficient - # NOTE: a node may have several IP addresses, one for each - # network interface. `get_ip()` might return any of them, - # while they might not work for communication inside the node - # if the network setup is complicated. Using the loopback address - # solves this issue, as it always works for communication inside - # the node. - driver_ip = "127.0.0.1" + if should_use_gloo_loopback(worker_ips): + driver_ip = get_loopback_ip() distributed_init_method = get_distributed_init_method(driver_ip, get_open_port()) # Initialize the actual workers inside worker wrapper. @@ -341,6 +366,17 @@ def merge_lora_weights(self) -> None: if response["status"] != "lora_adapter_merged": raise RuntimeError(f"Worker {i} failed to merge LoRA weights") + def set_log_queue(self, log_queue: Queue | None) -> None: + """Keep the driver-side queue locally. + + ``multiprocessing.Queue`` is not picklable across Ray nodes, so worker + logs stay in the Ray session log dir instead of being forwarded. + """ + self._log_queue = log_queue + + def clear_log_queue(self) -> None: + self._log_queue = None + def collective_rpc(self, method: str | Callable, timeout: float | None = None, diff --git a/mkdocs.yml b/mkdocs.yml index 8374954bc2..3a67cb01c4 100644 --- a/mkdocs.yml +++ b/mkdocs.yml @@ -152,6 +152,8 @@ nav: - Getting Started: - Installation: getting_started/installation.md - Apple Silicon FastWan: getting_started/installation/mps.md + - NVIDIA DGX Spark: getting_started/installation/spark.md + - Pair two DGX Sparks: getting_started/installation/spark_pair.md - Quick Start: getting_started/quick_start.md - V1 API: getting_started/v1_api.md - Cookbook: