Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
694a15162f | ||
|
|
64df158881 | ||
|
|
d05b59db37 | ||
|
|
8d4e05acef | ||
|
|
4be9192615 |
@@ -0,0 +1,728 @@
|
||||
# MiniMax H3 component disaggregation on two DGX Sparks
|
||||
|
||||
This runbook tests the component-disaggregated MiniMax H3 inference path:
|
||||
|
||||
- Spark A: text/multimodal encoder, video VAE, audio VAE, and the driver.
|
||||
- Spark B: DiT denoising only.
|
||||
- Ray transports CUDA-contiguous conditioning and latent tensors between the two roles.
|
||||
- Each role uses one GPU with tensor parallelism and sequence parallelism disabled.
|
||||
- All role components remain resident until the generator shuts down.
|
||||
|
||||
The first test is a synchronous correctness smoke test. A second test exercises
|
||||
the bounded asynchronous request pipeline.
|
||||
|
||||
## Important: public SSH ports are not Ray node identities
|
||||
|
||||
It is fine for both Sparks to share one public or NAT IP while SSH exposes them
|
||||
through different ports:
|
||||
|
||||
| Purpose | Spark A | Spark B |
|
||||
|---|---|---|
|
||||
| SSH from another machine | PUBLIC_IP:SSH_PORT_A | PUBLIC_IP:SSH_PORT_B |
|
||||
| Ray and FastVideo | unique internal A_IP | unique internal B_IP |
|
||||
| Component role | encoder/decoder | DiT |
|
||||
|
||||
Do not pass PUBLIC_IP, SSH_PORT_A, or SSH_PORT_B to
|
||||
h3_encoder_node_ip or h3_dit_node_ip. The current runtime pins actors using
|
||||
Ray resources named node:A_IP and node:B_IP and rejects equal node addresses.
|
||||
|
||||
Use one of the following for A_IP and B_IP:
|
||||
|
||||
- The two QSFP/RoCE addresses.
|
||||
- Two distinct LAN addresses that are mutually reachable.
|
||||
- Two distinct Tailscale or WireGuard addresses.
|
||||
|
||||
If Ray itself reports the same NodeManagerAddress for both physical machines,
|
||||
stop. Port differences alone cannot select the two component roles. Configure
|
||||
unique internal or overlay addresses before continuing.
|
||||
|
||||
## 1. Connect and identify the internal interfaces
|
||||
|
||||
From the machine used for SSH:
|
||||
|
||||
~~~bash
|
||||
ssh -p SSH_PORT_A USER@PUBLIC_IP
|
||||
ssh -p SSH_PORT_B USER@PUBLIC_IP
|
||||
~~~
|
||||
|
||||
On each Spark:
|
||||
|
||||
~~~bash
|
||||
hostname
|
||||
ip -br addr
|
||||
ibdev2netdev
|
||||
~~~
|
||||
|
||||
The examples below use:
|
||||
|
||||
~~~text
|
||||
Spark A: 192.168.23.1
|
||||
Spark B: 192.168.23.2
|
||||
Ray head: Spark A, port 6379
|
||||
~~~
|
||||
|
||||
Replace these values with the real addresses.
|
||||
|
||||
If the direct QSFP interface does not already have IPv4 addresses, first find
|
||||
the interface name with ibdev2netdev. The commonly used interface is
|
||||
enp1s0f1np1.
|
||||
|
||||
On Spark A:
|
||||
|
||||
~~~bash
|
||||
export H3_QSFP_IFACE=enp1s0f1np1
|
||||
sudo nmcli device set "$H3_QSFP_IFACE" managed no
|
||||
sudo ip addr replace 192.168.23.1/24 dev "$H3_QSFP_IFACE"
|
||||
sudo ip link set "$H3_QSFP_IFACE" mtu 9000 up
|
||||
~~~
|
||||
|
||||
On Spark B:
|
||||
|
||||
~~~bash
|
||||
export H3_QSFP_IFACE=enp1s0f1np1
|
||||
sudo nmcli device set "$H3_QSFP_IFACE" managed no
|
||||
sudo ip addr replace 192.168.23.2/24 dev "$H3_QSFP_IFACE"
|
||||
sudo ip link set "$H3_QSFP_IFACE" mtu 9000 up
|
||||
~~~
|
||||
|
||||
Verify both directions:
|
||||
|
||||
~~~bash
|
||||
# Spark A
|
||||
ping -c 3 192.168.23.2
|
||||
|
||||
# Spark B
|
||||
ping -c 3 192.168.23.1
|
||||
~~~
|
||||
|
||||
These manually assigned addresses do not survive a reboot unless they are
|
||||
added to the persistent network configuration.
|
||||
|
||||
## 2. Put the same FastVideo code and environment on both Sparks
|
||||
|
||||
The component-disaggregation implementation must exist on both machines. A
|
||||
stock checkout that predates this change will not recognize the new flags.
|
||||
|
||||
On both Sparks:
|
||||
|
||||
~~~bash
|
||||
cd /path/to/FastVideo
|
||||
source .venv/bin/activate
|
||||
|
||||
git rev-parse HEAD
|
||||
UV_TORCH_BACKEND=cu130 uv pip install -e .
|
||||
~~~
|
||||
|
||||
If the model requires Hugging Face authentication, export the token before
|
||||
starting Ray so the Ray worker processes inherit it:
|
||||
|
||||
~~~bash
|
||||
export HF_TOKEN=YOUR_TOKEN
|
||||
~~~
|
||||
|
||||
Verify the software and GPU on both machines:
|
||||
|
||||
~~~bash
|
||||
python - <<'PY'
|
||||
import ray
|
||||
import torch
|
||||
from fastvideo.worker.minimax_h3_disaggregated import (
|
||||
MiniMaxH3DisaggregatedExecutor,
|
||||
)
|
||||
|
||||
print("Ray:", ray.__version__)
|
||||
print("Torch:", torch.__version__)
|
||||
print("CUDA:", torch.version.cuda)
|
||||
print("CUDA available:", torch.cuda.is_available())
|
||||
print("GPU:", torch.cuda.get_device_name(0))
|
||||
print("Disaggregated executor import: OK")
|
||||
PY
|
||||
~~~
|
||||
|
||||
Both machines must use the same Ray version and should report NVIDIA GB10.
|
||||
FastVideo currently requires Ray 2.49.1 or newer.
|
||||
|
||||
For local model directories, use the same absolute path on both machines. A
|
||||
Hugging Face model ID can be used instead; each role downloads only its owned
|
||||
component directories plus small required metadata.
|
||||
|
||||
## 3. Start the Ray cluster
|
||||
|
||||
Stop an old cluster first only if it is safe to terminate all Ray work on that
|
||||
machine:
|
||||
|
||||
~~~bash
|
||||
ray stop
|
||||
~~~
|
||||
|
||||
### Spark A: Ray head and encoder/decoder node
|
||||
|
||||
~~~bash
|
||||
cd /path/to/FastVideo
|
||||
source .venv/bin/activate
|
||||
|
||||
export A_IP=192.168.23.1
|
||||
export B_IP=192.168.23.2
|
||||
export FASTVIDEO_HOST_IP="$A_IP"
|
||||
export RAY_memory_monitor_refresh_ms=0
|
||||
export RAY_memory_usage_threshold=1.0
|
||||
export FASTVIDEO_STAGE_LOGGING=1
|
||||
|
||||
ray start \
|
||||
--head \
|
||||
--node-ip-address="$A_IP" \
|
||||
--port=6379 \
|
||||
--num-gpus=1 \
|
||||
--disable-usage-stats \
|
||||
--object-store-memory=2147483648 \
|
||||
--memory=4294967296
|
||||
~~~
|
||||
|
||||
### Spark B: DiT node
|
||||
|
||||
~~~bash
|
||||
cd /path/to/FastVideo
|
||||
source .venv/bin/activate
|
||||
|
||||
export A_IP=192.168.23.1
|
||||
export B_IP=192.168.23.2
|
||||
export FASTVIDEO_HOST_IP="$B_IP"
|
||||
export RAY_memory_monitor_refresh_ms=0
|
||||
export RAY_memory_usage_threshold=1.0
|
||||
export FASTVIDEO_STAGE_LOGGING=1
|
||||
|
||||
ray start \
|
||||
--address="$A_IP:6379" \
|
||||
--node-ip-address="$B_IP" \
|
||||
--num-gpus=1 \
|
||||
--disable-usage-stats \
|
||||
--object-store-memory=2147483648 \
|
||||
--memory=4294967296
|
||||
~~~
|
||||
|
||||
The Ray port 6379 is the head/GCS port. It is unrelated to either SSH port.
|
||||
Ray uses additional internal service ports, so forwarding only SSH and 6379
|
||||
through NAT is not a substitute for a mutually reachable internal network.
|
||||
|
||||
The Ray memory monitor is disabled because GB10 unified-memory usage during
|
||||
large checkpoint loading can otherwise be mistaken for host-memory pressure.
|
||||
|
||||
## 4. Verify Ray sees two distinct one-GPU nodes
|
||||
|
||||
Run on Spark A:
|
||||
|
||||
~~~bash
|
||||
export A_IP=192.168.23.1
|
||||
export RAY_ADDRESS="$A_IP:6379"
|
||||
|
||||
ray status
|
||||
|
||||
python - <<'PY'
|
||||
import os
|
||||
import ray
|
||||
|
||||
ray.init(address=os.environ["RAY_ADDRESS"])
|
||||
|
||||
print("Live nodes:")
|
||||
for node in ray.nodes():
|
||||
if node["Alive"]:
|
||||
print({
|
||||
"node_id": node["NodeID"],
|
||||
"address": node["NodeManagerAddress"],
|
||||
"gpu": node["Resources"].get("GPU", 0),
|
||||
})
|
||||
|
||||
print("Placement resources:")
|
||||
print(sorted(
|
||||
key for key in ray.cluster_resources()
|
||||
if key.startswith("node:")
|
||||
))
|
||||
PY
|
||||
~~~
|
||||
|
||||
Expected placement resources:
|
||||
|
||||
~~~text
|
||||
node:192.168.23.1
|
||||
node:192.168.23.2
|
||||
~~~
|
||||
|
||||
Expected total capacity from ray status:
|
||||
|
||||
~~~text
|
||||
0.0/2.0 GPU
|
||||
~~~
|
||||
|
||||
Do not continue if only one node appears, either GPU is unavailable, or both
|
||||
NodeManagerAddress values are identical.
|
||||
|
||||
## 5. Create the synchronous smoke-test configuration
|
||||
|
||||
Run on Spark A:
|
||||
|
||||
~~~bash
|
||||
export A_IP=192.168.23.1
|
||||
export B_IP=192.168.23.2
|
||||
export H3_MODEL_PATH=MiniMaxAI/MiniMax-H3
|
||||
|
||||
tee /tmp/h3-disaggregated-smoke.yaml >/dev/null <<YAML
|
||||
generator:
|
||||
model_path: "$H3_MODEL_PATH"
|
||||
|
||||
engine:
|
||||
# num_gpus is one per independent component role. The specialized
|
||||
# runtime still reserves two physical GPUs, one on each named node.
|
||||
num_gpus: 1
|
||||
execution_backend: ray
|
||||
use_fsdp_inference: false
|
||||
|
||||
parallelism:
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
|
||||
offload:
|
||||
dit: false
|
||||
dit_layerwise: false
|
||||
text_encoder: false
|
||||
image_encoder: false
|
||||
vae: false
|
||||
pin_cpu_memory: false
|
||||
lazy_module_load: false
|
||||
|
||||
compile:
|
||||
enabled: false
|
||||
text_encoder_enabled: false
|
||||
vae_enabled: false
|
||||
audio_vae_enabled: false
|
||||
|
||||
pipeline:
|
||||
experimental:
|
||||
h3_disaggregated: true
|
||||
h3_encoder_node_ip: "$A_IP"
|
||||
h3_dit_node_ip: "$B_IP"
|
||||
h3_ray_address: "$A_IP:6379"
|
||||
|
||||
h3_sequential_load: false
|
||||
vae_parallel_encode: false
|
||||
vae_parallel_decode: false
|
||||
video_decode_backend: h3-vae
|
||||
|
||||
# Conservative correctness-first backend for the base H3 checkpoint.
|
||||
attention_backend: TORCH_SDPA
|
||||
|
||||
request:
|
||||
prompt: >-
|
||||
A cinematic wide shot of an alpine lake at sunrise, with gentle mist
|
||||
drifting above the water.
|
||||
negative_prompt: ""
|
||||
|
||||
sampling:
|
||||
seed: 2026
|
||||
height: 256
|
||||
width: 256
|
||||
# H3 requires at least five seconds. 124 is the smallest normal legal
|
||||
# frame count satisfying its 17n+5 causal-VAE geometry.
|
||||
num_frames: 124
|
||||
fps: 24
|
||||
# Two sigma points produce one DiT forward, suitable for plumbing checks.
|
||||
num_inference_steps: 2
|
||||
guidance_scale: 1.0
|
||||
batch_cfg: false
|
||||
|
||||
output:
|
||||
output_path: outputs/h3-disaggregated-smoke/
|
||||
save_video: true
|
||||
return_frames: false
|
||||
YAML
|
||||
~~~
|
||||
|
||||
Inspect the resolved file before launching:
|
||||
|
||||
~~~bash
|
||||
sed -n '1,220p' /tmp/h3-disaggregated-smoke.yaml
|
||||
~~~
|
||||
|
||||
Critical settings:
|
||||
|
||||
- num_gpus must be 1, not 2.
|
||||
- tp_size and sp_size must both be 1.
|
||||
- execution_backend must be ray.
|
||||
- h3_encoder_node_ip and h3_dit_node_ip must be distinct Ray node IPs.
|
||||
- FSDP, lazy loading, sequential loading, VAE parallelism, and TAEH3 must be off.
|
||||
- The full H3 video VAE is selected with video_decode_backend: h3-vae.
|
||||
|
||||
## 6. Run the synchronous correctness smoke test
|
||||
|
||||
On Spark A:
|
||||
|
||||
~~~bash
|
||||
cd /path/to/FastVideo
|
||||
source .venv/bin/activate
|
||||
|
||||
export A_IP=192.168.23.1
|
||||
export RAY_ADDRESS="$A_IP:6379"
|
||||
export FASTVIDEO_HOST_IP="$A_IP"
|
||||
|
||||
fastvideo generate --config /tmp/h3-disaggregated-smoke.yaml
|
||||
~~~
|
||||
|
||||
Successful actor startup validates both placement and component isolation.
|
||||
The expected resident module sets are:
|
||||
|
||||
~~~text
|
||||
Spark A / encoder_decoder:
|
||||
audio_vae, processor, scheduler, text_encoder, tokenizer, vae
|
||||
|
||||
Spark B / dit:
|
||||
audio_scheduler, scheduler, transformer
|
||||
~~~
|
||||
|
||||
The output should be written under:
|
||||
|
||||
~~~text
|
||||
outputs/h3-disaggregated-smoke/
|
||||
~~~
|
||||
|
||||
The first startup can be slow because the two large component sets are loaded
|
||||
and, when using a Hugging Face model ID, downloaded into each node's cache.
|
||||
|
||||
## 7. Exercise three-request pipeline concurrency
|
||||
|
||||
This test directly uses the disaggregated executor's bounded lookahead API. It
|
||||
keeps both actors alive, reports their health, and submits three requests.
|
||||
|
||||
Run on Spark A:
|
||||
|
||||
~~~bash
|
||||
cd /path/to/FastVideo
|
||||
source .venv/bin/activate
|
||||
|
||||
export A_IP=192.168.23.1
|
||||
export RAY_ADDRESS="$A_IP:6379"
|
||||
export FASTVIDEO_HOST_IP="$A_IP"
|
||||
|
||||
python - <<'PY'
|
||||
import time
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.pipelines import ForwardBatch
|
||||
|
||||
generator = VideoGenerator.from_file("/tmp/h3-disaggregated-smoke.yaml")
|
||||
|
||||
try:
|
||||
print("Worker health:")
|
||||
for receipt in generator.executor.runtime.health():
|
||||
print(receipt)
|
||||
|
||||
prompts = [
|
||||
"An alpine lake at sunrise with thin mist.",
|
||||
"Ocean waves beneath a cloudy blue sky.",
|
||||
"A quiet forest path illuminated by morning light.",
|
||||
]
|
||||
|
||||
batches = [
|
||||
ForwardBatch(
|
||||
data_type="video",
|
||||
prompt=prompt,
|
||||
negative_prompt="",
|
||||
height=256,
|
||||
width=256,
|
||||
num_frames=124,
|
||||
fps=24,
|
||||
num_inference_steps=2,
|
||||
guidance_scale=1.0,
|
||||
batch_cfg=False,
|
||||
num_videos_per_prompt=1,
|
||||
seed=2026 + index,
|
||||
save_video=False,
|
||||
return_frames=False,
|
||||
)
|
||||
for index, prompt in enumerate(prompts)
|
||||
]
|
||||
|
||||
started = time.perf_counter()
|
||||
for index, output in enumerate(
|
||||
generator.executor.iter_forward(batches),
|
||||
start=1,
|
||||
):
|
||||
audio = output.extra["audio"]
|
||||
print(
|
||||
f"request={index}",
|
||||
f"video={tuple(output.output.shape)}",
|
||||
f"audio={tuple(audio.shape)}",
|
||||
f"sample_rate={output.extra['audio_sample_rate']}",
|
||||
f"elapsed={time.perf_counter() - started:.1f}s",
|
||||
)
|
||||
finally:
|
||||
generator.shutdown()
|
||||
PY
|
||||
~~~
|
||||
|
||||
Expected health properties:
|
||||
|
||||
~~~text
|
||||
role: encoder_decoder
|
||||
node_ip: 192.168.23.1
|
||||
all_resident: True
|
||||
|
||||
role: dit
|
||||
node_ip: 192.168.23.2
|
||||
all_resident: True
|
||||
~~~
|
||||
|
||||
The scheduling order is:
|
||||
|
||||
~~~text
|
||||
Spark A encodes request N+1
|
||||
Spark B denoises request N
|
||||
Spark A decodes request N-1
|
||||
~~~
|
||||
|
||||
Spark A serializes its own encode and decode operations, while those operations
|
||||
overlap with Spark B's independent denoising work. Spark B schedules the next
|
||||
denoise before the driver waits for the previous decode result.
|
||||
|
||||
## 8. Observe utilization and confirm components are not reloaded
|
||||
|
||||
In a separate SSH session on each Spark:
|
||||
|
||||
~~~bash
|
||||
nvidia-smi dmon -s pucvmet
|
||||
~~~
|
||||
|
||||
Ray worker logs are normally under:
|
||||
|
||||
~~~text
|
||||
/tmp/ray/session_latest/logs/
|
||||
~~~
|
||||
|
||||
Inspect relevant lifecycle messages:
|
||||
|
||||
~~~bash
|
||||
grep -R -E \
|
||||
"Loading pipeline modules|Released MiniMax|Reloading MiniMax|component-disaggregated" \
|
||||
/tmp/ray/session_latest/logs/ | tail -n 100
|
||||
~~~
|
||||
|
||||
For one generator lifetime:
|
||||
|
||||
- Each actor should initialize its pipeline once.
|
||||
- No request should release or reload the text encoder, DiT, or VAEs.
|
||||
- Spark B should continue denoising without waiting for Spark A's VAE decode.
|
||||
|
||||
### Measure the two payload handoffs
|
||||
|
||||
Enable transfer profiling on the **driver before creating the generator**:
|
||||
|
||||
~~~bash
|
||||
FASTVIDEO_H3_PROFILE_TRANSFERS=1 \
|
||||
fastvideo generate --config /tmp/h3-disaggregated-smoke.yaml
|
||||
~~~
|
||||
|
||||
The driver passes this setting to both actors, including when connecting to an
|
||||
existing Ray cluster. It also applies to `iter_forward` and its async wrapper.
|
||||
Look for `[RAY_TRANSFER]`, `[RAY_STAGE]`, and `[RAY_PAYLOAD]` in the worker logs.
|
||||
Every line includes a `request_id` to correlate concurrent requests.
|
||||
|
||||
Each request logs an `A_TO_B` encoded-payload receipt on the DiT worker and a
|
||||
`B_TO_A` denoised-payload receipt on the encoder/decoder worker:
|
||||
|
||||
| Field | What it measures |
|
||||
|---|---|
|
||||
| `source_wait_s` | Receiver wait until the producer's result is available somewhere in Ray. Includes any remaining producer work, serialization, and readiness notification. |
|
||||
| `object_fetch_s` | Wait to make that ready object available in the receiver's local object store. Includes Ray transfer/control overhead and any object-store restore; a cached object can be nearly immediate. |
|
||||
| `materialize_s` | Local `ray.get` deserialization and tensor reconstruction, including waiting for CUDA copies to complete. |
|
||||
| `receive_s` | `object_fetch_s + materialize_s`, excluding the upstream wait. |
|
||||
| `tensor_bytes` | Logical bytes in the conditioning/latent tensors and all five layout tensors. Excludes serialization metadata, backing-storage overhead, and network protocol bytes. |
|
||||
| `object_fetch_mb_s` | Logical tensor MB divided by `object_fetch_s`; an effective payload rate, not measured NIC bandwidth. |
|
||||
|
||||
`[RAY_STAGE]` separately reports synchronized `elapsed_s` for `encode`,
|
||||
`denoise`, and `decode`. Large `object_fetch_s` points toward the Ray data path;
|
||||
large `materialize_s` points toward reconstruction/device-copy cost. Large
|
||||
`source_wait_s` alone does not imply a slow network: compare it with the
|
||||
producer's stage time. Actor queueing before a receiving method starts is not
|
||||
part of these receiver timers. Sender serialization occurs after the producer's
|
||||
stage timer ends and may overlap other work.
|
||||
|
||||
The timers use a monotonic clock on the receiving Spark, so they do not require
|
||||
synchronized clocks between machines. Profiling passes a nested ObjectRef and
|
||||
uses `ray.wait(fetch_local=False)`, then `ray.wait(fetch_local=True)`, then
|
||||
`ray.get` to separate the three phases. See Ray's
|
||||
[object-reference behavior](https://docs.ray.io/en/latest/ray-core/objects.html#passing-object-arguments)
|
||||
and [wait semantics](https://docs.ray.io/en/latest/ray-core/api/doc/ray.wait.html).
|
||||
|
||||
This is an opt-in diagnostic: deferring the fetch until receiver entry changes
|
||||
prefetch overlap, and CUDA synchronization adds overhead. Compare a short warmed
|
||||
run with profiling disabled before drawing throughput conclusions. The default
|
||||
execution keeps Ray's automatic argument fetching. These application timings do
|
||||
not isolate pure time on the wire, and CUDA payloads alone do not establish that
|
||||
the transport uses RDMA or GPU-direct transfers.
|
||||
|
||||
For an external CUDA trace, set `FASTVIDEO_NVTX_PROFILE=1` in the Ray worker
|
||||
environment before starting Ray. With transfer profiling enabled, NVTX ranges
|
||||
include `h3.A_TO_B.source_wait`, `h3.A_TO_B.object_fetch`,
|
||||
`h3.A_TO_B.materialize` (and the corresponding `B_TO_A` ranges), plus
|
||||
`h3.encode`, `h3.denoise`, and `h3.decode`.
|
||||
|
||||
## 9. Optional FastH3/VSA test after correctness passes
|
||||
|
||||
The first smoke test intentionally uses the base MiniMax H3 checkpoint with
|
||||
TORCH_SDPA. For the FastH3 VSA preview checkpoint, change:
|
||||
|
||||
~~~yaml
|
||||
generator:
|
||||
model_path: FastVideo/FastVideo-FastH3-4-step-Preview-v1-VSA-DataFree
|
||||
|
||||
pipeline:
|
||||
experimental:
|
||||
attention_backend: VIDEO_SPARSE_ATTN_H3
|
||||
VSA_sparsity: 0.9
|
||||
VSA_tile_size: 64
|
||||
~~~
|
||||
|
||||
Use five sigma points for the four-forward preview schedule:
|
||||
|
||||
~~~yaml
|
||||
request:
|
||||
sampling:
|
||||
num_inference_steps: 5
|
||||
~~~
|
||||
|
||||
On GB10, do not enable the sm_100a FA4 path. Keep compilation disabled until
|
||||
the basic distributed path is known to work.
|
||||
|
||||
### Use FA4 for supported dense attention layers
|
||||
|
||||
The main FastH3 blocks keep `VIDEO_SPARSE_ATTN_H3`. The text refiner and video
|
||||
VAE use dense attention and can select ordinary BF16 FA4 with a compatible
|
||||
CuTe installation that supports GB10. This is separate from the sm_100a
|
||||
FP4/block-sparse kernels above. Installing FA4 alone does not enable it.
|
||||
|
||||
After a small FA4 forward test passes in the environment on each Spark, set
|
||||
the flag on the driver before creating a new generator:
|
||||
|
||||
~~~bash
|
||||
FASTVIDEO_FA4=1 python /tmp/fasth3_req_video.py
|
||||
~~~
|
||||
|
||||
The H3 executor forwards an explicitly set driver `FASTVIDEO_FA4` through both
|
||||
actors' `runtime_env`, before backend imports, including when Ray is already
|
||||
running. A value in `ray_runtime_env.env_vars` takes precedence over the driver
|
||||
shell. The `ray_non_carry_over_env_vars.json` exclusion policy can disable
|
||||
automatic forwarding. Recreate the generator and actors to apply changes.
|
||||
|
||||
Each actor logs `[H3_WORKER]` with its role, node IP, Python executable, and
|
||||
effective `FASTVIDEO_FA4` value before model loading. Check that both show
|
||||
`FASTVIDEO_FA4=1` and the expected Python environment. Dense backend selection
|
||||
then logs `Using FlashAttention-4 backend`; the main transformer also logs its
|
||||
VSA backend. An explicitly requested FA4 backend that cannot import now raises
|
||||
the original import cause instead of silently selecting SDPA. Other layer
|
||||
constraints, such as unsupported head dimensions, can still select SDPA.
|
||||
|
||||
## 10. Troubleshooting
|
||||
|
||||
### Both Sparks share one public IP
|
||||
|
||||
Different SSH ports are only a login mechanism. Use internal LAN, QSFP, or
|
||||
overlay-network addresses for Ray. Do not write PUBLIC_IP:SSH_PORT into either
|
||||
H3 node-IP field.
|
||||
|
||||
### Ray reports the same node IP twice
|
||||
|
||||
The current runtime cannot distinguish those nodes because it pins actors with
|
||||
node:IP resources and validates actual placement by IP. Configure distinct
|
||||
reachable addresses. If that is impossible, the runtime must be extended to
|
||||
accept Ray node IDs or explicit named per-role resources.
|
||||
|
||||
### Ray reports no live node resource
|
||||
|
||||
Confirm that FASTVIDEO_HOST_IP matched --node-ip-address before ray start on
|
||||
each machine:
|
||||
|
||||
~~~bash
|
||||
python - <<'PY'
|
||||
import ray
|
||||
ray.init(address="auto")
|
||||
print(sorted(k for k in ray.cluster_resources() if k.startswith("node:")))
|
||||
PY
|
||||
~~~
|
||||
|
||||
Restart Ray after correcting the environment.
|
||||
|
||||
### An actor remains pending
|
||||
|
||||
Check:
|
||||
|
||||
~~~bash
|
||||
ray status
|
||||
~~~
|
||||
|
||||
Each requested node must have one free GPU. Stop unrelated GPU jobs or stale
|
||||
Ray actors before retrying.
|
||||
|
||||
### Worker dies while loading checkpoint shards
|
||||
|
||||
Verify that these were exported before ray start on both nodes:
|
||||
|
||||
~~~bash
|
||||
export RAY_memory_monitor_refresh_ms=0
|
||||
export RAY_memory_usage_threshold=1.0
|
||||
~~~
|
||||
|
||||
Also confirm that no other large process is consuming the GB10's unified
|
||||
memory.
|
||||
|
||||
### Import or unknown-argument error
|
||||
|
||||
The new implementation is missing from one node or that node is using a
|
||||
different Python environment. On both machines:
|
||||
|
||||
~~~bash
|
||||
which python
|
||||
which ray
|
||||
python -c "from fastvideo.worker.minimax_h3_disaggregated import MiniMaxH3DisaggregatedExecutor; print('OK')"
|
||||
~~~
|
||||
|
||||
### Hugging Face download or authentication error
|
||||
|
||||
Ensure the model is accessible from both nodes. Export HF_TOKEN before starting
|
||||
Ray, or pre-populate each node's Hugging Face cache.
|
||||
|
||||
### TAEH3 validation error
|
||||
|
||||
TAEH3 intentionally omits the full video VAE and is therefore incompatible
|
||||
with this topology's resident video-VAE requirement. Use h3-vae.
|
||||
|
||||
### num_gpus, TP, or SP validation error
|
||||
|
||||
For component disaggregation, use:
|
||||
|
||||
~~~yaml
|
||||
num_gpus: 1
|
||||
tp_size: 1
|
||||
sp_size: 1
|
||||
~~~
|
||||
|
||||
Although two physical GPUs are reserved, each component role is an independent
|
||||
one-GPU process rather than one two-rank model-parallel world.
|
||||
|
||||
## 11. Stop the cluster
|
||||
|
||||
The generator shuts down its two actors, but it does not stop the shared Ray
|
||||
cluster. When testing is complete, run this on both machines:
|
||||
|
||||
~~~bash
|
||||
ray stop
|
||||
~~~
|
||||
|
||||
## Relevant implementation files
|
||||
|
||||
- fastvideo/pipelines/basic/minimax_h3/disaggregated.py
|
||||
- fastvideo/worker/minimax_h3_disaggregated.py
|
||||
- fastvideo/fastvideo_args.py
|
||||
- fastvideo/worker/executor.py
|
||||
- fastvideo/tests/worker/test_minimax_h3_disaggregated.py
|
||||
|
||||
The existing docs/getting_started/installation/spark_pair.md describes the
|
||||
separate sequence-parallel two-Spark mode. Do not copy its num_gpus=2 and
|
||||
sp_size=2 settings into this component-disaggregated mode.
|
||||
@@ -48,6 +48,7 @@ if TYPE_CHECKING:
|
||||
FASTVIDEO_TRACE_STEPS: str = ""
|
||||
FASTVIDEO_SERVER_DEV_MODE: bool = False
|
||||
FASTVIDEO_STAGE_LOGGING: bool = False
|
||||
FASTVIDEO_H3_PROFILE_TRANSFERS: bool = False
|
||||
FASTVIDEO_CFG_GATE_STEP: float = 1.0
|
||||
FASTVIDEO_HOST_IP: str = ""
|
||||
FASTVIDEO_LOOPBACK_IP: str = ""
|
||||
@@ -343,6 +344,11 @@ environment_variables: dict[str, Callable[[], Any]] = {
|
||||
"FASTVIDEO_STAGE_LOGGING":
|
||||
lambda: bool(int(os.getenv("FASTVIDEO_STAGE_LOGGING", "0"))),
|
||||
|
||||
# Opt-in receiver-side Ray fetch/materialization and synchronized stage
|
||||
# timings for component-disaggregated H3. This changes prefetch overlap.
|
||||
"FASTVIDEO_H3_PROFILE_TRANSFERS":
|
||||
lambda: bool(int(os.getenv("FASTVIDEO_H3_PROFILE_TRANSFERS", "0"))),
|
||||
|
||||
# CFG gating fraction for stale-uncond reuse (Adaptive Guidance / LinearAG
|
||||
# variant — Castillo et al. 2023, arXiv:2312.12487). Float in [0, 1].
|
||||
# Interpretation: for step index `i < len(timesteps) * X`, run both
|
||||
|
||||
@@ -184,6 +184,16 @@ class FastVideoArgs:
|
||||
# keeps every component resident.
|
||||
lazy_module_load: bool | None = None
|
||||
|
||||
# Component-disaggregated MiniMax-H3 inference. The encoder/decoder and
|
||||
# DiT roles are each pinned to one Ray node and keep their component set
|
||||
# resident for the lifetime of the worker. ``num_gpus`` remains one here:
|
||||
# it describes each role, while the runtime reserves one GPU on each of
|
||||
# the two named nodes.
|
||||
h3_disaggregated: bool = False
|
||||
h3_encoder_node_ip: str | None = None
|
||||
h3_dit_node_ip: str | None = None
|
||||
h3_ray_address: str | 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
|
||||
@@ -488,7 +498,7 @@ class FastVideoArgs:
|
||||
parser.add_argument(
|
||||
"--distributed-executor-backend",
|
||||
type=str,
|
||||
choices=["mp"],
|
||||
choices=["mp", "ray"],
|
||||
default=FastVideoArgs.distributed_executor_backend,
|
||||
help="The distributed executor backend to use",
|
||||
)
|
||||
@@ -756,6 +766,30 @@ class FastVideoArgs:
|
||||
"Omit for auto (on for unified-memory devices such as GB10; off on discrete GPUs). "
|
||||
"Pass --no-h3-sequential-load to keep the encoder resident for later generate() calls.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--h3-disaggregated",
|
||||
action=StoreBoolean,
|
||||
default=FastVideoArgs.h3_disaggregated,
|
||||
help="Run MiniMax-H3 as persistent encoder/decoder and DiT workers on two Ray nodes.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--h3-encoder-node-ip",
|
||||
type=str,
|
||||
default=FastVideoArgs.h3_encoder_node_ip,
|
||||
help="Ray node IP for the persistent MiniMax-H3 encoder/decoder worker.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--h3-dit-node-ip",
|
||||
type=str,
|
||||
default=FastVideoArgs.h3_dit_node_ip,
|
||||
help="Ray node IP for the persistent MiniMax-H3 DiT worker.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--h3-ray-address",
|
||||
type=str,
|
||||
default=FastVideoArgs.h3_ray_address,
|
||||
help="Ray cluster address for MiniMax-H3 disaggregation. Unset uses RAY_ADDRESS, then auto.",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--video-decode-backend",
|
||||
type=str,
|
||||
@@ -968,6 +1002,8 @@ class FastVideoArgs:
|
||||
if self.hsdp_shard_dim == -1:
|
||||
self.hsdp_shard_dim = self.num_gpus
|
||||
|
||||
self._check_h3_disaggregated_args()
|
||||
|
||||
assert self.sp_size <= self.num_gpus and self.num_gpus % self.sp_size == 0, "num_gpus must >= and be divisible by sp_size"
|
||||
assert self.hsdp_replicate_dim <= self.num_gpus and self.num_gpus % self.hsdp_replicate_dim == 0, "num_gpus must >= and be divisible by hsdp_replicate_dim"
|
||||
assert self.hsdp_shard_dim <= self.num_gpus and self.num_gpus % self.hsdp_shard_dim == 0, "num_gpus must >= and be divisible by hsdp_shard_dim"
|
||||
@@ -990,6 +1026,56 @@ class FastVideoArgs:
|
||||
self.pipeline_config.vae_config.load_encoder = True
|
||||
self.preprocess_config.check_preprocess_config()
|
||||
|
||||
def _check_h3_disaggregated_args(self) -> None:
|
||||
"""Fail early on topology settings incompatible with role isolation."""
|
||||
if not self.h3_disaggregated:
|
||||
return
|
||||
if self.mode != ExecutionMode.INFERENCE:
|
||||
raise ValueError("MiniMax-H3 component disaggregation supports inference mode only, not preprocessing "
|
||||
"or training modes.")
|
||||
if self.distributed_executor_backend != "ray":
|
||||
raise ValueError("MiniMax-H3 component disaggregation requires distributed_executor_backend='ray'.")
|
||||
|
||||
supported_overrides = {None, "MiniMaxH3ModularPipeline", "MiniMaxH3Ref2VAModularPipeline"}
|
||||
if self.override_pipeline_cls_name not in supported_overrides:
|
||||
raise ValueError("MiniMax-H3 component disaggregation does not support pipeline override "
|
||||
f"{self.override_pipeline_cls_name!r}.")
|
||||
# A bare PipelineConfig is unresolved (common for the direct legacy
|
||||
# constructor), so leave model-family resolution to the H3 workers.
|
||||
# A concrete config subclass, however, is an unambiguous early signal.
|
||||
if type(self.pipeline_config) is not PipelineConfig:
|
||||
from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig
|
||||
if not isinstance(self.pipeline_config, MiniMaxH3PipelineConfig):
|
||||
raise ValueError("MiniMax-H3 component disaggregation requires MiniMaxH3PipelineConfig, got "
|
||||
f"{type(self.pipeline_config).__name__}.")
|
||||
|
||||
encoder_ip = (self.h3_encoder_node_ip or "").strip()
|
||||
dit_ip = (self.h3_dit_node_ip or "").strip()
|
||||
if not encoder_ip or not dit_ip:
|
||||
raise ValueError("MiniMax-H3 component disaggregation requires h3_encoder_node_ip and h3_dit_node_ip.")
|
||||
if encoder_ip == dit_ip:
|
||||
raise ValueError("MiniMax-H3 encoder/decoder and DiT workers must use different Ray node IPs.")
|
||||
self.h3_encoder_node_ip = encoder_ip
|
||||
self.h3_dit_node_ip = dit_ip
|
||||
|
||||
if self.h3_ray_address is not None:
|
||||
self.h3_ray_address = self.h3_ray_address.strip()
|
||||
if not self.h3_ray_address:
|
||||
raise ValueError("h3_ray_address must be non-empty when provided.")
|
||||
if self.num_gpus != 1 or self.tp_size != 1 or self.sp_size != 1:
|
||||
raise ValueError("MiniMax-H3 component disaggregation requires num_gpus=1, tp_size=1, and sp_size=1 "
|
||||
"per worker; it reserves one GPU on each of the two Ray nodes.")
|
||||
if self.use_fsdp_inference:
|
||||
raise ValueError("MiniMax-H3 component disaggregation does not support FSDP inference.")
|
||||
if self.h3_sequential_load is True or self.lazy_module_load is True:
|
||||
raise ValueError("MiniMax-H3 component disaggregation keeps role components resident; sequential and "
|
||||
"lazy module loading must be disabled.")
|
||||
if self.vae_parallel_encode or self.vae_parallel_decode:
|
||||
raise ValueError("MiniMax-H3 component disaggregation does not support sequence-parallel VAE execution.")
|
||||
if self.video_decode_backend != "h3-vae":
|
||||
raise ValueError("MiniMax-H3 component disaggregation requires video_decode_backend='h3-vae' so the "
|
||||
"encoder/decoder worker owns the resident video VAE.")
|
||||
|
||||
def _resolve_device_offload_conflicts(self) -> None:
|
||||
"""Resolve offload modes after device-local policy has been applied."""
|
||||
from fastvideo.platforms import current_platform
|
||||
|
||||
@@ -5,11 +5,27 @@ from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import (
|
||||
MiniMaxH3ModularPipeline,
|
||||
MiniMaxH3Ref2VAModularPipeline,
|
||||
)
|
||||
from fastvideo.pipelines.basic.minimax_h3.disaggregated import (
|
||||
MINIMAX_H3_WIRE_SCHEMA_VERSION,
|
||||
MiniMaxH3DenoisedState,
|
||||
MiniMaxH3DiTPipeline,
|
||||
MiniMaxH3EncodedState,
|
||||
MiniMaxH3EncoderDecoderPipeline,
|
||||
MiniMaxH3RefDiTPipeline,
|
||||
MiniMaxH3RefEncoderDecoderPipeline,
|
||||
)
|
||||
from fastvideo.pipelines.basic.minimax_h3.reference import MiniMaxH3Reference
|
||||
|
||||
__all__ = [
|
||||
"EntryClass",
|
||||
"MINIMAX_H3_WIRE_SCHEMA_VERSION",
|
||||
"MiniMaxH3DenoisedState",
|
||||
"MiniMaxH3DiTPipeline",
|
||||
"MiniMaxH3EncodedState",
|
||||
"MiniMaxH3EncoderDecoderPipeline",
|
||||
"MiniMaxH3ModularPipeline",
|
||||
"MiniMaxH3RefDiTPipeline",
|
||||
"MiniMaxH3RefEncoderDecoderPipeline",
|
||||
"MiniMaxH3Ref2VAModularPipeline",
|
||||
"MiniMaxH3Reference",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,357 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Component-disaggregated MiniMax H3 stage runners and wire contracts."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, replace
|
||||
from uuid import uuid4
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.basic.minimax_h3.minimax_h3_pipeline import (
|
||||
MiniMaxH3BasePipeline,
|
||||
_apply_h3_checkpoint_arch_configs,
|
||||
_use_taeh3_t2va,
|
||||
)
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import MiniMaxH3PackedLayout
|
||||
from fastvideo.pipelines.basic.minimax_h3.stages import (
|
||||
MiniMaxH3AudioDecodingStage,
|
||||
MiniMaxH3DenoisingStage,
|
||||
MiniMaxH3LatentPreparationStage,
|
||||
MiniMaxH3VideoDecodingStage,
|
||||
)
|
||||
from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_latent_preparation import MINIMAX_H3_LAYOUT_KEY
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, PipelineLoggingInfo
|
||||
|
||||
MINIMAX_H3_WIRE_SCHEMA_VERSION = 1
|
||||
|
||||
|
||||
def _cuda_tensor(value: torch.Tensor, name: str, dimensions: int) -> torch.Tensor:
|
||||
if not isinstance(value, torch.Tensor) or value.ndim != dimensions:
|
||||
shape = None if not isinstance(value, torch.Tensor) else tuple(value.shape)
|
||||
raise ValueError(f"MiniMax-H3 wire field {name!r} must be a {dimensions}D tensor, got {shape}.")
|
||||
return value.detach().to(device="cuda", non_blocking=True).contiguous()
|
||||
|
||||
|
||||
def _cuda_layout(layout: MiniMaxH3PackedLayout) -> MiniMaxH3PackedLayout:
|
||||
if not isinstance(layout, MiniMaxH3PackedLayout):
|
||||
raise TypeError("MiniMax-H3 encode output is missing its packed layout.")
|
||||
return replace(
|
||||
layout,
|
||||
position_ids=_cuda_tensor(layout.position_ids, "layout.position_ids", 2),
|
||||
token_tags=_cuda_tensor(layout.token_tags, "layout.token_tags", 1),
|
||||
video_indices=_cuda_tensor(layout.video_indices, "layout.video_indices", 1),
|
||||
audio_indices=_cuda_tensor(layout.audio_indices, "layout.audio_indices", 1),
|
||||
text_indices=_cuda_tensor(layout.text_indices, "layout.text_indices", 1),
|
||||
)
|
||||
|
||||
|
||||
def _validate_schema(schema_version: int) -> None:
|
||||
if schema_version != MINIMAX_H3_WIRE_SCHEMA_VERSION:
|
||||
raise ValueError("Unsupported MiniMax-H3 wire schema version "
|
||||
f"{schema_version}; expected {MINIMAX_H3_WIRE_SCHEMA_VERSION}.")
|
||||
|
||||
|
||||
def _raw_latent_shape(batch: ForwardBatch, stage: str) -> tuple[int, int, int, int, int]:
|
||||
shape = batch.raw_latent_shape
|
||||
if shape is None or len(shape) != 5:
|
||||
raise ValueError(f"MiniMax-H3 {stage} must produce a five-dimensional raw latent shape.")
|
||||
return (int(shape[0]), int(shape[1]), int(shape[2]), int(shape[3]), int(shape[4]))
|
||||
|
||||
|
||||
def _validate_cuda_layout(layout: MiniMaxH3PackedLayout) -> None:
|
||||
if not isinstance(layout, MiniMaxH3PackedLayout):
|
||||
raise TypeError("MiniMax-H3 wire payload requires a packed layout.")
|
||||
expected_shapes = {
|
||||
"position_ids": (layout.sequence_length, 3),
|
||||
"token_tags": (layout.sequence_length, ),
|
||||
}
|
||||
for name, tensor in (
|
||||
("position_ids", layout.position_ids),
|
||||
("token_tags", layout.token_tags),
|
||||
("video_indices", layout.video_indices),
|
||||
("audio_indices", layout.audio_indices),
|
||||
("text_indices", layout.text_indices),
|
||||
):
|
||||
if not tensor.is_cuda or not tensor.is_contiguous():
|
||||
raise ValueError(f"MiniMax-H3 wire field layout.{name} must be a contiguous CUDA tensor.")
|
||||
expected = expected_shapes.get(name)
|
||||
if expected is not None and tuple(tensor.shape) != expected:
|
||||
raise ValueError(
|
||||
f"MiniMax-H3 wire field layout.{name} has shape {tuple(tensor.shape)}, expected {expected}.")
|
||||
if expected is None and tensor.ndim != 1:
|
||||
raise ValueError(f"MiniMax-H3 wire field layout.{name} must be one-dimensional.")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MiniMaxH3EncodedState:
|
||||
"""Minimal CUDA payload sent from the encoder/VAE node to the DiT node."""
|
||||
|
||||
request_id: str
|
||||
prompt_embeds: torch.Tensor
|
||||
video_latents: torch.Tensor
|
||||
audio_latents: torch.Tensor
|
||||
layout: MiniMaxH3PackedLayout
|
||||
raw_latent_shape: tuple[int, int, int, int, int]
|
||||
num_inference_steps: int
|
||||
vsa_sparsity: float
|
||||
vsa_mode: str = "exempt"
|
||||
vsa_dense_first_n_steps: int = 0
|
||||
vsa_dense_layers: tuple[int, ...] = ()
|
||||
logging_info: PipelineLoggingInfo | None = None
|
||||
schema_version: int = MINIMAX_H3_WIRE_SCHEMA_VERSION
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
_validate_schema(self.schema_version)
|
||||
if not self.request_id:
|
||||
raise ValueError("MiniMax-H3 wire payload requires a non-empty request_id.")
|
||||
if len(self.raw_latent_shape) != 5 or min(self.raw_latent_shape) <= 0:
|
||||
raise ValueError(f"Invalid MiniMax-H3 raw latent shape: {self.raw_latent_shape}.")
|
||||
if self.num_inference_steps <= 0:
|
||||
raise ValueError("MiniMax-H3 wire payload requires at least one denoising step.")
|
||||
if self.vsa_mode not in {"exempt", "compete"}:
|
||||
raise ValueError(f"vsa_mode must be 'exempt' or 'compete', got {self.vsa_mode!r}.")
|
||||
for name, tensor, dimensions in (
|
||||
("prompt_embeds", self.prompt_embeds, 3),
|
||||
("video_latents", self.video_latents, 2),
|
||||
("audio_latents", self.audio_latents, 2),
|
||||
):
|
||||
if not tensor.is_cuda or not tensor.is_contiguous() or tensor.ndim != dimensions:
|
||||
raise ValueError(f"MiniMax-H3 wire field {name!r} must be a contiguous CUDA {dimensions}D tensor.")
|
||||
_validate_cuda_layout(self.layout)
|
||||
if self.video_latents.shape[0] != self.layout.video_indices.numel():
|
||||
raise ValueError("MiniMax-H3 video latent rows do not match layout.video_indices.")
|
||||
if self.audio_latents.shape[0] != self.layout.audio_indices.numel():
|
||||
raise ValueError("MiniMax-H3 audio latent rows do not match layout.audio_indices.")
|
||||
if self.prompt_embeds.shape[1] != self.layout.text_indices.numel():
|
||||
raise ValueError("MiniMax-H3 prompt embedding rows do not match layout.text_indices.")
|
||||
|
||||
@classmethod
|
||||
def from_batch(cls, batch: ForwardBatch, request_id: str | None = None) -> MiniMaxH3EncodedState:
|
||||
if len(batch.prompt_embeds) != 1 or batch.latents is None or batch.audio_latents is None:
|
||||
raise ValueError("MiniMax-H3 encode must produce one prompt embedding and both latent streams.")
|
||||
layout = batch.extra.get(MINIMAX_H3_LAYOUT_KEY)
|
||||
return cls(
|
||||
request_id=request_id if request_id is not None else uuid4().hex,
|
||||
prompt_embeds=_cuda_tensor(batch.prompt_embeds[0], "prompt_embeds", 3),
|
||||
video_latents=_cuda_tensor(batch.latents, "video_latents", 2),
|
||||
audio_latents=_cuda_tensor(batch.audio_latents, "audio_latents", 2),
|
||||
layout=_cuda_layout(layout),
|
||||
raw_latent_shape=_raw_latent_shape(batch, "encode"),
|
||||
num_inference_steps=int(batch.num_inference_steps),
|
||||
vsa_sparsity=float(batch.VSA_sparsity),
|
||||
vsa_mode=str(batch.extra.get("vsa_mode", "exempt")),
|
||||
vsa_dense_first_n_steps=int(batch.extra.get("vsa_dense_first_n_steps", 0)),
|
||||
vsa_dense_layers=tuple(int(layer) for layer in batch.extra.get("vsa_dense_layers", ())),
|
||||
logging_info=batch.logging_info,
|
||||
)
|
||||
|
||||
def to_batch(self) -> ForwardBatch:
|
||||
# Pickle/Ray reconstruction does not invoke dataclass __post_init__.
|
||||
# Revalidate the schema and CUDA tensor contract at the receiving boundary.
|
||||
self.__post_init__()
|
||||
return ForwardBatch(
|
||||
data_type="video",
|
||||
prompt_embeds=[self.prompt_embeds],
|
||||
latents=self.video_latents,
|
||||
audio_latents=self.audio_latents,
|
||||
raw_latent_shape=self.raw_latent_shape,
|
||||
num_inference_steps=self.num_inference_steps,
|
||||
VSA_sparsity=self.vsa_sparsity,
|
||||
extra={
|
||||
MINIMAX_H3_LAYOUT_KEY: self.layout,
|
||||
"request_id": self.request_id,
|
||||
"vsa_mode": self.vsa_mode,
|
||||
"vsa_dense_first_n_steps": self.vsa_dense_first_n_steps,
|
||||
"vsa_dense_layers": self.vsa_dense_layers,
|
||||
},
|
||||
logging_info=self.logging_info or PipelineLoggingInfo(),
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MiniMaxH3DenoisedState:
|
||||
"""Minimal CUDA payload returned by the DiT node for VAE decoding."""
|
||||
|
||||
request_id: str
|
||||
video_latents: torch.Tensor
|
||||
audio_latents: torch.Tensor
|
||||
layout: MiniMaxH3PackedLayout
|
||||
raw_latent_shape: tuple[int, int, int, int, int]
|
||||
logging_info: PipelineLoggingInfo | None = None
|
||||
schema_version: int = MINIMAX_H3_WIRE_SCHEMA_VERSION
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
_validate_schema(self.schema_version)
|
||||
if not self.request_id:
|
||||
raise ValueError("MiniMax-H3 wire payload requires a non-empty request_id.")
|
||||
if len(self.raw_latent_shape) != 5 or min(self.raw_latent_shape) <= 0:
|
||||
raise ValueError(f"Invalid MiniMax-H3 raw latent shape: {self.raw_latent_shape}.")
|
||||
for name, tensor in (("video_latents", self.video_latents), ("audio_latents", self.audio_latents)):
|
||||
if not tensor.is_cuda or not tensor.is_contiguous() or tensor.ndim != 2:
|
||||
raise ValueError(f"MiniMax-H3 wire field {name!r} must be a contiguous CUDA 2D tensor.")
|
||||
_validate_cuda_layout(self.layout)
|
||||
if self.video_latents.shape[0] != self.layout.video_indices.numel():
|
||||
raise ValueError("MiniMax-H3 video latent rows do not match layout.video_indices.")
|
||||
if self.audio_latents.shape[0] != self.layout.audio_indices.numel():
|
||||
raise ValueError("MiniMax-H3 audio latent rows do not match layout.audio_indices.")
|
||||
|
||||
@classmethod
|
||||
def from_batch(cls, batch: ForwardBatch, request_id: str) -> MiniMaxH3DenoisedState:
|
||||
if batch.latents is None or batch.audio_latents is None:
|
||||
raise ValueError("MiniMax-H3 denoise must return both latent streams.")
|
||||
return cls(
|
||||
request_id=request_id,
|
||||
video_latents=_cuda_tensor(batch.latents, "video_latents", 2),
|
||||
audio_latents=_cuda_tensor(batch.audio_latents, "audio_latents", 2),
|
||||
layout=_cuda_layout(batch.extra.get(MINIMAX_H3_LAYOUT_KEY)),
|
||||
raw_latent_shape=_raw_latent_shape(batch, "denoise"),
|
||||
logging_info=batch.logging_info,
|
||||
)
|
||||
|
||||
def to_batch(self) -> ForwardBatch:
|
||||
# Revalidate after transport; unpickling itself skips __post_init__.
|
||||
self.__post_init__()
|
||||
return ForwardBatch(
|
||||
data_type="video",
|
||||
latents=self.video_latents,
|
||||
audio_latents=self.audio_latents,
|
||||
raw_latent_shape=self.raw_latent_shape,
|
||||
extra={
|
||||
MINIMAX_H3_LAYOUT_KEY: self.layout,
|
||||
"request_id": self.request_id,
|
||||
},
|
||||
logging_info=self.logging_info or PipelineLoggingInfo(),
|
||||
)
|
||||
|
||||
|
||||
class _MiniMaxH3ResidentRolePipeline(MiniMaxH3BasePipeline):
|
||||
"""Base that disables H3's single-worker release/reload lifecycle."""
|
||||
|
||||
_lazy_module_names: tuple[str, ...] = ()
|
||||
|
||||
def _defer_denoise_modules(self, fastvideo_args: FastVideoArgs) -> bool:
|
||||
del fastvideo_args
|
||||
return False
|
||||
|
||||
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
del batch, fastvideo_args
|
||||
raise RuntimeError("Component-disaggregated H3 pipelines expose encode(), denoise(), and decode().")
|
||||
|
||||
|
||||
class MiniMaxH3EncoderDecoderPipeline(_MiniMaxH3ResidentRolePipeline):
|
||||
"""Persistent Qwen + video/audio VAE worker for T2VA and FL2VA."""
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "processor", "vae", "audio_vae", "scheduler"]
|
||||
|
||||
@classmethod
|
||||
def get_hf_download_allow_patterns(cls) -> list[str]:
|
||||
patterns = super().get_hf_download_allow_patterns() or []
|
||||
# Packing uses the checkpoint's DiT patch geometry, but this role must
|
||||
# never download or instantiate DiT weights. The text-encoder loader
|
||||
# also consults transformer/config.json for connector RoPE metadata.
|
||||
geometry_dir = cls._extra_config_module_map.get("transformer", "transformer")
|
||||
metadata = {"transformer/config.json", f"{geometry_dir}/config.json"}
|
||||
return [*patterns, *(pattern for pattern in sorted(metadata) if pattern not in patterns)]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
_apply_h3_checkpoint_arch_configs(self.model_path, fastvideo_args, self._extra_config_module_map)
|
||||
shift = getattr(self.get_module("scheduler"), "shift", None)
|
||||
if shift is None or float(shift) != 12.0:
|
||||
raise ValueError(f"MiniMax-H3 video scheduler must expose shift=12, got {shift}.")
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
ref2va = self._ref2va
|
||||
self._add_condition_stages(fastvideo_args, ref2va=ref2va)
|
||||
vae = self.get_module("vae")
|
||||
audio_vae = self.get_module("audio_vae")
|
||||
scheduler = self.get_module("scheduler")
|
||||
if vae is None or audio_vae is None or scheduler is None:
|
||||
raise RuntimeError("MiniMax-H3 encoder/decoder worker requires both VAEs and the video scheduler.")
|
||||
self.add_stage(
|
||||
"latent_preparation_stage",
|
||||
MiniMaxH3LatentPreparationStage(vae=vae, audio_vae=audio_vae, scheduler=scheduler, ref2va=ref2va),
|
||||
)
|
||||
use_taeh3 = _use_taeh3_t2va(fastvideo_args, ref2va=ref2va)
|
||||
self.add_stage("video_decoding_stage", MiniMaxH3VideoDecodingStage(vae=None if use_taeh3 else vae))
|
||||
self.add_stage("audio_decoding_stage", MiniMaxH3AudioDecodingStage(audio_vae=audio_vae))
|
||||
|
||||
def encode(self, batch: ForwardBatch, request_id: str | None = None) -> MiniMaxH3EncodedState:
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
for name in ("input_preparation_stage", "conditioning_stage", "latent_preparation_stage"):
|
||||
batch = self._stage_name_mapping[name](batch, self.fastvideo_args)
|
||||
return MiniMaxH3EncodedState.from_batch(batch, request_id=request_id)
|
||||
|
||||
def decode(self, state: MiniMaxH3DenoisedState) -> ForwardBatch:
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
batch = state.to_batch()
|
||||
for name in ("video_decoding_stage", "audio_decoding_stage"):
|
||||
batch = self._stage_name_mapping[name](batch, self.fastvideo_args)
|
||||
batch.extra["request_id"] = state.request_id
|
||||
return batch
|
||||
|
||||
|
||||
class MiniMaxH3RefEncoderDecoderPipeline(MiniMaxH3EncoderDecoderPipeline):
|
||||
"""Persistent encoder/decoder worker for Ref2VA requests."""
|
||||
|
||||
_extra_config_module_map = {"transformer": "transformer_ref"}
|
||||
_ref2va_default = True
|
||||
|
||||
|
||||
class MiniMaxH3DiTPipeline(_MiniMaxH3ResidentRolePipeline):
|
||||
"""Persistent DiT-only worker for T2VA and FL2VA."""
|
||||
|
||||
_required_config_modules = ["transformer", "scheduler", "audio_scheduler"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
_apply_h3_checkpoint_arch_configs(self.model_path, fastvideo_args, self._extra_config_module_map)
|
||||
for module_name, modality, expected_shift in (
|
||||
("scheduler", "video", 12.0),
|
||||
("audio_scheduler", "audio", 3.0),
|
||||
):
|
||||
shift = getattr(self.get_module(module_name), "shift", None)
|
||||
if shift is None or float(shift) != expected_shift:
|
||||
raise ValueError(f"MiniMax-H3 {modality} scheduler must expose shift={expected_shift:g}, got {shift}.")
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
del fastvideo_args
|
||||
transformer = self.get_module("transformer")
|
||||
scheduler = self.get_module("scheduler")
|
||||
audio_scheduler = self.get_module("audio_scheduler")
|
||||
if transformer is None or scheduler is None or audio_scheduler is None:
|
||||
raise RuntimeError("MiniMax-H3 DiT worker requires transformer and both schedulers.")
|
||||
self.add_stage(
|
||||
"denoising_stage",
|
||||
MiniMaxH3DenoisingStage(
|
||||
transformer=transformer,
|
||||
scheduler=scheduler,
|
||||
audio_scheduler=audio_scheduler,
|
||||
),
|
||||
)
|
||||
|
||||
def denoise(self, state: MiniMaxH3EncodedState) -> MiniMaxH3DenoisedState:
|
||||
if not self.post_init_called:
|
||||
self.post_init()
|
||||
batch = self._stage_name_mapping["denoising_stage"](state.to_batch(), self.fastvideo_args)
|
||||
return MiniMaxH3DenoisedState.from_batch(batch, request_id=state.request_id)
|
||||
|
||||
|
||||
class MiniMaxH3RefDiTPipeline(MiniMaxH3DiTPipeline):
|
||||
"""Persistent Ref2VA DiT worker using the checkpoint transformer_ref partition."""
|
||||
|
||||
_extra_config_module_map = {"transformer": "transformer_ref"}
|
||||
_ref2va_default = True
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MINIMAX_H3_WIRE_SCHEMA_VERSION",
|
||||
"MiniMaxH3DenoisedState",
|
||||
"MiniMaxH3DiTPipeline",
|
||||
"MiniMaxH3EncodedState",
|
||||
"MiniMaxH3EncoderDecoderPipeline",
|
||||
"MiniMaxH3RefDiTPipeline",
|
||||
"MiniMaxH3RefEncoderDecoderPipeline",
|
||||
]
|
||||
+15
-12
@@ -306,33 +306,36 @@ class CudaPlatformBase(Platform):
|
||||
|
||||
target_backend = AttentionBackendEnum.FLASH_ATTN
|
||||
if not cls.has_device_capability(80):
|
||||
logger.info("Cannot use FlashAttention-2 backend for Volta and Turing "
|
||||
logger.info("Cannot use FlashAttention backend for Volta and Turing "
|
||||
"GPUs.")
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
elif dtype not in (torch.float16, torch.bfloat16):
|
||||
logger.info("Cannot use FlashAttention-2 backend for dtype other than "
|
||||
logger.info("Cannot use FlashAttention backend for dtype other than "
|
||||
"torch.float16 or torch.bfloat16.")
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
# FlashAttn is valid for the model, checking if the package is
|
||||
# installed.
|
||||
# Import the actual backend so its FA4/FA3/FA2 selection determines
|
||||
# availability; a separate FA2 package probe can reject valid installs.
|
||||
if target_backend == AttentionBackendEnum.FLASH_ATTN:
|
||||
try:
|
||||
import flash_attn # noqa: F401
|
||||
|
||||
from fastvideo.attention.backends.flash_attn import ( # noqa: F401
|
||||
FlashAttentionBackend)
|
||||
|
||||
supported_sizes = \
|
||||
FlashAttentionBackend.get_supported_head_sizes()
|
||||
if head_size not in supported_sizes:
|
||||
logger.info("Cannot use FlashAttention-2 backend for head size %d.", head_size)
|
||||
logger.info("Cannot use FlashAttention backend for head size %d.", head_size)
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
except ImportError:
|
||||
logger.info("Cannot use FlashAttention-2 backend because the "
|
||||
"flash_attn package is not found. "
|
||||
"Make sure that flash_attn was built and installed "
|
||||
"(on by default).")
|
||||
except ImportError as e:
|
||||
if envs.FASTVIDEO_FA4:
|
||||
raise RuntimeError(
|
||||
f"FASTVIDEO_FA4=1 but the FlashAttention backend failed to import ({type(e).__name__}: {e}). "
|
||||
"Verify the FA4 install and its dependencies in this worker's Python environment, "
|
||||
"or unset FASTVIDEO_FA4 to allow automatic fallback.") from e
|
||||
logger.info(
|
||||
"Cannot import FlashAttention backend (FASTVIDEO_FA4=0): %s: %s. "
|
||||
"Falling back to Torch SDPA.",
|
||||
type(e).__name__, str(e))
|
||||
target_backend = AttentionBackendEnum.TORCH_SDPA
|
||||
|
||||
if target_backend == AttentionBackendEnum.TORCH_SDPA:
|
||||
|
||||
@@ -0,0 +1,102 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU regressions for the CUDA selector's FlashAttention import policy.
|
||||
|
||||
Exercise the real resolver while faking device capability and backend imports;
|
||||
no CUDA device or FlashAttention installation is required.
|
||||
"""
|
||||
|
||||
import builtins
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.platforms import cuda
|
||||
from fastvideo.platforms.cuda import NonNvmlCudaPlatform
|
||||
from fastvideo.platforms.interface import AttentionBackendEnum
|
||||
|
||||
FLASH_ATTN_CLS = "fastvideo.attention.backends.flash_attn.FlashAttentionBackend"
|
||||
SDPA_CLS = "fastvideo.attention.backends.sdpa.SDPABackend"
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _fake_device(monkeypatch):
|
||||
monkeypatch.delenv("FASTVIDEO_ATTENTION_BACKEND", raising=False)
|
||||
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
|
||||
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device=None: (12, 1))
|
||||
|
||||
|
||||
def _fake_backend_import(monkeypatch, error=None):
|
||||
original_import = builtins.__import__
|
||||
imported = []
|
||||
backend = SimpleNamespace(
|
||||
FlashAttentionBackend=SimpleNamespace(get_supported_head_sizes=lambda: [64, 128]),
|
||||
)
|
||||
|
||||
def import_backend(name, globals=None, locals=None, fromlist=(), level=0):
|
||||
if name == "fastvideo.attention.backends.flash_attn":
|
||||
imported.append(name)
|
||||
if error is not None:
|
||||
raise error
|
||||
return backend
|
||||
if name == "flash_attn":
|
||||
# FA3-only environments need not provide the FA2 parent package.
|
||||
raise ModuleNotFoundError("No module named 'flash_attn'")
|
||||
return original_import(name, globals, locals, fromlist, level)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", import_backend)
|
||||
return imported
|
||||
|
||||
|
||||
def _resolve(*, selected_backend=None, head_size=128, dtype=torch.bfloat16):
|
||||
return NonNvmlCudaPlatform.get_attn_backend_cls(selected_backend, head_size, dtype)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("fa4", ["0", "1"])
|
||||
def test_usable_backend_does_not_require_separate_fa2_package_probe(monkeypatch, fa4):
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", fa4)
|
||||
imported = _fake_backend_import(monkeypatch)
|
||||
|
||||
assert _resolve() == FLASH_ATTN_CLS
|
||||
assert imported == ["fastvideo.attention.backends.flash_attn"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("error", [
|
||||
ModuleNotFoundError("No module named 'flash_attn'"),
|
||||
ImportError("cannot import name 'flash_attn_varlen_func' from 'flash_attn.cute'"),
|
||||
])
|
||||
def test_automatic_import_failure_logs_actual_reason_and_falls_back(monkeypatch, error):
|
||||
messages = []
|
||||
monkeypatch.setattr(cuda.logger, "info", lambda message, *args: messages.append(message % args))
|
||||
_fake_backend_import(monkeypatch, error)
|
||||
|
||||
assert _resolve() == SDPA_CLS
|
||||
assert any(f"FASTVIDEO_FA4=0): {type(error).__name__}: {error}" in message for message in messages)
|
||||
|
||||
|
||||
def test_explicit_fa4_import_failure_does_not_silently_select_sdpa(monkeypatch):
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
error = ImportError("cannot import name 'flash_attn_varlen_func' from 'flash_attn.cute'")
|
||||
_fake_backend_import(monkeypatch, error)
|
||||
|
||||
with pytest.raises(RuntimeError, match="FASTVIDEO_FA4=1 but the FlashAttention backend failed to import") as exc_info:
|
||||
_resolve()
|
||||
|
||||
assert exc_info.value.__cause__ is error
|
||||
assert str(error) in str(exc_info.value)
|
||||
assert "worker's Python environment" in str(exc_info.value)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("kwargs, capability, import_expected", [
|
||||
({"dtype": torch.float32}, (12, 1), False),
|
||||
({"head_size": 48}, (12, 1), True),
|
||||
({}, (7, 5), False),
|
||||
({"selected_backend": AttentionBackendEnum.TORCH_SDPA}, (12, 1), False),
|
||||
])
|
||||
def test_fa4_opt_in_preserves_layer_and_explicit_sdpa_fallbacks(monkeypatch, kwargs, capability, import_expected):
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
monkeypatch.setattr(torch.cuda, "get_device_capability", lambda device=None: capability)
|
||||
imported = _fake_backend_import(monkeypatch)
|
||||
|
||||
assert _resolve(**kwargs) == SDPA_CLS
|
||||
assert bool(imported) is import_expected
|
||||
@@ -0,0 +1,651 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CUDA contracts for the two-node MiniMax-H3 component pipeline."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from dataclasses import replace
|
||||
import pickle
|
||||
from types import SimpleNamespace
|
||||
from typing import Any
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.pipelines.minimax_h3 import MiniMaxH3PipelineConfig
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.pipelines.basic.minimax_h3.disaggregated import (
|
||||
MINIMAX_H3_WIRE_SCHEMA_VERSION,
|
||||
MiniMaxH3DenoisedState,
|
||||
MiniMaxH3DiTPipeline,
|
||||
MiniMaxH3EncodedState,
|
||||
MiniMaxH3EncoderDecoderPipeline,
|
||||
MiniMaxH3RefDiTPipeline,
|
||||
MiniMaxH3RefEncoderDecoderPipeline,
|
||||
)
|
||||
from fastvideo.pipelines.basic.minimax_h3.packing import MiniMaxH3PackedLayout
|
||||
from fastvideo.pipelines.basic.minimax_h3.stages.minimax_h3_latent_preparation import MINIMAX_H3_LAYOUT_KEY
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
import fastvideo.worker.minimax_h3_disaggregated as disaggregated_runtime
|
||||
from fastvideo.worker.minimax_h3_disaggregated import (
|
||||
MiniMaxH3DisaggregatedExecutor,
|
||||
RayMiniMaxH3DisaggregatedRuntime,
|
||||
_MiniMaxH3DiTActor,
|
||||
_MiniMaxH3EncoderDecoderActor,
|
||||
_node_resource,
|
||||
_resident_role_args,
|
||||
_validate_topology,
|
||||
)
|
||||
from fastvideo.worker.executor import Executor
|
||||
|
||||
|
||||
def _layout(*, noncontiguous: bool = False) -> MiniMaxH3PackedLayout:
|
||||
position_ids = (
|
||||
torch.arange(27, dtype=torch.float64, device="cuda")
|
||||
.reshape(3, 9)
|
||||
.transpose(0, 1)
|
||||
)
|
||||
if not noncontiguous:
|
||||
position_ids = position_ids.contiguous()
|
||||
return MiniMaxH3PackedLayout(
|
||||
sequence_length=9,
|
||||
position_ids=position_ids,
|
||||
token_tags=torch.tensor(
|
||||
[1, 1, 1, 2, 2, 0, 0, 0, 0],
|
||||
device="cuda",
|
||||
),
|
||||
video_indices=torch.tensor([5, 6, 7, 8], device="cuda"),
|
||||
audio_indices=torch.tensor([3, 4], device="cuda"),
|
||||
text_indices=torch.tensor([0, 1, 2], device="cuda"),
|
||||
num_condition_video_rows=0,
|
||||
num_condition_audio_rows=0,
|
||||
num_video_latent_frames=2,
|
||||
latent_height=2,
|
||||
latent_width=2,
|
||||
num_audio_latents=1,
|
||||
)
|
||||
|
||||
|
||||
def _encoded_batch(*, request_id: str | None = None) -> ForwardBatch:
|
||||
prompt_embeds = (
|
||||
torch.arange(12, dtype=torch.float32, device="cuda")
|
||||
.reshape(1, 4, 3)
|
||||
.transpose(1, 2)
|
||||
)
|
||||
video_latents = (
|
||||
torch.arange(20, dtype=torch.float32, device="cuda")
|
||||
.reshape(5, 4)
|
||||
.transpose(0, 1)
|
||||
)
|
||||
audio_latents = (
|
||||
torch.arange(12, dtype=torch.float32, device="cuda")
|
||||
.reshape(6, 2)
|
||||
.transpose(0, 1)
|
||||
)
|
||||
extra: dict[str, Any] = {
|
||||
MINIMAX_H3_LAYOUT_KEY: _layout(noncontiguous=True),
|
||||
"vsa_mode": "compete",
|
||||
"vsa_dense_first_n_steps": 2,
|
||||
"vsa_dense_layers": [1, 7],
|
||||
}
|
||||
if request_id is not None:
|
||||
extra["request_id"] = request_id
|
||||
return ForwardBatch(
|
||||
data_type="video",
|
||||
prompt="must not cross the wire",
|
||||
prompt_embeds=[prompt_embeds],
|
||||
latents=video_latents,
|
||||
audio_latents=audio_latents,
|
||||
raw_latent_shape=(1, 16, 2, 2, 2),
|
||||
num_inference_steps=8,
|
||||
VSA_sparsity=0.75,
|
||||
extra=extra,
|
||||
)
|
||||
|
||||
|
||||
def test_encoded_wire_state_is_minimal_cuda_contiguous_and_pickleable() -> None:
|
||||
batch = _encoded_batch()
|
||||
assert not batch.prompt_embeds[0].is_contiguous()
|
||||
assert not batch.latents.is_contiguous()
|
||||
assert not batch.audio_latents.is_contiguous()
|
||||
assert not batch.extra[MINIMAX_H3_LAYOUT_KEY].position_ids.is_contiguous()
|
||||
|
||||
state = MiniMaxH3EncodedState.from_batch(batch, request_id="request-7")
|
||||
restored = pickle.loads(pickle.dumps(state))
|
||||
|
||||
assert restored.request_id == "request-7"
|
||||
assert restored.schema_version == MINIMAX_H3_WIRE_SCHEMA_VERSION
|
||||
assert set(vars(restored)) == {
|
||||
"request_id",
|
||||
"prompt_embeds",
|
||||
"video_latents",
|
||||
"audio_latents",
|
||||
"layout",
|
||||
"raw_latent_shape",
|
||||
"num_inference_steps",
|
||||
"vsa_sparsity",
|
||||
"vsa_mode",
|
||||
"vsa_dense_first_n_steps",
|
||||
"vsa_dense_layers",
|
||||
"logging_info",
|
||||
"schema_version",
|
||||
}
|
||||
for tensor in (
|
||||
restored.prompt_embeds,
|
||||
restored.video_latents,
|
||||
restored.audio_latents,
|
||||
restored.layout.position_ids,
|
||||
restored.layout.token_tags,
|
||||
restored.layout.video_indices,
|
||||
restored.layout.audio_indices,
|
||||
restored.layout.text_indices,
|
||||
):
|
||||
assert tensor.device.type == "cuda"
|
||||
assert tensor.is_contiguous()
|
||||
|
||||
roundtrip = restored.to_batch()
|
||||
assert torch.equal(roundtrip.prompt_embeds[0], batch.prompt_embeds[0])
|
||||
assert torch.equal(roundtrip.latents, batch.latents)
|
||||
assert torch.equal(roundtrip.audio_latents, batch.audio_latents)
|
||||
assert roundtrip.raw_latent_shape == batch.raw_latent_shape
|
||||
assert roundtrip.num_inference_steps == 8
|
||||
assert roundtrip.VSA_sparsity == 0.75
|
||||
assert roundtrip.extra["vsa_mode"] == "compete"
|
||||
assert roundtrip.extra["vsa_dense_first_n_steps"] == 2
|
||||
assert roundtrip.extra["vsa_dense_layers"] == (1, 7)
|
||||
|
||||
|
||||
def test_denoised_wire_state_roundtrips_only_decode_inputs() -> None:
|
||||
encoded = MiniMaxH3EncodedState.from_batch(_encoded_batch(), request_id="request-8")
|
||||
denoised_batch = encoded.to_batch()
|
||||
denoised_batch.latents = denoised_batch.latents + 1
|
||||
denoised_batch.audio_latents = denoised_batch.audio_latents - 1
|
||||
|
||||
state = MiniMaxH3DenoisedState.from_batch(denoised_batch, request_id=encoded.request_id)
|
||||
restored = pickle.loads(pickle.dumps(state))
|
||||
roundtrip = restored.to_batch()
|
||||
|
||||
assert restored.request_id == "request-8"
|
||||
assert set(vars(restored)) == {
|
||||
"request_id",
|
||||
"video_latents",
|
||||
"audio_latents",
|
||||
"layout",
|
||||
"raw_latent_shape",
|
||||
"logging_info",
|
||||
"schema_version",
|
||||
}
|
||||
assert torch.equal(roundtrip.latents, denoised_batch.latents)
|
||||
assert torch.equal(roundtrip.audio_latents, denoised_batch.audio_latents)
|
||||
assert roundtrip.raw_latent_shape == (1, 16, 2, 2, 2)
|
||||
assert roundtrip.prompt_embeds == []
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("mutation", "message"),
|
||||
[
|
||||
({"schema_version": 999}, "schema version"),
|
||||
({"request_id": ""}, "non-empty request_id"),
|
||||
({"raw_latent_shape": (1, 16, 0, 2, 2)}, "raw latent shape"),
|
||||
({"num_inference_steps": 0}, "at least one denoising step"),
|
||||
({"vsa_mode": "unknown"}, "vsa_mode"),
|
||||
({"video_latents": torch.zeros(3, 5, device="cuda")}, "video latent rows"),
|
||||
({"audio_latents": torch.zeros(1, 6, device="cuda")}, "audio latent rows"),
|
||||
({"prompt_embeds": torch.zeros(1, 2, 4, device="cuda")}, "prompt embedding rows"),
|
||||
],
|
||||
)
|
||||
def test_encoded_wire_state_rejects_incompatible_payloads(mutation: dict[str, Any], message: str) -> None:
|
||||
state = MiniMaxH3EncodedState.from_batch(_encoded_batch(), request_id="valid")
|
||||
with pytest.raises(ValueError, match=message):
|
||||
replace(state, **mutation)
|
||||
|
||||
|
||||
def test_wire_state_rejects_noncontiguous_or_malformed_layout() -> None:
|
||||
state = MiniMaxH3EncodedState.from_batch(_encoded_batch(), request_id="valid")
|
||||
with pytest.raises(ValueError, match="layout.position_ids must be a contiguous CUDA tensor"):
|
||||
replace(state, layout=_layout(noncontiguous=True))
|
||||
|
||||
malformed = replace(_layout(), token_tags=torch.zeros(8, dtype=torch.long, device="cuda"))
|
||||
with pytest.raises(ValueError, match=r"layout.token_tags.*expected \(9,\)"):
|
||||
replace(state, layout=malformed)
|
||||
|
||||
|
||||
def test_from_batch_requires_all_stage_boundary_fields() -> None:
|
||||
with pytest.raises(ValueError, match="one prompt embedding and both latent streams"):
|
||||
MiniMaxH3EncodedState.from_batch(ForwardBatch(data_type="video"))
|
||||
|
||||
batch = _encoded_batch()
|
||||
batch.raw_latent_shape = None
|
||||
with pytest.raises(ValueError, match="five-dimensional raw latent shape"):
|
||||
MiniMaxH3EncodedState.from_batch(batch)
|
||||
|
||||
batch = _encoded_batch()
|
||||
batch.extra.pop(MINIMAX_H3_LAYOUT_KEY)
|
||||
with pytest.raises(TypeError, match="missing its packed layout"):
|
||||
MiniMaxH3EncodedState.from_batch(batch)
|
||||
|
||||
|
||||
def test_resident_roles_have_isolated_component_sets_and_no_lazy_modules() -> None:
|
||||
assert set(MiniMaxH3EncoderDecoderPipeline._required_config_modules) == {
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"processor",
|
||||
"vae",
|
||||
"audio_vae",
|
||||
"scheduler",
|
||||
}
|
||||
assert set(MiniMaxH3DiTPipeline._required_config_modules) == {
|
||||
"transformer",
|
||||
"scheduler",
|
||||
"audio_scheduler",
|
||||
}
|
||||
assert MiniMaxH3EncoderDecoderPipeline._lazy_module_names == ()
|
||||
assert MiniMaxH3DiTPipeline._lazy_module_names == ()
|
||||
|
||||
|
||||
def test_role_downloads_include_only_owned_weights_plus_required_dit_metadata() -> None:
|
||||
encoder_dirs = set(MiniMaxH3EncoderDecoderPipeline.get_hf_download_component_dirs())
|
||||
assert encoder_dirs == {
|
||||
"text_encoder",
|
||||
"tokenizer",
|
||||
"processor",
|
||||
"vae",
|
||||
"audio_vae",
|
||||
"scheduler",
|
||||
}
|
||||
encoder_patterns = MiniMaxH3EncoderDecoderPipeline.get_hf_download_allow_patterns()
|
||||
assert "transformer/config.json" in encoder_patterns
|
||||
assert "transformer/**" not in encoder_patterns
|
||||
|
||||
ref_encoder_patterns = MiniMaxH3RefEncoderDecoderPipeline.get_hf_download_allow_patterns()
|
||||
assert "transformer/config.json" in ref_encoder_patterns
|
||||
assert "transformer_ref/config.json" in ref_encoder_patterns
|
||||
assert "transformer/**" not in ref_encoder_patterns
|
||||
assert "transformer_ref/**" not in ref_encoder_patterns
|
||||
|
||||
assert set(MiniMaxH3DiTPipeline.get_hf_download_component_dirs()) == {
|
||||
"transformer",
|
||||
"scheduler",
|
||||
"audio_scheduler",
|
||||
}
|
||||
assert set(MiniMaxH3RefDiTPipeline.get_hf_download_component_dirs()) == {
|
||||
"transformer_ref",
|
||||
"scheduler",
|
||||
"audio_scheduler",
|
||||
}
|
||||
|
||||
|
||||
def _role_source_args() -> SimpleNamespace:
|
||||
return SimpleNamespace(
|
||||
num_gpus=8,
|
||||
tp_size=2,
|
||||
sp_size=4,
|
||||
hsdp_replicate_dim=2,
|
||||
hsdp_shard_dim=4,
|
||||
ray_placement_group=object(),
|
||||
ray_runtime_env={"env_vars": {"ORIGINAL": "1"}},
|
||||
distributed_executor_backend="ray",
|
||||
use_fsdp_inference=True,
|
||||
dit_cpu_offload=True,
|
||||
dit_layerwise_offload=True,
|
||||
text_encoder_cpu_offload=True,
|
||||
image_encoder_cpu_offload=True,
|
||||
vae_cpu_offload=True,
|
||||
lazy_module_load=True,
|
||||
h3_sequential_load=True,
|
||||
vae_parallel_encode=True,
|
||||
vae_parallel_decode=True,
|
||||
lora_path="adapter.safetensors",
|
||||
enable_torch_compile=True,
|
||||
enable_torch_compile_text_encoder=True,
|
||||
enable_torch_compile_vae=True,
|
||||
enable_torch_compile_audio_vae=True,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("role", ["encoder_decoder", "dit"])
|
||||
def test_resident_role_args_force_single_gpu_persistent_execution(role: str) -> None:
|
||||
source = _role_source_args()
|
||||
args = _resident_role_args(source, role=role)
|
||||
|
||||
assert (args.num_gpus, args.tp_size, args.sp_size) == (1, 1, 1)
|
||||
assert (args.hsdp_replicate_dim, args.hsdp_shard_dim) == (1, 1)
|
||||
assert args.distributed_executor_backend == "mp"
|
||||
assert args.ray_placement_group is None
|
||||
assert args.ray_runtime_env is None
|
||||
assert args.use_fsdp_inference is False
|
||||
assert args.lazy_module_load is False
|
||||
assert args.h3_sequential_load is False
|
||||
assert args.vae_parallel_encode is False
|
||||
assert args.vae_parallel_decode is False
|
||||
assert all(
|
||||
getattr(args, name) is False for name in (
|
||||
"dit_cpu_offload",
|
||||
"dit_layerwise_offload",
|
||||
"text_encoder_cpu_offload",
|
||||
"image_encoder_cpu_offload",
|
||||
"vae_cpu_offload",
|
||||
)
|
||||
)
|
||||
assert source.num_gpus == 8
|
||||
assert source.lazy_module_load is True
|
||||
assert source.ray_runtime_env == {"env_vars": {"ORIGINAL": "1"}}
|
||||
|
||||
if role == "encoder_decoder":
|
||||
assert args.lora_path is None
|
||||
assert args.enable_torch_compile is False
|
||||
assert args.enable_torch_compile_text_encoder is True
|
||||
else:
|
||||
assert args.lora_path == "adapter.safetensors"
|
||||
assert args.enable_torch_compile is True
|
||||
assert args.enable_torch_compile_text_encoder is False
|
||||
assert args.enable_torch_compile_vae is False
|
||||
assert args.enable_torch_compile_audio_vae is False
|
||||
|
||||
|
||||
def test_resident_role_args_reject_unknown_role() -> None:
|
||||
with pytest.raises(ValueError, match="Unknown MiniMax-H3 worker role"):
|
||||
_resident_role_args(_role_source_args(), role="both")
|
||||
|
||||
|
||||
def test_topology_requires_two_distinct_live_ray_nodes() -> None:
|
||||
resources = {
|
||||
_node_resource("10.0.0.1"): 1.0,
|
||||
_node_resource("10.0.0.2"): 1.0,
|
||||
"CPU": 64.0,
|
||||
}
|
||||
_validate_topology("10.0.0.1", "10.0.0.2", resources)
|
||||
|
||||
with pytest.raises(ValueError, match="both encoder and DiT node IPs"):
|
||||
_validate_topology("", "10.0.0.2", resources)
|
||||
with pytest.raises(ValueError, match="must use different Ray nodes"):
|
||||
_validate_topology("10.0.0.1", "10.0.0.1", resources)
|
||||
with pytest.raises(RuntimeError, match=r"10\.0\.0\.3.*available node IPs"):
|
||||
_validate_topology("10.0.0.1", "10.0.0.3", resources)
|
||||
|
||||
|
||||
def test_executor_selector_routes_disaggregated_h3_away_from_full_pipeline_ray_workers() -> None:
|
||||
args = SimpleNamespace(h3_disaggregated=True, distributed_executor_backend="ray")
|
||||
assert Executor.get_class(args) is MiniMaxH3DisaggregatedExecutor
|
||||
|
||||
|
||||
def _valid_disaggregated_args(**overrides: Any) -> FastVideoArgs:
|
||||
values: dict[str, Any] = {
|
||||
"model_path": "unused/for-this-test",
|
||||
"pipeline_config": MiniMaxH3PipelineConfig(),
|
||||
"distributed_executor_backend": "ray",
|
||||
"h3_disaggregated": True,
|
||||
"h3_encoder_node_ip": "10.0.0.1",
|
||||
"h3_dit_node_ip": "10.0.0.2",
|
||||
}
|
||||
values.update(overrides)
|
||||
return FastVideoArgs(**values)
|
||||
|
||||
|
||||
def test_disaggregated_args_are_opt_in_and_trim_node_addresses() -> None:
|
||||
args = _valid_disaggregated_args(
|
||||
h3_encoder_node_ip=" 10.0.0.1 ",
|
||||
h3_dit_node_ip=" 10.0.0.2 ",
|
||||
h3_ray_address=" 10.0.0.1:6379 ",
|
||||
)
|
||||
|
||||
assert args.h3_encoder_node_ip == "10.0.0.1"
|
||||
assert args.h3_dit_node_ip == "10.0.0.2"
|
||||
assert args.h3_ray_address == "10.0.0.1:6379"
|
||||
assert Executor.get_class(args) is MiniMaxH3DisaggregatedExecutor
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("overrides", "message"),
|
||||
[
|
||||
({"distributed_executor_backend": "mp"}, "requires distributed_executor_backend='ray'"),
|
||||
({"h3_dit_node_ip": "10.0.0.1"}, "must use different Ray node"),
|
||||
({"h3_encoder_node_ip": None}, "requires h3_encoder_node_ip"),
|
||||
({"num_gpus": 2}, "requires num_gpus=1"),
|
||||
({"tp_size": 2, "num_gpus": 2}, "requires num_gpus=1"),
|
||||
({"sp_size": 2, "num_gpus": 2}, "requires num_gpus=1"),
|
||||
({"use_fsdp_inference": True}, "does not support FSDP"),
|
||||
({"lazy_module_load": True}, "keeps role components resident"),
|
||||
({"h3_sequential_load": True}, "keeps role components resident"),
|
||||
({"vae_parallel_decode": True}, "does not support sequence-parallel VAE"),
|
||||
({"video_decode_backend": "taeh3"}, "requires video_decode_backend='h3-vae'"),
|
||||
],
|
||||
)
|
||||
def test_disaggregated_args_reject_incompatible_execution(overrides: dict[str, Any], message: str) -> None:
|
||||
with pytest.raises(ValueError, match=message):
|
||||
_valid_disaggregated_args(**overrides)
|
||||
|
||||
|
||||
class _FakeRef:
|
||||
|
||||
def __init__(self, kind: str, ordinal: int, value: Any = None) -> None:
|
||||
self.kind = kind
|
||||
self.ordinal = ordinal
|
||||
self.value = value
|
||||
|
||||
|
||||
class _RemoteMethod:
|
||||
|
||||
def __init__(self, name: str, events: list[tuple[str, tuple[Any, ...]]], result_factory=None) -> None:
|
||||
self.name = name
|
||||
self.events = events
|
||||
self.calls: list[tuple[Any, ...]] = []
|
||||
self.refs: list[_FakeRef] = []
|
||||
self.result_factory = result_factory
|
||||
|
||||
def remote(self, *args: Any) -> _FakeRef:
|
||||
self.calls.append(args)
|
||||
self.events.append((self.name, args))
|
||||
value = None if self.result_factory is None else self.result_factory(*args)
|
||||
ref = _FakeRef(self.name, len(self.calls), value)
|
||||
self.refs.append(ref)
|
||||
return ref
|
||||
|
||||
|
||||
class _FakeEncoderActor:
|
||||
|
||||
def __init__(self, events: list[tuple[str, tuple[Any, ...]]]) -> None:
|
||||
self.encode = _RemoteMethod("encode", events)
|
||||
self.decode = _RemoteMethod(
|
||||
"decode",
|
||||
events,
|
||||
lambda denoised_ref: ForwardBatch(data_type="video", extra={"decoded_from": denoised_ref}),
|
||||
)
|
||||
|
||||
|
||||
class _FakeDiTActor:
|
||||
|
||||
def __init__(self, events: list[tuple[str, tuple[Any, ...]]]) -> None:
|
||||
self.denoise = _RemoteMethod("denoise", events)
|
||||
|
||||
|
||||
class _FakeRay:
|
||||
|
||||
def __init__(self, events: list[tuple[str, tuple[Any, ...]]]) -> None:
|
||||
self.events = events
|
||||
self.get_calls: list[_FakeRef] = []
|
||||
self.wait_calls: list[tuple[list[_FakeRef], int, bool]] = []
|
||||
|
||||
def get(self, ref: _FakeRef | list[_FakeRef]):
|
||||
self.get_calls.append(ref)
|
||||
self.events.append(("get", (ref, )))
|
||||
return [item.value for item in ref] if isinstance(ref, list) else ref.value
|
||||
|
||||
def wait(self, refs: list[_FakeRef], *, num_returns: int, fetch_local: bool):
|
||||
self.wait_calls.append((refs, num_returns, fetch_local))
|
||||
self.events.append(("wait", (refs[0], )))
|
||||
return refs[:num_returns], refs[num_returns:]
|
||||
|
||||
def kill(self, actor: Any, *, no_restart: bool) -> None:
|
||||
self.events.append(("kill", (actor, no_restart)))
|
||||
|
||||
|
||||
def _fake_runtime(monkeypatch):
|
||||
events: list[tuple[str, tuple[Any, ...]]] = []
|
||||
fake_ray = _FakeRay(events)
|
||||
monkeypatch.setattr(disaggregated_runtime, "ray", fake_ray)
|
||||
runtime = RayMiniMaxH3DisaggregatedRuntime.__new__(RayMiniMaxH3DisaggregatedRuntime)
|
||||
runtime._closed = False
|
||||
runtime._profile_transfers = False
|
||||
runtime.encoder_decoder = _FakeEncoderActor(events)
|
||||
runtime.dit = _FakeDiTActor(events)
|
||||
return runtime, fake_ray, events
|
||||
|
||||
|
||||
def test_submit_builds_actor_dag_without_fetching_intermediate_payloads(monkeypatch) -> None:
|
||||
runtime, fake_ray, _ = _fake_runtime(monkeypatch)
|
||||
batch = _encoded_batch(request_id="batch-id")
|
||||
|
||||
decoded_ref = runtime.submit(batch)
|
||||
|
||||
encode_args = runtime.encoder_decoder.encode.calls[0]
|
||||
assert encode_args == (batch, "batch-id")
|
||||
assert runtime.dit.denoise.calls[0] == (runtime.encoder_decoder.encode.refs[0], )
|
||||
assert runtime.encoder_decoder.decode.calls[0] == (runtime.dit.denoise.refs[0], )
|
||||
assert decoded_ref.kind == "decode"
|
||||
assert fake_ray.get_calls == []
|
||||
|
||||
|
||||
def test_execute_forward_fetches_only_the_final_decode(monkeypatch) -> None:
|
||||
runtime, fake_ray, _ = _fake_runtime(monkeypatch)
|
||||
|
||||
output = runtime.execute_forward(_encoded_batch(), request_id="explicit-id")
|
||||
|
||||
assert isinstance(output, ForwardBatch)
|
||||
assert runtime.encoder_decoder.encode.calls[0][1] == "explicit-id"
|
||||
assert [ref.kind for ref in fake_ray.get_calls] == ["decode"]
|
||||
|
||||
|
||||
def test_iter_forward_overlaps_encode_denoise_decode_with_bounded_lookahead(monkeypatch) -> None:
|
||||
runtime, fake_ray, events = _fake_runtime(monkeypatch)
|
||||
batches = [_encoded_batch(request_id=f"request-{index}") for index in range(3)]
|
||||
|
||||
outputs = list(runtime.iter_forward(batches))
|
||||
|
||||
assert len(outputs) == 3
|
||||
assert [call[1] for call in runtime.encoder_decoder.encode.calls] == [
|
||||
"request-0",
|
||||
"request-1",
|
||||
"request-2",
|
||||
]
|
||||
assert [ref.kind for ref in fake_ray.get_calls] == ["decode", "decode", "decode"]
|
||||
assert len(fake_ray.wait_calls) == 2
|
||||
assert all(refs[0].kind == "denoise" and num_returns == 1 and fetch_local is False
|
||||
for refs, num_returns, fetch_local in fake_ray.wait_calls)
|
||||
|
||||
event_names = [name for name, _ in events]
|
||||
assert event_names == [
|
||||
"encode",
|
||||
"denoise",
|
||||
"encode",
|
||||
"wait",
|
||||
"decode",
|
||||
"denoise",
|
||||
"get",
|
||||
"encode",
|
||||
"wait",
|
||||
"decode",
|
||||
"denoise",
|
||||
"get",
|
||||
"decode",
|
||||
"get",
|
||||
]
|
||||
|
||||
|
||||
def test_iter_forward_async_exposes_the_same_bounded_pipeline(monkeypatch) -> None:
|
||||
runtime, fake_ray, _ = _fake_runtime(monkeypatch)
|
||||
batches = [_encoded_batch(request_id=f"async-{index}") for index in range(2)]
|
||||
|
||||
async def collect() -> list[ForwardBatch]:
|
||||
return [batch async for batch in runtime.iter_forward_async(batches)]
|
||||
|
||||
outputs = asyncio.run(collect())
|
||||
|
||||
assert len(outputs) == 2
|
||||
assert [call[1] for call in runtime.encoder_decoder.encode.calls] == ["async-0", "async-1"]
|
||||
assert [ref.kind for ref in fake_ray.get_calls] == ["decode", "decode"]
|
||||
|
||||
|
||||
class _ActorHandle:
|
||||
"""Ray-faithful handle: only methods the actor class defines are callable."""
|
||||
|
||||
def __init__(self, actor_cls: type, events: list[tuple[str, tuple[Any, ...]]], results: dict[str, Any]) -> None:
|
||||
self._actor_cls = actor_cls
|
||||
self._events = events
|
||||
self._results = results
|
||||
|
||||
def __getattr__(self, name: str) -> _RemoteMethod:
|
||||
if name.startswith("_") or not callable(getattr(self._actor_cls, name, None)):
|
||||
raise AttributeError(f"'{self._actor_cls.__name__}' actor has no method {name!r}")
|
||||
return _RemoteMethod(name, self._events, lambda *args: self._results.get(name))
|
||||
|
||||
|
||||
def _real_interface_runtime(monkeypatch, results: dict[str, Any] | None = None):
|
||||
events: list[tuple[str, tuple[Any, ...]]] = []
|
||||
monkeypatch.setattr(disaggregated_runtime, "ray", _FakeRay(events))
|
||||
runtime = RayMiniMaxH3DisaggregatedRuntime.__new__(RayMiniMaxH3DisaggregatedRuntime)
|
||||
runtime._closed = False
|
||||
runtime._profile_transfers = False
|
||||
results = results or {}
|
||||
runtime.encoder_decoder = _ActorHandle(_MiniMaxH3EncoderDecoderActor, events, results)
|
||||
runtime.dit = _ActorHandle(_MiniMaxH3DiTActor, events, results)
|
||||
return runtime, events
|
||||
|
||||
|
||||
def test_close_shuts_down_both_real_actor_interfaces_then_kills_them(monkeypatch) -> None:
|
||||
runtime, events = _real_interface_runtime(monkeypatch)
|
||||
|
||||
runtime.close()
|
||||
runtime.close() # idempotent
|
||||
|
||||
names = [name for name, _ in events]
|
||||
assert names == ["shutdown", "shutdown", "get", "kill", "kill"]
|
||||
|
||||
|
||||
def test_executor_lora_rpcs_reach_the_dit_actor_and_check_receipts(monkeypatch) -> None:
|
||||
receipts = {
|
||||
"set_lora_adapter": {"status": "lora_adapter_set"},
|
||||
"unmerge_lora_weights": {"status": "lora_adapter_unmerged"},
|
||||
"merge_lora_weights": {"status": "lora_adapter_merged"},
|
||||
}
|
||||
runtime, events = _real_interface_runtime(monkeypatch, receipts)
|
||||
executor = MiniMaxH3DisaggregatedExecutor.__new__(MiniMaxH3DisaggregatedExecutor)
|
||||
executor.runtime = runtime
|
||||
|
||||
executor.set_lora_adapter("fast", "adapter.safetensors", strength=0.5, accumulate=True)
|
||||
executor.unmerge_lora_weights()
|
||||
executor.merge_lora_weights()
|
||||
|
||||
calls = [(name, args) for name, args in events if name != "get"]
|
||||
assert calls == [
|
||||
("set_lora_adapter", ("fast", "adapter.safetensors", 0.5, True)),
|
||||
("unmerge_lora_weights", ()),
|
||||
("merge_lora_weights", ()),
|
||||
]
|
||||
receipts["merge_lora_weights"] = {"status": "failed"}
|
||||
with pytest.raises(RuntimeError, match="lora_adapter_merged"):
|
||||
executor.merge_lora_weights()
|
||||
runtime._closed = True # skip actor shutdown in __del__
|
||||
|
||||
|
||||
def test_dit_actor_lora_rpcs_delegate_to_the_resident_pipeline() -> None:
|
||||
calls: list[tuple[str, tuple[Any, ...], dict[str, Any]]] = []
|
||||
|
||||
def record(name: str):
|
||||
return lambda *args, **kwargs: calls.append((name, args, kwargs))
|
||||
|
||||
actor = _MiniMaxH3DiTActor.__new__(_MiniMaxH3DiTActor)
|
||||
actor.pipeline = SimpleNamespace(
|
||||
set_lora_adapter=record("set"),
|
||||
unmerge_lora_weights=record("unmerge"),
|
||||
merge_lora_weights=record("merge"),
|
||||
)
|
||||
|
||||
assert actor.set_lora_adapter("fast", "adapter.safetensors", 0.5, True) == {"status": "lora_adapter_set"}
|
||||
assert actor.unmerge_lora_weights() == {"status": "lora_adapter_unmerged"}
|
||||
assert actor.merge_lora_weights() == {"status": "lora_adapter_merged"}
|
||||
assert calls == [
|
||||
("set", ("fast", "adapter.safetensors"), {"strength": 0.5, "accumulate": True}),
|
||||
("unmerge", (), {}),
|
||||
("merge", (), {}),
|
||||
]
|
||||
@@ -0,0 +1,109 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU checks for settings applied before H3 Ray actor imports."""
|
||||
|
||||
from copy import deepcopy
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
|
||||
import fastvideo.worker.minimax_h3_disaggregated as runtime_module
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def clean_fa4_environment(monkeypatch):
|
||||
monkeypatch.delenv("FASTVIDEO_FA4", raising=False)
|
||||
monkeypatch.setattr(runtime_module, "RAY_NON_CARRY_OVER_ENV_VARS", set())
|
||||
|
||||
|
||||
@pytest.mark.parametrize("setting", ["0", "1"])
|
||||
def test_actor_env_carries_explicit_driver_fa4_without_node_identity(monkeypatch, setting):
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", setting)
|
||||
monkeypatch.setenv("FASTVIDEO_HOST_IP", "192.168.23.1")
|
||||
monkeypatch.setenv("CUDA_VISIBLE_DEVICES", "3")
|
||||
monkeypatch.setenv("NCCL_SOCKET_IFNAME", "driver-nic")
|
||||
args = SimpleNamespace(ray_runtime_env=None)
|
||||
assert runtime_module._h3_actor_runtime_env(args) == {"env_vars": {"FASTVIDEO_FA4": setting}}
|
||||
|
||||
|
||||
def test_unset_driver_flag_preserves_worker_defaults():
|
||||
assert runtime_module._h3_actor_runtime_env(SimpleNamespace(ray_runtime_env=None)) == {}
|
||||
|
||||
|
||||
def test_configured_actor_environment_wins_without_mutating_input(monkeypatch):
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
configured = {"env_vars": {"FASTVIDEO_FA4": "0", "EXISTING": "keep"}, "working_dir": "/tmp/project"}
|
||||
original = deepcopy(configured)
|
||||
actual = runtime_module._h3_actor_runtime_env(SimpleNamespace(ray_runtime_env=configured))
|
||||
assert actual == {"env_vars": original["env_vars"]}
|
||||
actual["env_vars"]["EXISTING"] = "changed"
|
||||
assert configured == original
|
||||
|
||||
|
||||
def test_driver_flag_merges_with_existing_environment(monkeypatch):
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
configured = {"env_vars": {"EXISTING": "keep"}, "pip": ["example-package"]}
|
||||
actual = runtime_module._h3_actor_runtime_env(SimpleNamespace(ray_runtime_env=configured))
|
||||
assert actual == {"env_vars": {"EXISTING": "keep", "FASTVIDEO_FA4": "1"}}
|
||||
assert "FASTVIDEO_FA4" not in configured["env_vars"]
|
||||
|
||||
|
||||
def test_actor_environment_inherits_job_packages_instead_of_reusing_local_paths(monkeypatch):
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
configured = {"working_dir": "/tmp/project", "py_modules": ["/tmp/local-module"]}
|
||||
assert runtime_module._h3_actor_runtime_env(SimpleNamespace(ray_runtime_env=configured)) == {
|
||||
"env_vars": {"FASTVIDEO_FA4": "1"}
|
||||
}
|
||||
|
||||
|
||||
def test_driver_flag_respects_non_carry_over_policy(monkeypatch):
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
monkeypatch.setattr(runtime_module, "RAY_NON_CARRY_OVER_ENV_VARS", {"FASTVIDEO_FA4"})
|
||||
assert runtime_module._h3_actor_runtime_env(SimpleNamespace(ray_runtime_env=None)) == {}
|
||||
|
||||
|
||||
@pytest.mark.parametrize("initialized", [False, True])
|
||||
def test_both_actor_creation_options_receive_fa4_even_with_existing_ray(monkeypatch, initialized):
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
options = []
|
||||
init_calls = []
|
||||
|
||||
class ActorBuilder:
|
||||
def options(self, **kwargs):
|
||||
options.append(kwargs)
|
||||
return self
|
||||
|
||||
def remote(self, *args):
|
||||
return object()
|
||||
|
||||
fake_ray = SimpleNamespace(
|
||||
is_initialized=lambda: initialized,
|
||||
init=lambda **kwargs: init_calls.append(kwargs),
|
||||
cluster_resources=lambda: {"node:192.168.23.1": 1.0, "node:192.168.23.2": 1.0},
|
||||
remote=lambda cls: ActorBuilder(),
|
||||
)
|
||||
monkeypatch.setattr(runtime_module, "ray", fake_ray)
|
||||
monkeypatch.setattr(runtime_module, "assert_ray_available", lambda: None)
|
||||
monkeypatch.setattr(runtime_module.RayMiniMaxH3DisaggregatedRuntime, "_validate_workers", lambda self: None)
|
||||
configured = {"env_vars": {"EXISTING": "keep"}}
|
||||
runtime_module.RayMiniMaxH3DisaggregatedRuntime(
|
||||
SimpleNamespace(ray_runtime_env=configured),
|
||||
encoder_node_ip="192.168.23.1",
|
||||
dit_node_ip="192.168.23.2",
|
||||
)
|
||||
assert len(init_calls) == (0 if initialized else 1)
|
||||
assert len(options) == 2
|
||||
for actor_options in options:
|
||||
assert actor_options["runtime_env"] == {"env_vars": {"EXISTING": "keep", "FASTVIDEO_FA4": "1"}}
|
||||
assert configured == {"env_vars": {"EXISTING": "keep"}}
|
||||
|
||||
|
||||
def test_worker_diagnostic_reports_actual_python_and_fa4(monkeypatch):
|
||||
messages = []
|
||||
monkeypatch.setenv("FASTVIDEO_FA4", "1")
|
||||
monkeypatch.setattr(runtime_module, "get_ip", lambda: "192.168.23.2")
|
||||
monkeypatch.setattr(runtime_module.logger, "info", lambda fmt, *args: messages.append(fmt % args))
|
||||
runtime_module._log_h3_worker_environment("dit")
|
||||
line, = messages
|
||||
assert "role=dit node_ip=192.168.23.2" in line
|
||||
assert f"python={runtime_module.sys.executable}" in line
|
||||
assert "FASTVIDEO_FA4=1" in line
|
||||
@@ -0,0 +1,195 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""CPU regressions for H3 Ray timing; no Ray cluster or CUDA device required."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
import fastvideo.worker.minimax_h3_disaggregated as runtime_module
|
||||
|
||||
|
||||
def _state():
|
||||
return SimpleNamespace(
|
||||
request_id="request-7",
|
||||
prompt_embeds=torch.zeros(1, 3, 4),
|
||||
video_latents=torch.zeros(4, 5),
|
||||
audio_latents=torch.zeros(2, 6),
|
||||
layout=SimpleNamespace(
|
||||
position_ids=torch.zeros(9, 3, dtype=torch.float64),
|
||||
token_tags=torch.zeros(9, dtype=torch.int64),
|
||||
text_indices=torch.zeros(3, dtype=torch.int64),
|
||||
video_indices=torch.zeros(4, dtype=torch.int64),
|
||||
audio_indices=torch.zeros(2, dtype=torch.int64),
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def timing(monkeypatch):
|
||||
clock = SimpleNamespace(now=0.0)
|
||||
events = []
|
||||
messages = []
|
||||
result = _state()
|
||||
|
||||
def wait(refs, *, num_returns, fetch_local):
|
||||
assert num_returns == 1
|
||||
events.append(("wait", fetch_local))
|
||||
clock.now += 2.0 if fetch_local else 10.0
|
||||
return refs, []
|
||||
|
||||
def get(ref):
|
||||
events.append(("get", ref))
|
||||
clock.now += 3.0
|
||||
return result
|
||||
|
||||
def synchronize():
|
||||
events.append(("sync",))
|
||||
clock.now += 5.0
|
||||
|
||||
monkeypatch.setattr(runtime_module.time, "perf_counter", lambda: clock.now)
|
||||
monkeypatch.setattr(runtime_module, "ray", SimpleNamespace(wait=wait, get=get))
|
||||
monkeypatch.setattr(torch.cuda, "is_available", lambda: True)
|
||||
monkeypatch.setattr(torch.cuda, "synchronize", synchronize)
|
||||
monkeypatch.setattr(runtime_module.envs, "FASTVIDEO_NVTX_PROFILE", False)
|
||||
monkeypatch.setattr(runtime_module.logger, "info", lambda fmt, *args: messages.append(fmt % args))
|
||||
return SimpleNamespace(clock=clock, events=events, messages=messages, result=result)
|
||||
|
||||
|
||||
def test_payload_total_includes_every_layout_tensor(timing):
|
||||
assert runtime_module._log_h3_state("encoded", timing.result) == 536
|
||||
assert any("layout.position_ids:" in line for line in timing.messages)
|
||||
assert any("layout.token_tags:" in line for line in timing.messages)
|
||||
assert all("request_id=request-7" in line for line in timing.messages)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("direction", ["A_TO_B", "B_TO_A"])
|
||||
def test_transfer_separates_production_fetch_and_completed_device_copy(timing, direction):
|
||||
ref = object()
|
||||
result = runtime_module._receive_h3_state(runtime_module._H3TransferRef(ref, "request-7"), direction)
|
||||
assert result is timing.result
|
||||
assert timing.events == [("sync",), ("wait", False), ("wait", True), ("get", ref), ("sync",)]
|
||||
line, = timing.messages
|
||||
assert f"direction={direction}" in line
|
||||
assert "request_id=request-7" in line
|
||||
assert "tensor_bytes=536" in line
|
||||
assert "source_wait_s=10.000000" in line
|
||||
assert "object_fetch_s=2.000000" in line
|
||||
assert "materialize_s=8.000000" in line
|
||||
assert "receive_s=10.000000" in line
|
||||
|
||||
|
||||
def test_disabled_profiling_does_not_wait_synchronize_or_read_clock(timing, monkeypatch):
|
||||
def unexpected_clock():
|
||||
pytest.fail("profiling disabled must not read the clock")
|
||||
|
||||
monkeypatch.setattr(runtime_module.time, "perf_counter", unexpected_clock)
|
||||
assert runtime_module._receive_h3_state(timing.result, "A_TO_B") is timing.result
|
||||
with runtime_module._h3_stage_timer("encode", "request-7", False):
|
||||
pass
|
||||
assert timing.events == []
|
||||
assert timing.messages == []
|
||||
|
||||
|
||||
def test_stage_time_includes_gpu_completion_and_excludes_previous_work(timing):
|
||||
with runtime_module._h3_stage_timer("denoise", "request-7", True):
|
||||
timing.clock.now += 7.0
|
||||
assert timing.events == [("sync",), ("sync",)]
|
||||
assert timing.messages == ["[RAY_STAGE] request_id=request-7 stage=denoise elapsed_s=12.000000"]
|
||||
|
||||
|
||||
def test_failed_transfer_propagates_without_logging_success(timing, monkeypatch):
|
||||
def fail(ref):
|
||||
raise RuntimeError("producer failed")
|
||||
|
||||
monkeypatch.setattr(runtime_module.ray, "get", fail)
|
||||
with pytest.raises(RuntimeError, match="producer failed"):
|
||||
runtime_module._receive_h3_state(runtime_module._H3TransferRef(object(), "request-7"), "A_TO_B")
|
||||
assert timing.messages == []
|
||||
|
||||
|
||||
def test_zero_fetch_duration_does_not_divide_by_zero(timing, monkeypatch):
|
||||
monkeypatch.setattr(runtime_module.time, "perf_counter", lambda: 0.0)
|
||||
runtime_module._receive_h3_state(runtime_module._H3TransferRef(object(), "request-7"), "A_TO_B")
|
||||
assert "object_fetch_mb_s=n/a" in timing.messages[0]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("stage,direction", [("denoise", "A_TO_B"), ("decode", "B_TO_A")])
|
||||
def test_actor_receives_before_model_stage(timing, stage, direction):
|
||||
actor_class = (runtime_module._MiniMaxH3DiTActor if stage == "denoise"
|
||||
else runtime_module._MiniMaxH3EncoderDecoderActor)
|
||||
actor = actor_class.__new__(actor_class)
|
||||
actor._profile_transfers = True
|
||||
|
||||
def forward(state):
|
||||
assert state is timing.result
|
||||
assert any(f"direction={direction}" in line for line in timing.messages)
|
||||
timing.clock.now += 7.0
|
||||
return state
|
||||
|
||||
actor.pipeline = SimpleNamespace(**{stage: forward})
|
||||
assert getattr(actor, stage)(runtime_module._H3TransferRef(object(), "request-7")) is timing.result
|
||||
assert any(f"stage={stage} elapsed_s=12.000000" in line for line in timing.messages)
|
||||
|
||||
|
||||
def _fake_runtime(monkeypatch, enabled):
|
||||
events = []
|
||||
|
||||
def method(name):
|
||||
def remote(*args):
|
||||
ref = SimpleNamespace(kind=name, args=args)
|
||||
events.append((name, ref))
|
||||
return ref
|
||||
|
||||
return SimpleNamespace(remote=remote)
|
||||
|
||||
def get(ref):
|
||||
# The driver must never materialize encode/denoise intermediates.
|
||||
assert ref.kind == "decode"
|
||||
events.append(("get", ref))
|
||||
return ref
|
||||
|
||||
def wait(refs, *, num_returns, fetch_local):
|
||||
assert refs[0].kind == "denoise"
|
||||
assert num_returns == 1 and fetch_local is False
|
||||
events.append(("wait", refs[0]))
|
||||
return refs, []
|
||||
|
||||
runtime = runtime_module.RayMiniMaxH3DisaggregatedRuntime.__new__(
|
||||
runtime_module.RayMiniMaxH3DisaggregatedRuntime)
|
||||
runtime._closed = False
|
||||
runtime._profile_transfers = enabled
|
||||
runtime.encoder_decoder = SimpleNamespace(encode=method("encode"), decode=method("decode"))
|
||||
runtime.dit = SimpleNamespace(denoise=method("denoise"))
|
||||
monkeypatch.setattr(runtime_module, "ray", SimpleNamespace(get=get, wait=wait))
|
||||
return runtime, events
|
||||
|
||||
|
||||
@pytest.mark.parametrize("enabled", [False, True])
|
||||
@pytest.mark.parametrize("stream", [False, True])
|
||||
def test_runtime_preserves_request_ids_and_bounded_driver_dag(monkeypatch, enabled, stream):
|
||||
runtime, events = _fake_runtime(monkeypatch, enabled)
|
||||
batches = [ForwardBatch(data_type="video", extra={"request_id": str(i)}) for i in range(3)]
|
||||
if stream:
|
||||
assert len(list(runtime.iter_forward(batches))) == 3
|
||||
assert [name for name, _ in events] == [
|
||||
"encode", "denoise", "encode", "wait", "decode", "denoise", "get",
|
||||
"encode", "wait", "decode", "denoise", "get", "decode", "get",
|
||||
]
|
||||
else:
|
||||
runtime.execute_forward(batches[0], request_id="explicit")
|
||||
assert [name for name, _ in events] == ["encode", "denoise", "decode", "get"]
|
||||
|
||||
for name, result_ref in events:
|
||||
if name not in {"denoise", "decode"}:
|
||||
continue
|
||||
arg, = result_ref.args
|
||||
assert isinstance(arg, runtime_module._H3TransferRef) is enabled
|
||||
source_ref = arg.ref if enabled else arg
|
||||
assert source_ref.kind == ("encode" if name == "denoise" else "denoise")
|
||||
if enabled:
|
||||
expected_id = source_ref.args[1] if name == "denoise" else source_ref.args[0].request_id
|
||||
assert arg.request_id == expected_id
|
||||
if not stream:
|
||||
assert arg.request_id == "explicit"
|
||||
@@ -32,6 +32,9 @@ class Executor(ABC):
|
||||
|
||||
@staticmethod
|
||||
def get_class(fastvideo_args: FastVideoArgs) -> type["Executor"]:
|
||||
if fastvideo_args.h3_disaggregated:
|
||||
from fastvideo.worker.minimax_h3_disaggregated import MiniMaxH3DisaggregatedExecutor
|
||||
return cast(type["Executor"], MiniMaxH3DisaggregatedExecutor)
|
||||
if fastvideo_args.distributed_executor_backend == "mp":
|
||||
from fastvideo.worker.multiproc_executor import MultiprocExecutor
|
||||
return cast(type["Executor"], MultiprocExecutor)
|
||||
|
||||
@@ -0,0 +1,559 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Ray data-plane for component-disaggregated MiniMax H3 inference."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import asyncio
|
||||
from collections.abc import AsyncIterator, Iterable, Iterator
|
||||
from contextlib import contextmanager, suppress
|
||||
from copy import deepcopy
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
from queue import Queue
|
||||
import sys
|
||||
import time
|
||||
from typing import Any, cast
|
||||
from uuid import uuid4
|
||||
|
||||
import torch
|
||||
|
||||
import fastvideo.envs as envs
|
||||
from fastvideo.distributed import cleanup_dist_env_and_memory
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.basic.minimax_h3.disaggregated import (
|
||||
MiniMaxH3DenoisedState,
|
||||
MiniMaxH3DiTPipeline,
|
||||
MiniMaxH3EncodedState,
|
||||
MiniMaxH3EncoderDecoderPipeline,
|
||||
MiniMaxH3RefDiTPipeline,
|
||||
MiniMaxH3RefEncoderDecoderPipeline,
|
||||
)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.lazy_module import is_lazy_module
|
||||
from fastvideo.profiler import nvtx_range
|
||||
from fastvideo.utils import get_ip, get_open_port
|
||||
from fastvideo.worker.executor import Executor
|
||||
from fastvideo.worker.ray_env import RAY_NON_CARRY_OVER_ENV_VARS
|
||||
from fastvideo.worker.ray_utils import assert_ray_available, ray
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
def _h3_state_tensors(state: Any) -> Iterator[tuple[str, torch.Tensor]]:
|
||||
for field_name in ("prompt_embeds", "video_latents", "audio_latents"):
|
||||
tensor = getattr(state, field_name, None)
|
||||
if tensor is not None:
|
||||
yield field_name, tensor
|
||||
|
||||
layout = getattr(state, "layout", None)
|
||||
if layout is not None:
|
||||
for field_name in ("position_ids", "token_tags", "text_indices", "video_indices", "audio_indices"):
|
||||
tensor = getattr(layout, field_name, None)
|
||||
if tensor is not None:
|
||||
yield f"layout.{field_name}", tensor
|
||||
|
||||
|
||||
def _log_h3_state(name: str, state: Any) -> int:
|
||||
total_bytes = 0
|
||||
for field_name, tensor in _h3_state_tensors(state):
|
||||
size = tensor.numel() * tensor.element_size()
|
||||
total_bytes += size
|
||||
logger.info(
|
||||
"[RAY_PAYLOAD] %s.%s: %.2f MB shape=%s dtype=%s device=%s request_id=%s",
|
||||
name,
|
||||
field_name,
|
||||
size / 1e6,
|
||||
tuple(tensor.shape),
|
||||
tensor.dtype,
|
||||
tensor.device,
|
||||
state.request_id,
|
||||
)
|
||||
|
||||
logger.info(
|
||||
"[RAY_PAYLOAD] %s TOTAL: %.2f MB request_id=%s (logical tensor bytes)",
|
||||
name,
|
||||
total_bytes / 1e6,
|
||||
state.request_id,
|
||||
)
|
||||
return total_bytes
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _H3TransferRef:
|
||||
"""Keep the ref nested so Ray does not fetch it before actor entry."""
|
||||
|
||||
ref: Any
|
||||
request_id: str
|
||||
|
||||
|
||||
def _synchronize_h3_device() -> None:
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.synchronize()
|
||||
|
||||
|
||||
def _receive_h3_state(state: Any, direction: str) -> Any:
|
||||
if not isinstance(state, _H3TransferRef):
|
||||
return state
|
||||
|
||||
# A receiver-local clock avoids clock skew between the two Sparks. Drain
|
||||
# earlier device work before timing, then wait for production WITHOUT
|
||||
# fetching locally so producer compute is not counted as object fetch.
|
||||
_synchronize_h3_device()
|
||||
started = time.perf_counter()
|
||||
with nvtx_range(f"h3.{direction}.source_wait"):
|
||||
ray.wait([state.ref], num_returns=1, fetch_local=False)
|
||||
source_ready = time.perf_counter()
|
||||
with nvtx_range(f"h3.{direction}.object_fetch"):
|
||||
ray.wait([state.ref], num_returns=1, fetch_local=True)
|
||||
fetched = time.perf_counter()
|
||||
with nvtx_range(f"h3.{direction}.materialize"):
|
||||
result = ray.get(state.ref)
|
||||
_synchronize_h3_device()
|
||||
materialized = time.perf_counter()
|
||||
|
||||
tensor_bytes = sum(tensor.numel() * tensor.element_size() for _, tensor in _h3_state_tensors(result))
|
||||
object_fetch_s = fetched - source_ready
|
||||
# This is an effective logical-payload rate, not measured NIC bandwidth.
|
||||
rate = f"{tensor_bytes / 1e6 / object_fetch_s:.2f}" if object_fetch_s > 0 else "n/a"
|
||||
logger.info(
|
||||
"[RAY_TRANSFER] request_id=%s direction=%s tensor_bytes=%d "
|
||||
"source_wait_s=%.6f object_fetch_s=%.6f materialize_s=%.6f "
|
||||
"receive_s=%.6f object_fetch_mb_s=%s",
|
||||
state.request_id,
|
||||
direction,
|
||||
tensor_bytes,
|
||||
source_ready - started,
|
||||
object_fetch_s,
|
||||
materialized - fetched,
|
||||
materialized - source_ready,
|
||||
rate,
|
||||
)
|
||||
return result
|
||||
|
||||
|
||||
@contextmanager
|
||||
def _h3_stage_timer(stage: str, request_id: str, enabled: bool) -> Iterator[None]:
|
||||
if not enabled:
|
||||
yield
|
||||
return
|
||||
_synchronize_h3_device()
|
||||
started = time.perf_counter()
|
||||
with nvtx_range(f"h3.{stage}"):
|
||||
yield
|
||||
_synchronize_h3_device()
|
||||
logger.info("[RAY_STAGE] request_id=%s stage=%s elapsed_s=%.6f", request_id, stage, time.perf_counter() - started)
|
||||
|
||||
|
||||
def _node_resource(node_ip: str) -> str:
|
||||
return f"node:{node_ip}"
|
||||
|
||||
|
||||
def _validate_topology(encoder_node_ip: str, dit_node_ip: str, resources: dict[str, float]) -> None:
|
||||
if not encoder_node_ip or not dit_node_ip:
|
||||
raise ValueError("MiniMax-H3 disaggregation requires both encoder and DiT node IPs.")
|
||||
if encoder_node_ip == dit_node_ip:
|
||||
raise ValueError("MiniMax-H3 encoder/decoder and DiT workers must use different Ray nodes.")
|
||||
missing = [ip for ip in (encoder_node_ip, dit_node_ip) if resources.get(_node_resource(ip), 0.0) <= 0]
|
||||
if missing:
|
||||
available = sorted(key.removeprefix("node:") for key in resources if key.startswith("node:"))
|
||||
raise RuntimeError(f"Ray has no live node resource for {missing}; available node IPs: {available}.")
|
||||
|
||||
|
||||
def _resident_role_args(source: FastVideoArgs, *, role: str) -> FastVideoArgs:
|
||||
"""Clone process-local args and force one persistent, non-parallel component role."""
|
||||
args = deepcopy(source)
|
||||
args.num_gpus = 1
|
||||
args.tp_size = 1
|
||||
args.sp_size = 1
|
||||
args.hsdp_replicate_dim = 1
|
||||
args.hsdp_shard_dim = 1
|
||||
args.ray_placement_group = None
|
||||
args.ray_runtime_env = None
|
||||
args.distributed_executor_backend = "mp"
|
||||
args.use_fsdp_inference = False
|
||||
args.dit_cpu_offload = False
|
||||
args.dit_layerwise_offload = False
|
||||
args.text_encoder_cpu_offload = False
|
||||
args.image_encoder_cpu_offload = False
|
||||
args.vae_cpu_offload = False
|
||||
args.lazy_module_load = False
|
||||
args.h3_sequential_load = False
|
||||
args.vae_parallel_encode = False
|
||||
args.vae_parallel_decode = False
|
||||
if role == "encoder_decoder":
|
||||
# Adapters target the transformer and belong exclusively on Spark B.
|
||||
args.lora_path = None
|
||||
args.enable_torch_compile = False
|
||||
elif role == "dit":
|
||||
args.enable_torch_compile_text_encoder = False
|
||||
args.enable_torch_compile_vae = False
|
||||
args.enable_torch_compile_audio_vae = False
|
||||
else:
|
||||
raise ValueError(f"Unknown MiniMax-H3 worker role: {role!r}.")
|
||||
return args
|
||||
|
||||
|
||||
def _bind_single_gpu_process() -> None:
|
||||
os.environ["LOCAL_RANK"] = "0"
|
||||
os.environ["RANK"] = "0"
|
||||
os.environ["WORLD_SIZE"] = "1"
|
||||
os.environ["MASTER_ADDR"] = "127.0.0.1"
|
||||
os.environ["MASTER_PORT"] = str(get_open_port())
|
||||
if torch.cuda.is_available():
|
||||
torch.cuda.set_device(0)
|
||||
|
||||
|
||||
def _is_ref2va(args: FastVideoArgs) -> bool:
|
||||
return args.override_pipeline_cls_name == "MiniMaxH3Ref2VAModularPipeline"
|
||||
|
||||
|
||||
def _request_id(batch: ForwardBatch) -> str:
|
||||
supplied = batch.extra.get("request_id")
|
||||
return str(supplied) if supplied is not None else uuid4().hex
|
||||
|
||||
|
||||
def _next_or_sentinel(iterator: Iterator[ForwardBatch], sentinel: object) -> ForwardBatch | object:
|
||||
return next(iterator, sentinel)
|
||||
|
||||
|
||||
def _h3_actor_runtime_env(fastvideo_args: FastVideoArgs) -> dict[str, Any]:
|
||||
"""Apply the driver's explicit FA4 choice before actors import backends.
|
||||
|
||||
Existing Ray services do not inherit later driver-shell exports. Use an
|
||||
actor runtime_env rather than changing os.environ inside the constructor:
|
||||
attention modules can already have resolved their implementation by then.
|
||||
An explicit ray_runtime_env value takes precedence over the driver shell.
|
||||
Other fields inherit from the Ray job: local working_dir/py_modules paths
|
||||
must be uploaded by ray.init and cannot be passed directly to actors.
|
||||
"""
|
||||
configured = fastvideo_args.ray_runtime_env or {}
|
||||
env_vars = dict(configured.get("env_vars", {}))
|
||||
fa4 = os.environ.get("FASTVIDEO_FA4")
|
||||
if fa4 is not None and "FASTVIDEO_FA4" not in RAY_NON_CARRY_OVER_ENV_VARS:
|
||||
env_vars.setdefault("FASTVIDEO_FA4", fa4)
|
||||
return {"env_vars": env_vars} if env_vars else {}
|
||||
|
||||
|
||||
def _log_h3_worker_environment(role: str) -> None:
|
||||
logger.info(
|
||||
"[H3_WORKER] role=%s node_ip=%s python=%s FASTVIDEO_FA4=%d",
|
||||
role,
|
||||
get_ip(),
|
||||
sys.executable,
|
||||
int(envs.FASTVIDEO_FA4),
|
||||
)
|
||||
|
||||
|
||||
class _MiniMaxH3ResidentActor:
|
||||
"""RPCs shared by both roles: readiness receipt and in-actor cleanup."""
|
||||
|
||||
_role: str
|
||||
pipeline: Any
|
||||
|
||||
def health(self) -> dict[str, Any]:
|
||||
modules = tuple(sorted(self.pipeline.modules))
|
||||
return {
|
||||
"ready": True,
|
||||
"role": self._role,
|
||||
"node_ip": get_ip(),
|
||||
"modules": modules,
|
||||
"all_resident": all(not is_lazy_module(module) for module in self.pipeline.modules.values()),
|
||||
}
|
||||
|
||||
def shutdown(self) -> dict[str, str]:
|
||||
self.pipeline = None
|
||||
cleanup_dist_env_and_memory(shutdown_ray=False)
|
||||
return {"status": "shutdown_complete"}
|
||||
|
||||
|
||||
class _MiniMaxH3EncoderDecoderActor(_MiniMaxH3ResidentActor):
|
||||
|
||||
_role = "encoder_decoder"
|
||||
|
||||
def __init__(self, fastvideo_args: FastVideoArgs, profile_transfers: bool = False) -> None:
|
||||
self._profile_transfers = profile_transfers
|
||||
_bind_single_gpu_process()
|
||||
_log_h3_worker_environment("encoder_decoder")
|
||||
args = _resident_role_args(fastvideo_args, role="encoder_decoder")
|
||||
pipeline_cls = MiniMaxH3RefEncoderDecoderPipeline if _is_ref2va(args) else MiniMaxH3EncoderDecoderPipeline
|
||||
self.pipeline = pipeline_cls(args.model_path, args)
|
||||
self.pipeline.post_init()
|
||||
|
||||
def encode(self, batch: ForwardBatch, request_id: str) -> MiniMaxH3EncodedState:
|
||||
with _h3_stage_timer("encode", request_id, self._profile_transfers):
|
||||
return self.pipeline.encode(batch, request_id=request_id)
|
||||
|
||||
def decode(self, state: MiniMaxH3DenoisedState | _H3TransferRef) -> ForwardBatch:
|
||||
state = _receive_h3_state(state, "B_TO_A")
|
||||
with _h3_stage_timer("decode", state.request_id, self._profile_transfers):
|
||||
return self.pipeline.decode(state)
|
||||
|
||||
|
||||
class _MiniMaxH3DiTActor(_MiniMaxH3ResidentActor):
|
||||
|
||||
_role = "dit"
|
||||
|
||||
def __init__(self, fastvideo_args: FastVideoArgs, profile_transfers: bool = False) -> None:
|
||||
self._profile_transfers = profile_transfers
|
||||
_bind_single_gpu_process()
|
||||
_log_h3_worker_environment("dit")
|
||||
args = _resident_role_args(fastvideo_args, role="dit")
|
||||
pipeline_cls = MiniMaxH3RefDiTPipeline if _is_ref2va(args) else MiniMaxH3DiTPipeline
|
||||
self.pipeline = pipeline_cls(args.model_path, args)
|
||||
self.pipeline.post_init()
|
||||
|
||||
def denoise(self, state: MiniMaxH3EncodedState | _H3TransferRef) -> MiniMaxH3DenoisedState:
|
||||
state = _receive_h3_state(state, "A_TO_B")
|
||||
if self._profile_transfers:
|
||||
_log_h3_state("A_TO_B encoded", state)
|
||||
with _h3_stage_timer("denoise", state.request_id, self._profile_transfers):
|
||||
result = self.pipeline.denoise(state)
|
||||
if self._profile_transfers:
|
||||
_log_h3_state("B_TO_A denoised", result)
|
||||
return result
|
||||
|
||||
# Adapters target the transformer, which lives only on this role. The
|
||||
# receipts match GPU workers so the executor can check them the same way.
|
||||
def set_lora_adapter(self,
|
||||
lora_nickname: str,
|
||||
lora_path: str | None = None,
|
||||
strength: float = 1.0,
|
||||
accumulate: bool = False) -> dict[str, str]:
|
||||
self.pipeline.set_lora_adapter(lora_nickname, lora_path, strength=strength, accumulate=accumulate)
|
||||
return {"status": "lora_adapter_set"}
|
||||
|
||||
def unmerge_lora_weights(self) -> dict[str, str]:
|
||||
self.pipeline.unmerge_lora_weights()
|
||||
return {"status": "lora_adapter_unmerged"}
|
||||
|
||||
def merge_lora_weights(self) -> dict[str, str]:
|
||||
self.pipeline.merge_lora_weights()
|
||||
return {"status": "lora_adapter_merged"}
|
||||
|
||||
|
||||
class RayMiniMaxH3DisaggregatedRuntime:
|
||||
"""Own one persistent encoder/decoder actor and one persistent DiT actor."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
*,
|
||||
encoder_node_ip: str,
|
||||
dit_node_ip: str,
|
||||
ray_address: str | None = None,
|
||||
) -> None:
|
||||
assert_ray_available()
|
||||
address = ray_address or os.environ.get("RAY_ADDRESS") or "auto"
|
||||
if not ray.is_initialized():
|
||||
ray.init(address=address, runtime_env=fastvideo_args.ray_runtime_env)
|
||||
_validate_topology(encoder_node_ip, dit_node_ip, ray.cluster_resources())
|
||||
|
||||
self._profile_transfers = envs.FASTVIDEO_H3_PROFILE_TRANSFERS
|
||||
actor_args = deepcopy(fastvideo_args)
|
||||
actor_args.ray_placement_group = None
|
||||
actor_args.ray_runtime_env = None
|
||||
common_options = {
|
||||
"num_cpus": 0,
|
||||
"num_gpus": 1,
|
||||
"max_restarts": 0,
|
||||
"runtime_env": _h3_actor_runtime_env(fastvideo_args),
|
||||
}
|
||||
self.encoder_node_ip = encoder_node_ip
|
||||
self.dit_node_ip = dit_node_ip
|
||||
self._closed = False
|
||||
self.encoder_decoder = None
|
||||
self.dit = None
|
||||
try:
|
||||
self.encoder_decoder = ray.remote(_MiniMaxH3EncoderDecoderActor).options(
|
||||
**common_options,
|
||||
resources={
|
||||
_node_resource(encoder_node_ip): 0.001
|
||||
},
|
||||
).remote(actor_args, self._profile_transfers)
|
||||
self.dit = ray.remote(_MiniMaxH3DiTActor).options(
|
||||
**common_options,
|
||||
resources={
|
||||
_node_resource(dit_node_ip): 0.001
|
||||
},
|
||||
).remote(actor_args, self._profile_transfers)
|
||||
self._validate_workers()
|
||||
except Exception:
|
||||
for actor in (self.encoder_decoder, self.dit):
|
||||
if actor is not None:
|
||||
with suppress(Exception):
|
||||
ray.kill(actor, no_restart=True)
|
||||
raise
|
||||
|
||||
def _validate_workers(self) -> None:
|
||||
health = self.health()
|
||||
expected = {
|
||||
"encoder_decoder": set(MiniMaxH3EncoderDecoderPipeline._required_config_modules),
|
||||
"dit": set(MiniMaxH3DiTPipeline._required_config_modules),
|
||||
}
|
||||
for receipt in health:
|
||||
if not receipt["all_resident"]:
|
||||
raise RuntimeError(f"MiniMax-H3 {receipt['role']} worker contains a deferred component.")
|
||||
if set(receipt["modules"]) != expected[receipt["role"]]:
|
||||
raise RuntimeError(f"MiniMax-H3 {receipt['role']} worker loaded {receipt['modules']}, "
|
||||
f"expected {tuple(sorted(expected[receipt['role']]))}.")
|
||||
actual = {receipt["role"]: receipt["node_ip"] for receipt in health}
|
||||
requested = {"encoder_decoder": self.encoder_node_ip, "dit": self.dit_node_ip}
|
||||
if actual != requested:
|
||||
raise RuntimeError(f"Ray placed MiniMax-H3 roles on {actual}, not the requested nodes {requested}.")
|
||||
|
||||
def health(self) -> list[dict[str, Any]]:
|
||||
if self.encoder_decoder is None or self.dit is None:
|
||||
raise RuntimeError("MiniMax-H3 disaggregated workers have not been created.")
|
||||
return ray.get([self.encoder_decoder.health.remote(), self.dit.health.remote()])
|
||||
|
||||
def _transfer_arg(self, ref: Any, request_id: str) -> Any:
|
||||
return _H3TransferRef(ref, request_id) if self._profile_transfers else ref
|
||||
|
||||
def submit(self, batch: ForwardBatch, *, request_id: str | None = None):
|
||||
"""Build a direct actor-to-actor DAG without materializing intermediates on the driver."""
|
||||
if self._closed:
|
||||
raise RuntimeError("MiniMax-H3 disaggregated runtime is closed.")
|
||||
if self.encoder_decoder is None or self.dit is None:
|
||||
raise RuntimeError("MiniMax-H3 disaggregated workers have not been created.")
|
||||
resolved_request_id = request_id if request_id is not None else _request_id(batch)
|
||||
encoded_ref = self.encoder_decoder.encode.remote(batch, resolved_request_id)
|
||||
denoised_ref = self.dit.denoise.remote(self._transfer_arg(encoded_ref, resolved_request_id))
|
||||
return self.encoder_decoder.decode.remote(self._transfer_arg(denoised_ref, resolved_request_id))
|
||||
|
||||
def execute_forward(self, batch: ForwardBatch, *, request_id: str | None = None) -> ForwardBatch:
|
||||
return ray.get(self.submit(batch, request_id=request_id))
|
||||
|
||||
async def execute_forward_async(self, batch: ForwardBatch, *, request_id: str | None = None) -> ForwardBatch:
|
||||
return await asyncio.to_thread(self.execute_forward, batch, request_id=request_id)
|
||||
|
||||
def iter_forward(self, batches: Iterable[ForwardBatch]) -> Iterator[ForwardBatch]:
|
||||
"""Run a bounded one-request lookahead pipeline over the two actors.
|
||||
|
||||
Spark A encodes the next request while Spark B denoises the current one.
|
||||
Once the current denoise completes, its decode is queued on Spark A and
|
||||
the next denoise starts immediately on Spark B. Intermediate ObjectRefs
|
||||
are never fetched by the driver.
|
||||
"""
|
||||
iterator = iter(batches)
|
||||
try:
|
||||
first = next(iterator)
|
||||
except StopIteration:
|
||||
return
|
||||
|
||||
if self.encoder_decoder is None or self.dit is None:
|
||||
raise RuntimeError("MiniMax-H3 disaggregated workers have not been created.")
|
||||
request_id = _request_id(first)
|
||||
encoded_ref = self.encoder_decoder.encode.remote(first, request_id)
|
||||
denoised_ref = self.dit.denoise.remote(self._transfer_arg(encoded_ref, request_id))
|
||||
for next_batch in iterator:
|
||||
next_request_id = _request_id(next_batch)
|
||||
next_encoded_ref = self.encoder_decoder.encode.remote(next_batch, next_request_id)
|
||||
ray.wait([denoised_ref], num_returns=1, fetch_local=False)
|
||||
decoded_ref = self.encoder_decoder.decode.remote(self._transfer_arg(denoised_ref, request_id))
|
||||
next_denoised_ref = self.dit.denoise.remote(self._transfer_arg(next_encoded_ref, next_request_id))
|
||||
yield ray.get(decoded_ref)
|
||||
denoised_ref = next_denoised_ref
|
||||
request_id = next_request_id
|
||||
yield ray.get(self.encoder_decoder.decode.remote(self._transfer_arg(denoised_ref, request_id)))
|
||||
|
||||
async def iter_forward_async(self, batches: Iterable[ForwardBatch]) -> AsyncIterator[ForwardBatch]:
|
||||
"""Asynchronously consume the bounded actor pipeline without blocking the event loop."""
|
||||
iterator = iter(self.iter_forward(batches))
|
||||
sentinel = object()
|
||||
while True:
|
||||
result = await asyncio.to_thread(_next_or_sentinel, iterator, sentinel)
|
||||
if result is sentinel:
|
||||
return
|
||||
yield cast(ForwardBatch, result)
|
||||
|
||||
def close(self) -> None:
|
||||
if self._closed:
|
||||
return
|
||||
self._closed = True
|
||||
actors = [actor for actor in (self.encoder_decoder, self.dit) if actor is not None]
|
||||
try:
|
||||
ray.get([actor.shutdown.remote() for actor in actors])
|
||||
finally:
|
||||
for actor in actors:
|
||||
with suppress(Exception):
|
||||
ray.kill(actor, no_restart=True)
|
||||
|
||||
|
||||
def _check_receipt(receipt: dict[str, Any], expected: str) -> None:
|
||||
if receipt.get("status") != expected:
|
||||
raise RuntimeError(f"MiniMax-H3 DiT worker did not confirm {expected}: {receipt}.")
|
||||
|
||||
|
||||
class MiniMaxH3DisaggregatedExecutor(Executor):
|
||||
"""VideoGenerator-compatible adapter over the two-role Ray runtime."""
|
||||
|
||||
def _init_executor(self) -> None:
|
||||
encoder_ip = self.fastvideo_args.h3_encoder_node_ip
|
||||
dit_ip = self.fastvideo_args.h3_dit_node_ip
|
||||
if encoder_ip is None or dit_ip is None:
|
||||
raise ValueError("Set h3_encoder_node_ip and h3_dit_node_ip for component disaggregation.")
|
||||
self.runtime = RayMiniMaxH3DisaggregatedRuntime(
|
||||
self.fastvideo_args,
|
||||
encoder_node_ip=encoder_ip,
|
||||
dit_node_ip=dit_ip,
|
||||
ray_address=self.fastvideo_args.h3_ray_address,
|
||||
)
|
||||
|
||||
def execute_forward(self, forward_batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
del fastvideo_args
|
||||
return self.runtime.execute_forward(forward_batch)
|
||||
|
||||
async def execute_forward_async(self, forward_batch: ForwardBatch) -> ForwardBatch:
|
||||
return await self.runtime.execute_forward_async(forward_batch)
|
||||
|
||||
def iter_forward(self, batches: Iterable[ForwardBatch]) -> Iterator[ForwardBatch]:
|
||||
return self.runtime.iter_forward(batches)
|
||||
|
||||
def iter_forward_async(self, batches: Iterable[ForwardBatch]) -> AsyncIterator[ForwardBatch]:
|
||||
return self.runtime.iter_forward_async(batches)
|
||||
|
||||
def set_lora_adapter(self,
|
||||
lora_nickname: str,
|
||||
lora_path: str | None = None,
|
||||
strength: float = 1.0,
|
||||
accumulate: bool = False) -> None:
|
||||
_check_receipt(
|
||||
ray.get(self.runtime.dit.set_lora_adapter.remote(lora_nickname, lora_path, strength, accumulate)),
|
||||
"lora_adapter_set")
|
||||
|
||||
def unmerge_lora_weights(self) -> None:
|
||||
_check_receipt(ray.get(self.runtime.dit.unmerge_lora_weights.remote()), "lora_adapter_unmerged")
|
||||
|
||||
def merge_lora_weights(self) -> None:
|
||||
_check_receipt(ray.get(self.runtime.dit.merge_lora_weights.remote()), "lora_adapter_merged")
|
||||
|
||||
def collective_rpc(self,
|
||||
method,
|
||||
timeout: float | None = None,
|
||||
args: tuple = (),
|
||||
kwargs: dict[str, Any] | None = None) -> list[Any]:
|
||||
raise NotImplementedError("Component-disaggregated H3 has role-specific RPC; use runtime actor handles.")
|
||||
|
||||
def set_log_queue(self, log_queue: Queue | None) -> None:
|
||||
# multiprocessing.Queue cannot cross Ray nodes. Actor logs remain in the
|
||||
# Ray session log, matching RayDistributedExecutor's behavior.
|
||||
self._log_queue = log_queue
|
||||
|
||||
def clear_log_queue(self) -> None:
|
||||
self._log_queue = None
|
||||
|
||||
def shutdown(self) -> None:
|
||||
runtime = getattr(self, "runtime", None)
|
||||
if runtime is not None:
|
||||
runtime.close()
|
||||
|
||||
def __del__(self) -> None:
|
||||
with suppress(Exception):
|
||||
self.shutdown()
|
||||
|
||||
|
||||
__all__ = [
|
||||
"MiniMaxH3DisaggregatedExecutor",
|
||||
"RayMiniMaxH3DisaggregatedRuntime",
|
||||
]
|
||||
Reference in New Issue
Block a user