Compare commits

...
Author SHA1 Message Date
JerryZhou54 bdaf2e2e45 Support custom action trajectories for validation 2026-02-03 20:23:13 +00:00
JerryZhou54 f3a7a37312 Overfitting running for MC 10 2026-02-03 05:11:48 +00:00
JerryZhou54 b1abb5c251 Update action_labels 2026-02-02 05:20:06 +00:00
RandNMR73 a80824255e Add data preprocess pipeline for WanGame 2026-02-02 04:38:24 +00:00
Kaiqin Kong caa1c402ba [bugfix] Double Normalization in Preprocessing Dataset (#1055) 2026-01-31 11:11:59 -08:00
XOR-op 38a6bd93d3 [chore]: update sageattn3 installation instructions (#1050) 2026-01-29 15:45:52 -08:00
alexzms b867ef7e7c [SP Sharding] Fix SP loss sharding on token axis (thw) with padding; add distributed correctness tests (#1045)
Fixes the sequence parallel sharding on t, now SP shards on t*h*w
2026-01-27 22:53:04 -08:00
William Lin 3ae58c277a [docs] Update design overview and add agents tutorial (#1044) 2026-01-27 15:56:58 -08:00
Kaiqin Kong 0c6862ca55 [feature] Add Matrix Game 2.0 training (#1017)
The CI tests are quite unstable, but since multiple CI tests indicates that each individual tests are passed, I think we can merge this.
2026-01-26 19:21:43 -08:00
XOR-opandWill Lin e8c854bcf1 [docs] Offloading instruction (#1022)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-01-26 13:40:28 -08:00
William Lin 06860e96fe [docs] Update runpod instructions (#1043) 2026-01-26 13:13:54 -08:00
83 changed files with 10097 additions and 553 deletions
Binary file not shown.

After

Width:  |  Height:  |  Size: 211 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 117 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 18 KiB

After

Width:  |  Height:  |  Size: 461 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 210 KiB

Binary file not shown.

Before

Width:  |  Height:  |  Size: 27 KiB

After

Width:  |  Height:  |  Size: 437 KiB

+2
View File
@@ -6,6 +6,8 @@ FastVideo provides highly optimized custom attention kernels to accelerate video
* **[Video Sparse Attention (VSA)](vsa/index.md)**: Sparse attention mechanism selecting top-k blocks.
* **[Sliding Tile Attention (STA)](sta/index.md)**: Optimized attention for window-based video generation.
* **Backend development guide**: See the developer guide at
[Attention Backend Development](../contributing/attention_backend.md).
## General Build Instructions
+176
View File
@@ -0,0 +1,176 @@
# Attention Backend Development
This guide is for contributors adding a new attention backend (kernel or
implementation) to FastVideo. If you just want to use existing kernels or build
fastvideo-kernel, see [Attention overview](../attention/index.md).
## When you need this guide
Use this guide when you are:
- Adding a new attention kernel or algorithm.
- Wiring an existing kernel into FastVideo's attention selection.
- Extending attention support to a new platform.
## 0) Choose a backend name and scope
Pick a backend name in `UPPER_SNAKE_CASE` and decide where it should run.
Example: `MY_NEW_ATTN`.
You will use this name in:
- `AttentionBackendEnum` (global list of backends).
- `get_name()` in your backend class (match the enum name).
- Platform selectors (CUDA/ROCm/MPS/NPU) to return your backend class.
## 1) Add enum + platform selection
1) Add your backend to `fastvideo/platforms/interface.py`:
```python
class AttentionBackendEnum(enum.Enum):
...
MY_NEW_ATTN = enum.auto()
```
1) Register it in platform selection (example: CUDA). Update
`fastvideo/platforms/cuda.py` inside `get_attn_backend_cls`:
```python
elif selected_backend == AttentionBackendEnum.MY_NEW_ATTN:
try:
from fastvideo.attention.backends.my_new_attn import MyNewAttnBackend
return "fastvideo.attention.backends.my_new_attn.MyNewAttnBackend"
except ImportError as e:
logger.error("Failed to import MY_NEW_ATTN backend: %s", str(e))
raise
```
If you want support on other platforms, add a similar branch in
`fastvideo/platforms/rocm.py`, `fastvideo/platforms/mps.py`, or `fastvideo/platforms/npu.py`.
## 2) Implement the backend
Create `fastvideo/attention/backends/my_new_attn.py` and implement the required
classes.
Minimal skeleton (no custom metadata):
```python
import torch
from dataclasses import dataclass
from fastvideo.attention.backends.abstract import (
AttentionBackend,
AttentionImpl,
AttentionMetadata,
AttentionMetadataBuilder,
)
class MyNewAttnBackend(AttentionBackend):
@staticmethod
def get_name() -> str:
return "MY_NEW_ATTN"
@staticmethod
def get_impl_cls() -> type["MyNewAttnImpl"]:
return MyNewAttnImpl
@staticmethod
def get_metadata_cls() -> type["AttentionMetadata"]:
return AttentionMetadata
@staticmethod
def get_builder_cls() -> type["AttentionMetadataBuilder"]:
return MyNewAttnMetadataBuilder
@dataclass
class MyNewAttnMetadata(AttentionMetadata):
current_timestep: int
class MyNewAttnMetadataBuilder(AttentionMetadataBuilder):
def __init__(self) -> None:
pass
def prepare(self) -> None:
pass
def build(self, current_timestep: int, **kwargs):
return MyNewAttnMetadata(current_timestep=current_timestep)
class MyNewAttnImpl(AttentionImpl):
def __init__(
self,
num_heads: int,
head_size: int,
softmax_scale: float,
causal: bool = False,
num_kv_heads: int | None = None,
prefix: str = "",
**extra_impl_args,
) -> None:
self.softmax_scale = softmax_scale
self.causal = causal
def forward(
self,
query: torch.Tensor,
key: torch.Tensor,
value: torch.Tensor,
attn_metadata: MyNewAttnMetadata,
) -> torch.Tensor:
# Implement attention
return torch.nn.functional.scaled_dot_product_attention(
query.transpose(1, 2),
key.transpose(1, 2),
value.transpose(1, 2),
is_causal=self.causal,
scale=self.softmax_scale,
).transpose(1, 2)
```
Optional:
- Implement `preprocess_qkv` / `postprocess_output` if your kernel needs tiling
or reshaping.
- Use `fastvideo.forward_context.get_forward_context()` if you need dynamic
per-step data (e.g., window sizes).
- Set `accept_output_buffer = True` if your backend writes into a provided
output buffer.
## 3) Wire into attention layers
Backends are used by `LocalAttention` and `DistributedAttention`. These layers
accept a `supported_attention_backends` tuple. If your backend should be
eligible, update the call sites that construct these layers (search for
`supported_attention_backends=`).
## 4) Add compiled kernels (optional)
If you have a custom CUDA kernel:
1) Add sources in `fastvideo-kernel/csrc/attention/`.
2) Register bindings in `fastvideo-kernel/csrc/common_extension.cpp`.
3) Add to `fastvideo-kernel/CMakeLists.txt` (and any feature flags).
4) Expose in `fastvideo-kernel/python/fastvideo_kernel/ops.py`.
5) Export in `fastvideo-kernel/python/fastvideo_kernel/__init__.py`.
Keep a Python/Triton fallback so the backend runs even when the kernel is not
available.
## 5) Testing and debugging
- Add a small parity test or microbenchmark comparing to SDPA.
- Force your backend with the env var:
`FASTVIDEO_ATTENTION_BACKEND=MY_NEW_ATTN`.
- Check logs from `fastvideo/attention/selector.py` to confirm selection.
## Checklist
- [ ] Added enum entry in `fastvideo/platforms/interface.py`.
- [ ] Implemented backend in `fastvideo/attention/backends/`.
- [ ] Registered selection in platform(s).
- [ ] Updated layer call sites to include the backend where appropriate.
- [ ] Added tests and documentation.
+480
View File
@@ -0,0 +1,480 @@
# FastVideo + Coding Agents
Coding agents are now strong at navigating large codebases and iterating fast
with parity tests and examples. This guide shows how to use them to add new
model pipelines and ship PRs in a production-grade video diffusion framework.
FastVideo is a great project to contribute to, with production-grade
infrastructure, active collaborations (including NVIDIA), and a pipeline design
and inference architecture that has been forked by [SGLang’s
multimodal generation stack](https://github.com/sgl-project/sglang/tree/main/python/sglang/multimodal_gen).
Goal: run the new pipeline with a minimal script like
`examples/inference/basic/basic.py`. In production, FastVideo can download
models automatically via `HF_HOME`; for development, use local directories so
agents can run scripts and tests deterministically. We standardize local paths
as:
- `official_weights/<model_name>/` for official checkpoints
- `converted_weights/<model_name>/` if conversion is required
## Tips when prompting the agent
When prompting the agent, include:
- This guide and the [FastVideo design overview](../design/overview.md).
- Exact file paths to edit.
- A closest reference example file in FastVideo.
- Expected behavior and acceptance criteria.
- Repro steps (command, inputs, logs).
- Constraints (performance, memory, compatibility).
- Local paths (e.g., `official_weights/<model_name>/` or
`converted_weights/<model_name>/`) for parity tests.
## FastVideo structure at a glance
Before diving in, scan these references:
- [Contributing overview](overview.md) for environment/setup context.
- [FastVideo design overview](../design/overview.md) for pipeline architecture, configs, and HF layout.
FastVideo maps a Diffusers-style repo into a pipeline like:
- `fastvideo/models/*`: model implementations (DiT, VAE, encoders, upsamplers).
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
weight name translation.
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
- `fastvideo/pipelines/basic/*`: end-to-end pipeline logic built from stages.
- `model_index.json`: the HF repo entrypoint that maps component names to
classes and weight files.
- Component loading happens in `VideoGenerator.from_pretrained`, which reads
`model_index.json`, resolves configs, and loads weights.
Minimal usage example (based on `examples/inference/basic/basic.py`):
```python
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
sampling = SamplingParam.from_pretrained(model_id)
sampling.num_frames = 45
video = generator.generate_video(
"A vibrant city street at sunset.",
sampling_param=sampling,
output_path="video_samples",
save_video=True,
)
```
## Some questions to ask yourself before starting
Answering these upfront clarifies the work and speeds up implementation.
### Is the model already supported by SGLang's multimodal generation stack?
If yes, you can port many components from SGLang. It is a FastVideo fork, so
interfaces line up, but you still need to swap layers/modules to match
FastVideo's architecture and attention stack.
If not, implement the model directly in FastVideo.
### Is there an official implementation of the model you are adding?
If yes, use it as the numerical reference. For example, LTX‑2 has an official
implementation here: https://github.com/Lightricks/LTX-2. Prefer official code
even if Diffusers also has one.
### Is there a HuggingFace repo for the model you are adding? Is it in Diffusers format?
If yes, load it directly in FastVideo after setting tensor mapping rules in the
config. Otherwise, convert the weights to Diffusers format. See [Weights and
Diffusers format](../design/overview.md#weights-and-diffusers-format) for details.
### Do I have official weights + local paths ready?
Standardize local paths as:
- `official_weights/<model_name>/` for official checkpoints
- `converted_weights/<model_name>/` if conversion is required (can be created later)
### What pipeline components are required for the model you are adding?
Usually you need a transformer (DiT), VAE, text encoder, and tokenizer. Some
models add extra components.
### What tasks does the model support?
Usually a video diffusion model supports text‑to‑video (T2V),
image‑to‑video (I2V), and video‑to‑video (V2V). Some add extra tasks (two‑stage
generation, keyframe interpolation), which require extra components.
It's usually easiest to start with a T2V pipeline and add the other tasks later.
You can refer to the [Pipeline system](../design/overview.md#pipeline-system)
section for more details.
### Am I able to generate videos with the official implementation?
These videos and prompts are your reference. Once the FastVideo pipeline works,
compare outputs to the official implementation. Due to seeding and other
factors, outputs may not match exactly, but they should be comparable.
## Workflow: adding a full pipeline
This is an example workflow for adding a full model pipeline (model + configs +
examples + tests). This guide is in active development; feedback is welcome.
!!! note
If you get stuck, refer to existing models/pipelines in FastVideo or ask in Slack.
### 0) Fetch official model's code and weights
Purpose:
- Keep official checkpoints and source code local so conversion, parity tests,
and reference runs are reproducible.
- Clone the official repo so you can use it as a numerical reference.
Action:
- Download official weights into `official_weights/<model_name>/`
(Diffusers format or not).
- Clone the official repo under the project root (e.g., `FastVideo/LTX-2/`).
- If a Diffusers-format HF repo already exists, you can skip manual weight
handling and download it directly with
`scripts/huggingface/download_hf.py`.
!!! note
This step is best done manually because large downloads can time out.
Example:
```bash
python scripts/huggingface/download_hf.py \
--repo_id Wan-AI/Wan2.1-T2V-1.3B-Diffusers \
--local_dir official_weights/Wan2.1-T2V-1.3B-Diffusers \
--repo_type model
```
### 1) Implement the model + config mapping
Purpose:
- Model weights are a dictionary of named tensors (`state_dict`). If the names
don’t line up with FastVideo’s module names, weights won’t load correctly (or
will silently load into the wrong layer).
- Official checkpoints often use different prefixes or module layouts than
FastVideo, so we translate names via the mapping (during load or conversion).
- Mapping aligns three things:
1. the official implementation’s module names,
2. the checkpoint `state_dict` keys,
3. FastVideo’s model classes and layer naming conventions.
- If names don’t align, weights won’t load; implement the FastVideo model and
define mapping rules first.
Action:
- Implement the FastVideo model + config mapping.
- Add/extend the model in `fastvideo/models/...` and config in
`fastvideo/configs/models/...` (including `param_names_mapping`).
- Reuse existing FastVideo layers/modules where possible.
- Use FastVideo’s attention layers:
- `DistributedAttention` only for full‑sequence self‑attention in the DiT.
- `LocalAttention` for cross‑attention and other attention layers.
- See the “Configuration System” and “Weights and Diffusers format” sections
in `docs/design/overview.md` for how these pieces connect.
- If you are using an agent, ask it to implement the model, config mapping,
and a parity test together so you can validate numerics immediately.
!!! note
After the first component is aligned and parity‑tested, open a **DRAFT PR**
on FastVideo so the rest of the pipeline work can build on top of it.
!!! note
If a Diffusers-format HF repo already exists and loads correctly, you can
skip conversion entirely (no conversion script needed) and just download it
with `scripts/huggingface/download_hf.py`. Otherwise, you may need a
conversion script + a `converted_weights/<model>/` staging directory.
Example (key renaming via arch config mapping, Wan2.1‑style):
```python
# Official model (simplified) in the upstream repo.
class OfficialWanTransformer(torch.nn.Module):
def __init__(self):
super().__init__()
self.patch_embedding = torch.nn.Conv3d(16, 1536, kernel_size=2, padding=0)
def forward(self, x):
return self.patch_embedding(x)
# FastVideo model (simplified) in fastvideo/models/dits/wanvideo.py
class PatchEmbed(torch.nn.Module):
def __init__(self):
super().__init__()
self.proj = torch.nn.Conv3d(16, 1536, kernel_size=2, padding=0)
def forward(self, x):
return self.proj(x)
class WanTransformer3DModel(torch.nn.Module):
def __init__(self):
super().__init__()
self.patch_embedding = PatchEmbed()
def forward(self, x):
return self.patch_embedding(x)
# Mapping defined in a config (simplified; see the real mapping in
# fastvideo/configs/models/dits/wanvideo.py)
param_names_mapping = {
r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$": r"blocks.\1.to_q.\2",
}
def apply_regex_map(state_dict, mapping):
# Pseudocode: apply regex substitutions in order
...
# Official checkpoint keys (example)
official = {
"patch_embedding.weight": ...,
"blocks.0.attn1.to_q.weight": ...,
}
# Apply mapping so keys match FastVideo modules
converted = apply_regex_map(official, param_names_mapping)
```
Optional helper (print a few checkpoint keys quickly):
```bash
python - <<'PY'
import safetensors.torch as st
keys = list(st.load_file("official_weights/<model>/transformer/diffusion_pytorch_model.safetensors").keys())
print(keys[:20])
PY
```
Example agent prompt (task request):
```
Please add the Wan2.1 T2V 1.3B Diffusers pipeline to FastVideo:
- Add a FastVideo native Wan2.1 DiT implementation + config mapping.
- Make sure to use the existing FastVideo layers and attention modules where possible.
- Add a parity test that loads the official model alongside the FastVideo model and compares outputs numerically with fixed seeds and inputs.
Paths:
- Official repo: Wan-AI/Wan2.1-T2V-1.3B-Diffusers
- Local download: official_weights/Wan2.1-T2V-1.3B-Diffusers
Mapping steps:
- Load the official DiT weights from
official_weights/Wan2.1-T2V-1.3B-Diffusers/transformer/diffusion_pytorch_model.safetensors.
- Instantiate the FastVideo DiT (`WanTransformer3DModel`) and compare
its `state_dict().keys()` to the official keys.
- Update `param_names_mapping` in
fastvideo/configs/models/dits/wanvideo.py to resolve missing/unexpected keys.
- Use `load_state_dict(strict=False)` during iteration to surface mismatches.
```
External examples of the same pattern:
- SGLang uses prefix-based routing in its weight loader to map checkpoint keys
into internal submodules (e.g., stripping a top-level prefix before delegating).
- vLLM includes model-specific renamers for certain checkpoints that adjust
key prefixes so weights match its internal naming.
### 2) Test numerical alignment with the official implementation
Purpose:
- Verify that the FastVideo component is numerically aligned with the official
implementation.
Action:
- Add or reuse a numerical parity test that loads the official model and the
FastVideo model and compares outputs.
- See examples in `tests/local_tests/` (e.g., `tests/local_tests/upsamplers/`)
and the commands in `tests/local_tests/README.md`.
- If there are discrepancies, add opt‑in logging to both models and compare
activation summaries (layer output sums, per‑stage logs).
- First align the loaded weights (validate `param_names_mapping`).
- Then align forward outputs using fixed seeds and inputs.
- Start with `atol=1e-4, rtol=1e-4` in `assert_close`.
- Keep dtype consistent (bf16 if available; otherwise fp32).
- If attention parity is unstable, align backends (e.g.,
`FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA`).
### 3) Repeat the process for each component
If the model requires additional components, repeat Steps 1–2 for each one.
For example, implement the VAE in `fastvideo/models/vaes/` and its config in
`fastvideo/configs/models/vaes/`, then add parity coverage for it.
### 4) Add a pipeline config + sample defaults
Purpose:
- `fastvideo/configs/pipelines/` describes pipeline wiring and model module
names.
- `fastvideo/configs/sample/` defines default runtime parameters.
Action:
- Add a new pipeline config + sampling params.
- Register them in `fastvideo/configs/pipelines/registry.py` and
`fastvideo/configs/sample/registry.py`.
### 5) Wire pipeline stages
Purpose:
- `fastvideo/pipelines/basic/<pipeline>/` contains the actual pipeline logic.
- `fastvideo/pipelines/stages/` holds reusable, testable stages.
Action:
- Build the pipeline using stages; keep new stages isolated and documented.
- Prefer opt‑in flags for expensive or optional steps.
### 6) Add pipeline‑level tests
Purpose:
- Ensure the end‑to‑end pipeline works and stays aligned as pieces evolve.
Action:
- Add a pipeline parity test under `tests/local_tests/pipelines/`.
- See the [Testing Guide](testing.md) for test conventions.
### 7) Add user‑facing examples
Purpose:
- `examples/inference/basic/` is the entry point for simple, runnable scripts.
Action:
- Provide a minimal “hello world” example plus advanced variations.
- Use fixed seeds and stable prompts.
- Run the example locally to confirm end‑to‑end behavior.
### 8) Add SSIM tests for CI checks
Purpose:
- Ensure visual similarity stays within expected bounds for regression testing.
- SSIM tests act as a higher‑level guardrail beyond unit/parity tests.
Action:
- Add SSIM tests under `fastvideo/tests/ssim/` and include reference videos
(see the structure in the Testing Guide).
- Use stable prompts/seeds and document any GPU‑specific requirements.
- Follow the [Testing Guide](testing.md) for reference video placement and
execution details.
### 9) Document it
Purpose:
- `docs/` is where users find the new pipeline usage and limitations.
Action:
- Add a short doc page or update an existing one.
- Mention any caveats (memory, speed, constraints).
## Common pitfalls when porting models
- **Attention backend mismatch**: parity can fail if the official model uses a
different attention backend (e.g., SDPA vs custom). Align backends before
debugging deeper issues.
- **Patchifier shape mistakes**: wrong patchification or reshape lengths can
silently corrupt outputs. Validate patch shapes early.
- **Mask handling**: attention masks must match the official behavior (padding,
causal masks, and broadcast shapes).
- **Scheduler / sigma schedule mismatch**: even small differences in schedules
or timestep shapes can cause noticeable drift.
## Diffusers vs manual conversion
If a model already ships in Diffusers format (with a proper `model_index.json`),
prefer downloading it directly and loading it via FastVideo. In that case:
- You usually **do not need** a conversion script.
- You still need a correct `param_names_mapping` if the internal module names
differ from FastVideo’s implementation.
If the model does **not** have a Diffusers-format repo:
- You will need a conversion script to rewrite `state_dict` keys into FastVideo
naming and stage the result (e.g., under `converted_weights/<model>/`).
- You may still use the official repo for reference parity and debugging.
In both cases, parity testing is required to validate correctness.
If you want to publish a Diffusers‑style repo after conversion, use
`scripts/checkpoint_conversion/create_hf_repo.py` to assemble a HuggingFace‑ready
directory before uploading.
## FAQ
**Q: Why do we implement the FastVideo model before conversion?**
A: You can’t define the key‑mapping rules until the FastVideo module names are
known. The implementation determines the target `state_dict` schema.
**Q: Do we always need a conversion script?**
A: No. If a Diffusers‑format repo exists and loads correctly, download it and
skip conversion.
**Q: How do I figure out `param_names_mapping` quickly?**
A: Load the official weights, instantiate the FastVideo model, and diff
`state_dict().keys()` on both sides. Add regex rules until missing/unexpected
keys are resolved. Agents can help you with this.
**Q: What if parity fails even after mapping?**
A: Align attention backends, sigma schedules, and timestep shapes first. Then
add opt‑in activation logging to locate the first divergent layer.
## Case study: LTX‑2 port (from PLAN.md)
The LTX‑2 port in `PLAN.md` shows the real sequence of steps and backtracking
that happened during integration. Use it as a reference for how parity work
actually unfolds:
- Ported components first (transformer, VAE, audio, text encoder).
- Added parity tests per component; used SDPA for reference parity.
- Added debug logging to compare per‑block activations and isolate divergence.
- Fixed cross‑attention reshape and patch grid bounds issues after logging.
- Aligned sigma schedule and masking behavior to match the official pipeline.
Recommendation:
- Keep raw step‑by‑step logs in your own local `PLAN.md` for large ports.
## Worked example: Wan2.1 T2V 1.3B pipeline
The Wan2.1 T2V 1.3B Diffusers pipeline is a good “standard” example for
FastVideo integration.
1. Verify model config + mapping.
- DiT mapping: `fastvideo/configs/models/dits/wanvideo.py`
- VAE: `fastvideo/models/vaes/wanvae.py`
- Text encoder: `fastvideo/models/encoders/t5.py`
2. Parity test the core components.
- Example tests: `fastvideo/tests/transformers/test_wanvideo.py`,
`fastvideo/tests/vaes/test_wan_vae.py`,
`fastvideo/tests/encoders/test_t5_encoder.py`
3. Pipeline wiring.
- Pipeline: `fastvideo/pipelines/basic/wan/wan_pipeline.py`
- Pipeline config: `fastvideo/configs/pipelines/wan.py`
- Sampling defaults: `fastvideo/configs/sample/wan.py`
4. Minimal example.
- Script: `examples/inference/basic/basic.py`
+43 -31
View File
@@ -1,16 +1,55 @@
# 📦 Developing FastVideo on RunPod
You can easily use the FastVideo Docker image as a custom container on [RunPod](https://www.runpod.io) for development or experimentation.
You can easily use the FastVideo Pod Template on [RunPod](https://www.runpod.io) for development or experimentation.
## Creating a new pod
Choose a GPU that supports CUDA 12.8
Pick 1 or 2 L40S GPU(s)
- Make sure you are using the correct RunPod account.
![RunPod Account Selection](../../assets/images/runpod_account.png)
- Use "Additional Filters" to select CUDA 12.8.
![RunPod CUDA selection](../../assets/images/runpod_cuda.png)
- Click "Deploy" and Pick a single A40 or RTX 4090 GPU.
![RunPod GPU Selection](../../assets/images/runpod_deploy.png)
- Select the "FastVideo" or "fastvideo-dev" Pod Template.
![RunPod Pod Template Selection](../../assets/images/runpod_create.png)
- Set the Pod name to "<name>-<FastVideo>-<date>".
- Finally, once the pod is deployed (will take a few minutes as the image is being pulled), you can SSH into it using "SSH exposed over TCP". You'll need to use the matching private ssh key you provided.
![RunPod SSH](../../assets/images/runpod_ssh.png)
## Working with the pod
After SSH'ing into your pod, you'll find the correct `uv` environment already activated and you should be in /FastVideo/ directory. Make sure to use /FastVideo/ for all your work.
To pull in the latest changes from the GitHub repo:
```bash
cd /FastVideo
git pull
```
Run your development workflows as usual:
```bash
# Run linters
pre-commit run --all-files
# Run tests
pytest tests/
```
Make sure to push your changes back to the GitHub repo as nothing will be saved to the pod when it is terminated.
After you are done with your work, you can terminate the pod by clicking the "Terminate" and "Delete" buttons. Remember if the pod is not completely deleted, Runpod will keep charging you for it.
## Extra Information:
If you need to customize the pod template this section has some useful information. For the most part you can leave the defaults of the FastVideo Pod Template.
When creating your pod template, use this image:
```
@@ -24,30 +63,3 @@ bash -c "apt update;DEBIAN_FRONTEND=noninteractive apt-get install openssh-serve
```
![RunPod template configuration](../../assets/images/runpod_template.png)
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
![RunPod ssh](../../assets/images/runpod_ssh.png)
## Working with the pod
After SSH'ing into your pod, you'll find the `fastvideo-dev` Conda environment already activated.
To pull in the latest changes from the GitHub repo:
```bash
cd /FastVideo
git pull
```
`If you have a persistent volume and want to keep your code changes, you can move /FastVideo to /workspace/FastVideo, or simply clone the repository there.`
Run your development workflows as usual:
```bash
# Run linters
pre-commit run --all-files
# Run tests
pytest tests/
```
+38 -30
View File
@@ -1,76 +1,84 @@
# 🛠️ Contributing to FastVideo
Thank you for your interest in contributing to FastVideo. We want to make the process as smooth for you as possible and this is a guide to help get you started!
Thank you for your interest in contributing to FastVideo. We want the process
to be smooth and beginner‑friendly, whether you are adding a new pipeline,
improving performance, or fixing a bug.
Our community is open to everyone and welcomes any contributions no matter how large or small.
## Quick prerequisites
# Developer Environment:
Do make sure you have CUDA 12.4 installed and supported. FastVideo currently only supports Linux and CUDA GPUs, but we hope to support other platforms in the future.
- **OS**: Linux is the primary development target (WSL can work).
- **GPU**: NVIDIA GPU recommended for inference and training workflows.
- **CUDA**: Use a recent CUDA 12.x toolchain (see the installation guide for
the current recommendation).
We recommend using a fresh Python 3.10 Conda environment to develop FastVideo:
For a full install checklist, see `docs/getting_started/installation/gpu.md`.
## Local development (Conda + editable install)
Install Miniconda:
```
```bash
wget https://repo.anaconda.com/miniconda/Miniconda3-latest-Linux-x86_64.sh
bash Miniconda3-latest-Linux-x86_64.sh
source ~/.bashrc
```
Create and activate a Conda environment for FastVideo:
Create and activate a Conda environment:
```
```bash
conda create -n fastvideo python=3.12 -y
conda activate fastvideo
```
Install `uv` (optional, but recommended):
From instructions on [uv](https://astral.sh/uv/):
```
```bash
curl -LsSf https://astral.sh/uv/install.sh | sh
# or
# or
wget -qO- https://astral.sh/uv/install.sh | sh
```
Clone the FastVideo repository and go to the FastVideo directory:
Clone the repo:
```
```bash
git clone https://github.com/hao-ai-lab/FastVideo.git && cd FastVideo
```
Now you can install FastVideo and setup git hooks for running linting. By using `pre-commit`, the linters will run and have to pass before you'll be able to make a commit.
Install FastVideo in editable mode and set up hooks:
```bash
uv pip install -e .[dev]
# Can also install flash-attn (optional)
uv pip install flash-attn --no-build-isolation
# Optional: FlashAttention (builds native kernels)
uv pip install flash-attn --no-build-isolation
# Linting, formatting and static type checking
# Linting, formatting, static typing
pre-commit install --hook-type pre-commit --hook-type commit-msg
# You can manually run pre-commit with
pre-commit run --all-files
# Unit tests
pytest tests/
```
If you are on a Hopper GPU, you should also install [FA3](https://github.com/Dao-AILab/flash-attention) for much better performance:
If you are on a Hopper GPU, installing FlashAttention 3 can improve
performance (see `docs/inference/optimizations.md`).
```
git clone https://github.com/Dao-AILab/flash-attention.git && cd flash-attention/hopper
## Docker development (optional)
# make sure you have ninja installed
uv pip install ninja
python setup.py install
```
If you prefer a containerized environment, use the dev image documented in
`docs/contributing/developer_env/docker.md`.
## Testing
Please refer to the [Testing Guide](testing.md) for more information on how to add and run tests in FastVideo.
See the [Testing Guide](testing.md) for how to add and run tests in FastVideo.
## Attention backend development
If you are adding a new attention kernel or backend, follow
[Attention Backend Development](attention_backend.md).
## Contributing with coding agents
For a step‑by‑step workflow on adding pipelines or components with coding
agents, see `docs/contributing/coding_agents.md`.
+141 -387
View File
@@ -1,424 +1,178 @@
# 🔍 FastVideo Overview
# FastVideo Architecture Overview
This document outlines FastVideo's architecture for developers interested in framework internals or contributions. It serves as an onboarding guide for new contributors by providing an overview of the most important directories and files within the `fastvideo/` codebase.
This document summarizes how FastVideo is structured and how a Diffusers-style
model repo maps into a runnable pipeline. It is intended for contributors who
need the high-level layout and key entrypoints, not every internal detail.
## Table of Contents - Directory Structure and Files
## FastVideo structure at a glance
- [`fastvideo/pipelines/`](#pipeline-system) - Core diffusion pipeline components
- [`fastvideo/models/`](#model-components) - Model implementations
- [`dits/`](#transformer-models) - Transformer-based diffusion models
- [`vaes/`](#vae-variational-auto-encoder) - Variational autoencoders
- [`encoders/`](#text-and-image-encoders) - Text and image encoders
- [`schedulers/`](#schedulers) - Diffusion schedulers
- [`fastvideo/attention/`](#optimized-attention) - Optimized attention implementations
- [`fastvideo/distributed/`](#distributed-processing) - Distributed computing utilities
- [`fastvideo/layers/`](#tensor-parallelism) - Custom neural network layers
- [`fastvideo/platforms/`](#platforms) - Hardware platform abstractions
- [`fastvideo/worker/`](#executor-and-worker-system) - Multi-GPU process management
- [`fastvideo/fastvideo_args.py`](#fastvideoargs) - Argument handling
- [`fastvideo/forward_context.py`](#forward-context-management) - Forward pass context management
- `fastvideo/utils.py` - Utility functions
- [`fastvideo/logger.py`](#logger) - Logging infrastructure
FastVideo maps a Diffusers-style repo into a pipeline like this:
## Core Architecture
- `fastvideo/models/*`: model implementations (DiT, VAE, encoders, upsamplers).
- `fastvideo/configs/models/*`: arch configs and `param_names_mapping` for
weight name translation.
- `fastvideo/configs/pipelines/*`: pipeline wiring (component classes + names).
- `fastvideo/configs/sample/*`: default runtime sampling parameters.
- `fastvideo/pipelines/basic/*`: end-to-end pipelines.
- `fastvideo/pipelines/stages/*`: reusable pipeline stages.
- `fastvideo/models/loader/*`: component loaders for Diffusers-style repos.
- `model_index.json`: HF repo entrypoint mapping component names to classes.
FastVideo separates model components from execution logic with these principles:
Flow:
`model_index.json` -> component loaders -> model modules -> pipeline stages ->
sampling params.
- **Component Isolation**: Models (encoders, VAEs, transformers) are isolated from execution (pipelines, stages, distributed processing)
- **Modular Design**: Components can be independently replaced
- **Distributed Execution**: Supports various parallelism strategies (Tensor, Sequence)
- **Custom Attention Backends**: Components can support and use different Attention implementations
- **Pipeline Abstraction**: Consistent interface across diffusion models
## FastVideoArgs
The `FastVideoArgs` class in `fastvideo/fastvideo_args.py` serves as the central configuration system for FastVideo. It contains all parameters needed to control model loading, inference configuration, performance optimization settings, and more.
Key features include:
- **Command-line Interface**: Automatic conversion between CLI arguments and dataclass fields
- **Configuration Groups**: Organized by functional areas (model loading, video params, optimization settings)
- **Context Management**: Global access to current settings via `get_current_fastvideo_args()`
- **Parameter Validation**: Ensures valid combinations of settings
Common configuration areas:
- **Model paths and loading options**: `model_path`, `trust_remote_code`, `revision`
- **Distributed execution settings**: `num_gpus`, `tp_size`, `sp_size`
- **Video generation parameters**: `height`, `width`, `num_frames`, `num_inference_steps`
- **Precision settings**: Control computation precision for different components
Example usage:
Minimal usage (from `examples/inference/basic/basic.py`):
```python
# Load arguments from command line
fastvideo_args = prepare_fastvideo_args(sys.argv[1:])
from fastvideo import VideoGenerator
from fastvideo.configs.sample import SamplingParam
# Access parameters
model = load_model(fastvideo_args.model_path)
model_id = "Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # or official_weights/<model_name>/
generator = VideoGenerator.from_pretrained(model_id, num_gpus=1)
# Set as global context
with set_current_fastvideo_args(fastvideo_args):
# Code that requires access to these arguments
result = generate_video()
```
## Pipeline System
### `ComposedPipelineBase`
This foundational class provides:
- **Model Loading**: Automatically loads components from HuggingFace-Diffusers-compatible model directories
- **Stage Management**: Creates and orchestrates processing stages
- **Data Flow Coordination**: Ensures proper state flow between stages
```python
class MyCustomPipeline(ComposedPipelineBase):
_required_config_modules = [
"text_encoder", "tokenizer", "vae", "transformer", "scheduler"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
# Pipeline-specific initialization
pass
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage("input_validation_stage", InputValidationStage())
self.add_stage("text_encoding_stage", CLIPTextEncodingStage(
text_encoder=self.get_module("text_encoder"),
tokenizer=self.get_module("tokenizer")
))
# Additional stages...
```
### Pipeline Stages
Each stage handles a specific diffusion process component:
- **Input Validation**: Parameter verification
- **Text Encoding**: CLIP, LLaMA, or T5-based encoding
- **Image Encoding**: Image input processing
- **Timestep & Latent Preparation**: Setup for diffusion
- **Denoising**: Core diffusion loop
- **Decoding**: Latent-to-pixel conversion
Each stage implements a standard interface:
```python
def forward(self, batch: ForwardBatch, fastvideo_args: FastVideoArgs) -> ForwardBatch:
# Process batch and update state
return batch
```
![Pipeline execution and data flow](../assets/images/pipeline.png)
### ForwardBatch
Defined in `fastvideo/pipelines/pipeline_batch_info.py`, `ForwardBatch` encapsulates the data payload passed between pipeline stages. It typically holds:
- **Input Data**: Prompts, images, generation parameters
- **Intermediate State**: Embeddings, latents, timesteps, accumulated during stage execution
- **Output Storage**: Generated results and metadata
- **Configuration**: Sampling parameters, precision settings
This structure facilitates clear state transitions between stages.
## Model Components
The `fastvideo/models/` directory contains implementations of the core neural network models used in video diffusion:
### Transformer Models
Transformer networks perform the actual denoising during diffusion:
- **Location**: `fastvideo/models/dits/`
- **Examples**:
- `WanTransformer3DModel`
- `HunyuanVideoTransformer3DModel`
Features include:
- Text/image conditioning
- Standardized interface for model-specific optimizations
```python
def forward(
self,
latents, # [B, T, C, H, W]
encoder_hidden_states, # Text embeddings
timestep, # Current diffusion timestep
encoder_hidden_states_image=None, # Optional image embeddings
**kwargs
):
# Perform denoising computation
return noise_pred # Predicted noise residual
```
### VAE (Variational Auto-Encoder)
VAEs handle conversion between pixel space and latent space:
- **Location**: `fastvideo/models/vaes/`
- **Examples**:
- `AutoencoderKLWan`
- `AutoencoderKLHunyuanVideo`
These models compress image/video data to a more efficient latent representation (typically 4x-8x smaller in each dimension).
FastVideo's VAE implementations include:
- Efficient video batch processing
- Memory optimization
- Optional tiling for large frames
- Distributed weight support
### Text and Image Encoders
Encoders process conditioning inputs into embeddings:
- **Location**: `fastvideo/models/encoders/`
- **Text Encoders**:
- `CLIPTextModel`
- `LlamaModel`
- `UMT5EncoderModel`
- **Image Encoders**:
- `CLIPVisionModel`
FastVideo implements optimizations such as:
- Vocab parallelism for distributed processing
- Caching for common prompts
- Precision-tuned computation
### Schedulers
Schedulers manage the diffusion sampling process:
- **Location**: `fastvideo/models/schedulers/`
- **Examples**:
- `UniPCMultistepScheduler`
- `FlowMatchEulerDiscreteScheduler`
These components control:
- Diffusion timestep sequences
- Noise prediction to latent update conversions
- Quality/speed trade-offs
```python
def step(
self,
model_output: torch.Tensor,
timestep: torch.LongTensor,
sample: torch.Tensor,
**kwargs
) -> torch.Tensor:
# Process model output and update latents
# Return updated latents
return prev_sample
```
This diagram shows how models are discovered, validated, and loaded across entrypoints, executors, pipelines, and model loaders.
![Model loading flow](../assets/images/load_models.png)
## Optimized Attention
The `fastvideo/attention/` directory contains optimized attention implementations crucial for efficient video diffusion:
### Attention Backends
Multiple implementations with automatic selection:
- **FLASH_ATTN**: Optimized for supporting hardware
- **TORCH_SDPA**: Built-in PyTorch scaled dot-product attention
- **SLIDING_TILE_ATTN**: For very long sequences
```python
# Configure available attention backends for this layer
self.attn = LocalAttention(
num_heads=num_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=(_Backend.FLASH_ATTN, _Backend.TORCH_SDPA)
)
# Override via environment variable
# export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
```
![Attention backend selector design](../assets/images/attention_backend.png)
### Attention Patterns
Supports various patterns with memory optimization techniques:
- **Cross/Self/Temporal/Global-Local Attention**
- Chunking, progressive computation, optimized masking
## Distributed Processing
The `fastvideo/distributed/` directory contains implementations for distributed model execution:
### Tensor Parallelism
Tensor parallelism splits model weights across devices:
- **Implementation**: Through `RowParallelLinear` and `ColumnParallelLinear` layers
- **Use cases**: Will be used by encoder models as their sequence lengths are shorter and enables efficient sharding.
```python
# Tensor-parallel layers in a transformer block
from fastvideo.layers.linear import ColumnParallelLinear, RowParallelLinear
# Split along output dimension
self.qkv_proj = ColumnParallelLinear(
input_size=hidden_size,
output_size=3 * hidden_size,
bias=True,
gather_output=False
)
# Split along input dimension
self.out_proj = RowParallelLinear(
input_size=hidden_size,
output_size=hidden_size,
bias=True,
input_is_parallel=True
sampling = SamplingParam.from_pretrained(model_id)
sampling.num_frames = 45
video = generator.generate_video(
"A vibrant city street at sunset.",
sampling_param=sampling,
output_path="video_samples",
save_video=True,
)
```
### Sequence Parallelism
## Configuration system
Sequence parallelism splits sequences across devices:
FastVideo uses typed configs to keep model definitions, pipeline wiring, and
runtime parameters consistent:
- **Implementation**: Through `DistributedAttention` and sequence splitting
- **Use cases**: Long video sequences or high-resolution processing. Used by DiT models.
- `fastvideo/configs/models/`: architecture definitions, layer shapes, and
`param_names_mapping` rules for key renaming.
- `fastvideo/configs/pipelines/`: pipeline wiring and required components.
- `fastvideo/configs/sample/`: default sampling parameters (steps, frames,
guidance scale, resolution, fps).
- `fastvideo/configs/registry.py`: pipeline registry.
- `fastvideo/configs/sample/registry.py`: sampling param registry.
```python
# Distributed attention for long sequences
from fastvideo.attention import DistributedAttention
`FastVideoArgs` (in `fastvideo/fastvideo_args.py`) provides runtime settings and
is passed into pipeline construction and stages.
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=head_dim,
causal=False,
supported_attention_backends=(_Backend.SLIDING_TILE_ATTN, _Backend.FLASH_ATTN)
)
## Weights and Diffusers format
FastVideo follows the HuggingFace Diffusers repo layout. This keeps loaders
compatible with HF repos and makes it easy to add new components.
Typical Diffusers repo:
```
<model-repo>/
model_index.json
scheduler/
scheduler_config.json
transformer/ # or unet/ for image models
config.json
diffusion_pytorch_model.safetensors
vae/
config.json
diffusion_pytorch_model.safetensors
text_encoder/
config.json
model.safetensors
tokenizer/
tokenizer_config.json
tokenizer.json
```
### Communication Primitives
Key points:
Efficient distributed operations via AllGather, AllReduce, and synchronization mechanisms.
- `model_index.json` is the root map that tells FastVideo which components to
load and which classes implement them.
- Each component lives in its own folder with a `config.json` and weights.
- Weights are usually in `diffusion_pytorch_model.safetensors`.
Efficient communication primitives minimize distributed overhead:
Note on tensor names:
- **Sequence-Parallel AllGather**: Collects sequence chunks
- **Tensor-Parallel AllReduce**: Combines partial results
- **Distributed Synchronization**: Coordinates execution
Official checkpoints often use different `state_dict` names than FastVideo's
module layout. We translate tensor names via the DiT arch config mapping
(`param_names_mapping` under `fastvideo/configs/models/dits/`). This is similar
in spirit to name-translation layers used in systems like vLLM and SGLang.
## Forward Context Management
Example HF repo (Wan 2.1 T2V 1.3B Diffusers):
### ForwardContext
Defined in `fastvideo/forward_context.py`, `ForwardContext` manages execution-specific state *within* a forward pass, particularly for low-level optimizations. It is accessed via `get_forward_context()`.
- **Attention Metadata**: Configuration for optimized attention kernels (`attn_metadata`)
- **Profiling Data**: Potential hooks for performance metrics collection
This context-based approach enables:
- Dynamic optimization based on execution state (e.g., attention backend selection)
- Step-specific customizations within model components
Usage example:
```python
with set_forward_context(current_timestep, attn_metadata, fastvideo_args):
# During this forward pass, components can access context
# through get_forward_context()
output = model(inputs)
```
https://huggingface.co/Wan-AI/Wan2.1-T2V-1.3B-Diffusers/tree/main
```
## Executor and Worker System
Example `model_index.json` from that repo:
The `fastvideo/worker/` directory contains the distributed execution framework:
### Executor Abstraction
FastVideo implements a flexible execution model for distributed processing:
- **Executor Base Class**: An abstract base class defining the interface for all executors
- **MultiProcExecutor**: Primary implementation that spawns and manages worker processes
- **GPU Workers**: Handle actual model execution on individual GPUs
The MultiProcExecutor implementation:
1. Spawns worker processes for each GPU
2. Establishes communication channels via pipes
3. Coordinates distributed operations across workers
4. Handles graceful startup and shutdown of the process group
Each GPU worker:
1. Initializes the distributed environment
2. Builds the pipeline for the specified model
3. Executes requested operations on its assigned GPU
4. Manages local resources and communicates results back to the executor
This design allows FastVideo to efficiently utilize multiple GPUs while providing a simple, unified interface for model execution.
## Platforms
The `fastvideo/platforms/` directory provides hardware platform abstractions that enable FastVideo to run efficiently on different hardware configurations:
### Platform Abstraction
FastVideo's platform abstraction layer enables:
- **Hardware Detection**: Automatic detection of available hardware
- **Backend Selection**: Appropriate selection of compute kernels
- **Memory Management**: Efficient utilization of hardware-specific memory features
The primary components include:
- **Platform Interface**: Defines the common API for all platform implementations
- **CUDA Platform**: Optimized implementation for NVIDIA GPUs
- **Backend Enum**: Used throughout the codebase for feature selection
Usage example:
```python
from fastvideo.platforms import current_platform, _Backend
# Check hardware capabilities
if current_platform.supports_backend(_Backend.FLASH_ATTN):
# Use FlashAttention implementation
else:
# Fall back to standard implementation
```json
{
"_class_name": "WanPipeline",
"_diffusers_version": "0.33.0.dev0",
"scheduler": [
"diffusers",
"UniPCMultistepScheduler"
],
"text_encoder": [
"transformers",
"UMT5EncoderModel"
],
"tokenizer": [
"transformers",
"T5TokenizerFast"
],
"transformer": [
"diffusers",
"WanTransformer3DModel"
],
"vae": [
"diffusers",
"AutoencoderKLWan"
]
}
```
The platform system is designed to be extensible for future hardware targets.
How this maps to FastVideo:
## Logger
- `WanPipeline` -> `fastvideo/pipelines/basic/wan/wan_pipeline.py`
- `WanTransformer3DModel` -> `fastvideo/models/dits/wanvideo.py`
- `AutoencoderKLWan` -> `fastvideo/models/vaes/wanvae.py`
- `UMT5EncoderModel` -> `fastvideo/models/encoders/t5.py`
- `T5TokenizerFast` -> loaded via HF in `fastvideo/models/loader/`
- `UniPCMultistepScheduler` -> loaded via Diffusers scheduler utilities
- Pipeline defaults -> `fastvideo/configs/pipelines/wan.py`
- Sampling defaults -> `fastvideo/configs/sample/wan.py`
See [PR](https://github.com/hao-ai-lab/FastVideo/pull/356)
## Pipeline system
*TODO*: (help wanted) Add an environment variable that disables process-aware logging.
- `fastvideo/pipelines/basic/*` contains end-to-end pipelines for each model
family.
- `fastvideo/pipelines/stages/*` contains reusable, testable stages.
- Pipelines subclass `ComposedPipelineBase` and declare required components via
`_required_config_modules`.
- `ForwardBatch` (in `fastvideo/pipelines/pipeline_batch_info.py`) carries
prompts, latents, timesteps, and intermediate state across stages.
## Contributing to FastVideo
## Model components
If you're a new contributor, here are some common areas to explore:
- DiT models: `fastvideo/models/dits/`
- VAEs: `fastvideo/models/vaes/`
- Text/image encoders: `fastvideo/models/encoders/`
- Schedulers: `fastvideo/models/schedulers/`
- Upsamplers: `fastvideo/models/upsamplers/`
- Optional audio models: `fastvideo/models/audio/`
1. **Adding a new model**: Implement new model types in the appropriate subdirectory of `fastvideo/models/`
2. **Optimizing performance**: Look at attention implementations or memory management
3. **Adding a new pipeline**: Create a new pipeline subclass in `fastvideo/pipelines/`
4. **Hardware support**: Extend the `platforms` module for new hardware targets
## Attention and distributed execution
When adding code, follow these practices:
- Attention backends live in `fastvideo/attention/` and can be selected via
`FASTVIDEO_ATTENTION_BACKEND`.
- `LocalAttention` is used for cross-attention and most attention layers.
- `DistributedAttention` is used for full-sequence self-attention in the DiT.
- Tensor-parallel layers live in `fastvideo/layers/`.
- Sequence/tensor parallel utilities live in `fastvideo/distributed/`.
- Use type hints for better code readability
- Add appropriate docstrings
- Maintain the separation between model components and execution logic
- Follow existing patterns for distributed processing
## Related docs
- [Contributing overview](../contributing/overview.md)
- [Coding agents workflow](../contributing/coding_agents.md)
- [Testing guide](../contributing/testing.md)
+140
View File
@@ -0,0 +1,140 @@
# Offloading
This page describes how to use offloading techniques for inference to reduce GPU memory usage while maintaining acceptable performance.
## Default Behavior
```python
dit_cpu_offload: bool = True
use_fsdp_inference: bool = False
dit_layerwise_offload: bool = True
text_encoder_cpu_offload: bool = True
image_encoder_cpu_offload: bool = True
vae_cpu_offload: bool = True
pin_cpu_memory: bool = True
```
## Behavior Explanation
!!! note
For CLI usage, replace underscores (`_`) with hyphens (`-`).
### `use_fsdp_inference`
Enables [FSDP](https://docs.pytorch.org/tutorials/intermediate/FSDP_tutorial.html) for inference. The model weights are sharded across multiple GPUs to reduce memory usage per GPU, and weights are broadcast to all GPUs layer by layer during inference.
#### Performance Impact
FSDP inference introduces negligible performance overhead due to weight prefetching. Performance overhead may be visible when GPU interconnect is slow (e.g., multiple consumer-level GPUs connected by slow PCIe without GPU P2P support).
#### Usage Recommendation
We recommend enabling this option when multiple GPUs are available.
### `dit_cpu_offload`
Enables CPU offloading for FSDP inference. When enabled, the model weights are offloaded to CPU memory, and the weight of each layer is moved to GPU memory only when that layer is being computed.
#### Performance Impact
The PyTorch FSDP implementation does not overlap computation and data transfer perfectly for inference, so enabling this option will harm performance.
#### Usage Recommendation
This option only takes effect when FSDP is enabled. For single GPU usage, we recommend using `dit_layerwise_offload` instead.
### `dit_layerwise_offload`
This option is similar to `dit_cpu_offload`, but with two key differences:
1. It overlaps computation and PCIe data transfer
2. It only works for single GPU inference
#### Performance Impact
This option introduces negligible performance overhead.
#### Usage Recommendation
We recommend enabling this option for single GPU usage. This option is not compatible with FSDP.
### `text_encoder_cpu_offload`
When enabled, the text encoder model weights are offloaded to CPU memory, and text encoding is computed on CPU.
#### Performance Impact
This option significantly slows down text encoding computation, but text encoding is usually not the bottleneck.
#### Usage Recommendation
We recommend enabling this option only when OOM happens.
### `image_encoder_cpu_offload` and `vae_cpu_offload`
When enabled, the weights are stored in CPU memory and moved to GPU memory when the corresponding module is being computed. After computation, the weights are moved back to CPU memory.
#### Performance Impact
These options introduce performance overhead due to PCIe data transfer.
#### Usage Recommendation
We recommend enabling these options when OOM happens.
## General Recommendations
### Single GPU Inference
We recommend enabling `dit_layerwise_offload`. If OOM happens, also enable `image_encoder_cpu_offload` and `vae_cpu_offload`. If OOM still happens, consider enabling `text_encoder_cpu_offload`.
### Multi-GPU Inference
We recommend enabling `use_fsdp_inference` and disabling both `dit_layerwise_offload` and `dit_cpu_offload`. If OOM happens, consider enabling `text_encoder_cpu_offload`, `image_encoder_cpu_offload`, and `vae_cpu_offload`. If OOM still happens, consider enabling `dit_cpu_offload`.
## Examples
### Single GPU with Layerwise Offloading
```python
from fastvideo import VideoGenerator
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=1,
# Recommended for single GPU
dit_layerwise_offload=True,
# Enable if OOM happens
vae_cpu_offload=True,
image_encoder_cpu_offload=True,
text_encoder_cpu_offload=True,
# Speeds up CPU-GPU transfer
pin_cpu_memory=True,
)
prompt = "A curious raccoon peers through a vibrant field of yellow sunflowers."
video = generator.generate_video(prompt, output_path="output/", save_video=True)
```
### Multi-GPU with FSDP
```python
from fastvideo import VideoGenerator
generator = VideoGenerator.from_pretrained(
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers",
num_gpus=2,
# Recommended for multi-GPU
use_fsdp_inference=True,
dit_layerwise_offload=False,
dit_cpu_offload=False,
# Enable if OOM happens
vae_cpu_offload=True,
image_encoder_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
)
prompt = "A majestic lion strides across the golden savanna."
video = generator.generate_video(prompt, output_path="output/", save_video=True)
```
+2 -6
View File
@@ -103,7 +103,7 @@ python setup.py install # or pip install -e .
**`SAGE_ATTN_THREE`**
[SageAttention 3](https://huggingface.co/jt-zhang/SageAttention3) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
[SageAttention 3](https://github.com/thu-ml/SageAttention/tree/main/sageattention3_blackwell) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
#### Hardware Requirements
@@ -113,11 +113,7 @@ python setup.py install # or pip install -e .
Note that Sage Attention 3 requires `python>=3.13`, `torch>=2.8.0`, `CUDA >=12.8`. If you are using `uv` and using `torch==2.8.0` make sure that `sentencepiece==0.2.1` in the pyproject.toml file.
To use Sage Attention 3 in FastVideo, first get access to the SageAttention3 code, then move `sageattn/` and `setup.py` to the directory `fastvideo/attention/backends`, then install from using:
```bash
python setup.py install
```
To use Sage Attention 3 in FastVideo, follow the `README.md` in the linked repository to install the package from source.
## Teacache
+1 -2
View File
@@ -1,6 +1,5 @@
from fastvideo import VideoGenerator
from fastvideo.configs.pipelines.wan import MatrixGameI2V480PConfig
from fastvideo.models.dits.matrix_game.utils import create_action_presets
from fastvideo.models.dits.matrixgame.utils import create_action_presets
import torch
@@ -1,5 +1,5 @@
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
from fastvideo.models.dits.matrix_game.utils import get_current_action_async, expand_action_to_frames
from fastvideo.models.dits.matrixgame.utils import get_current_action_async, expand_action_to_frames
import torch
import asyncio
@@ -10,7 +10,7 @@ from fastapi import FastAPI, Request, HTTPException
from fastapi.responses import HTMLResponse, FileResponse
from fastvideo.entrypoints.streaming_generator import StreamingVideoGenerator
from fastvideo.models.dits.matrix_game.utils import expand_action_to_frames
from fastvideo.models.dits.matrixgame.utils import expand_action_to_frames
VARIANT_CONFIG = {
@@ -0,0 +1,93 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
MODEL_PATH="FastVideo/Matrix-Game-2.0-Foundation-Diffusers"
DATA_DIR="footsies-dataset/preprocessed/combined_parquet_dataset"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=8
# export CUDA_VISIBLE_DEVICES=4,5
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "matrixgame_finetune"
--output_dir "checkpoints/matrixgame_finetune"
--max_train_steps 1500
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 4
--num_latent_t 20
--num_height 352
--num_width 640
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 2
--tp_size 1
--hsdp_replicate_dim 4
--hsdp_shard_dim 2
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 100
--training_state_checkpointing_steps 100
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/matrixgame_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,27 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="Matrix-Game-2.0-Foundation-Diffusers"
DATA_MERGE_PATH="footsies-dataset/merge.txt"
OUTPUT_DIR="footsies-dataset/preprocessed/"
# export CUDA_VISIBLE_DEVICES=0
export MASTER_ADDR=localhost
export MASTER_PORT=29500
export RANK=0
export WORLD_SIZE=1
python fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 4 \
--seed 42 \
--max_height 352 \
--max_width 640 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--samples_per_file 4 \
--train_fps 25 \
--flush_frequency 4 \
--preprocess_task matrixgame
File diff suppressed because it is too large Load Diff
@@ -0,0 +1,18 @@
Total Files: 16
00. Hold [W] + Static
01. Hold [S] + Static
02. Hold [A] + Static
03. Hold [D] + Static
04. Hold [WA] + Static
05. Hold [WD] + Static
06. Hold [SA] + Static
07. Hold [SD] + Static
08. No Key + Hold [up]
09. No Key + Hold [down]
10. No Key + Hold [left]
11. No Key + Hold [right]
12. No Key + Hold [up_right]
13. No Key + Hold [up_left]
14. No Key + Hold [down_right]
15. No Key + Hold [down_left]
@@ -0,0 +1,93 @@
#!/bin/bash
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=offline
export TOKENIZERS_PARALLELISM=false
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
MODEL_PATH="weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers"
DATA_DIR="mc_wasd_10/preprocessed/combined_parquet_dataset"
VALIDATION_DATASET_FILE="mc_wasd_10/validation.json"
NUM_GPUS=4
# export CUDA_VISIBLE_DEVICES=0,1,2,3
# IP=[MASTER NODE IP]
# Training arguments
training_args=(
--tracker_project_name "wangame_1.3b_overfit"
--output_dir "wangame_1.3b_overfit"
--max_train_steps 1500
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 20
--num_height 352
--num_width 640
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wangame_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,120 @@
#!/bin/bash
#SBATCH --job-name=wangame_1.3b_overfit
#SBATCH --partition=main
#SBATCH --nodes=1
#SBATCH --ntasks=1
#SBATCH --ntasks-per-node=1
#SBATCH --gres=gpu:8
#SBATCH --cpus-per-task=128
#SBATCH --mem=1440G
#SBATCH --output=wangame_1.3b_overfit_output/wangame_1.3b_overfit_%j.out
#SBATCH --error=wangame_1.3b_overfit_output/wangame_1.3b_overfit_%j.err
#SBATCH --exclusive
# Basic Info
export NCCL_P2P_DISABLE=1
export TORCH_NCCL_ENABLE_MONITORING=0
export NCCL_DEBUG_SUBSYS=INIT,NET
# different cache dir for different processes
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
export MASTER_PORT=29500
export NODE_RANK=$SLURM_PROCID
nodes=( $(scontrol show hostnames $SLURM_JOB_NODELIST) )
export MASTER_ADDR=${nodes[0]}
export TOKENIZERS_PARALLELISM=false
# export WANDB_API_KEY="8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde"
export WANDB_API_KEY="50632ebd88ffd970521cec9ab4a1a2d7e85bfc45"
# export WANDB_API_KEY='your_wandb_api_key_here'
export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export FASTVIDEO_ATTENTION_BACKEND=FLASH_ATTN
source ~/conda/miniconda/bin/activate
conda activate wei-fv-distill
export HOME="/mnt/weka/home/hao.zhang/wei"
MODEL_PATH="weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers"
DATA_DIR="mc_wasd_10/preprocessed/combined_parquet_dataset"
VALIDATION_DATASET_FILE="examples/training/finetune/WanGame2.1_1.3b_i2v/validation.json"
# Configs
NUM_GPUS=8
# Training arguments
training_args=(
--tracker_project_name "wangame_1.3b_overfit"
--output_dir "wangame_1.3b_overfit"
--max_train_steps 15000
--train_batch_size 1
--train_sp_batch_size 1
--gradient_accumulation_steps 1
--num_latent_t 20
--num_height 352
--num_width 640
--num_frames 77
--enable_gradient_checkpointing_type "full"
)
# Parallel arguments
parallel_args=(
--num_gpus $NUM_GPUS
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
)
# Model arguments
model_args=(
--model_path $MODEL_PATH
--pretrained_model_name_or_path $MODEL_PATH
)
# Dataset arguments
dataset_args=(
--data_path "$DATA_DIR"
--dataloader_num_workers 1
)
# Validation arguments
validation_args=(
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 100
--validation_sampling_steps "40"
--validation_guidance_scale "1.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 2e-5
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000000
--training_state_checkpointing_steps 10000000
--weight_decay 1e-4
--max_grad_norm 1.0
)
# Miscellaneous arguments
miscellaneous_args=(
--inference_mode False
--checkpoints_total_limit 3
--training_cfg_rate 0.1
--multi_phased_distill_schedule "4000-1"
--not_apply_cfg_solver
--dit_precision "fp32"
--num_euler_timesteps 50
--ema_start_step 0
)
# If you do not have 32 GPUs and to fit in memory, you can: 1. increase sp_size. 2. reduce num_latent_t
torchrun \
--nnodes 1 \
--nproc_per_node $NUM_GPUS \
fastvideo/training/wangame_training_pipeline.py \
"${parallel_args[@]}" \
"${model_args[@]}" \
"${dataset_args[@]}" \
"${training_args[@]}" \
"${optimizer_args[@]}" \
"${validation_args[@]}" \
"${miscellaneous_args[@]}"
@@ -0,0 +1,193 @@
import os
import numpy as np
# Configuration
SCRIPT_DIR = os.path.dirname(os.path.abspath(__file__))
BASE_OUTPUT_DIR = os.path.join(SCRIPT_DIR, 'action')
VIDEO_OUTPUT_DIR = BASE_OUTPUT_DIR
os.makedirs(VIDEO_OUTPUT_DIR, exist_ok=True)
FRAME_COUNT = 81
CAM_VALUE = 0.1
# Action Mapping
KEY_TO_INDEX = {
'W': 0, 'S': 1, 'A': 2, 'D': 3,
}
VIEW_ACTION_TO_MOUSE = {
"stop": [0.0, 0.0],
"up": [CAM_VALUE, 0.0],
"down": [-CAM_VALUE, 0.0],
"left": [0.0, -CAM_VALUE],
"right": [0.0, CAM_VALUE],
"up_right": [CAM_VALUE, CAM_VALUE],
"up_left": [CAM_VALUE, -CAM_VALUE],
"down_right": [-CAM_VALUE, CAM_VALUE],
"down_left": [-CAM_VALUE, -CAM_VALUE],
}
def get_multihot_vector(keys_str):
"""Convert string like 'WA' to [1, 0, 1, 0, 0, 0]"""
vector = [0.0] * 6
if not keys_str:
return vector
for char in keys_str.upper():
if char in KEY_TO_INDEX:
vector[KEY_TO_INDEX[char]] = 1.0
return vector
def get_mouse_vector(view_str):
"""Convert view string to [x, y]"""
return VIEW_ACTION_TO_MOUSE.get(view_str.lower(), [0.0, 0.0])
def generate_sequence(key_seq, mouse_seq):
"""
Generates action arrays based on sequences.
"""
keyboard_arr = np.zeros((FRAME_COUNT, 6), dtype=np.float32)
mouse_arr = np.zeros((FRAME_COUNT, 2), dtype=np.float32)
mid_point = FRAME_COUNT // 2
# First Half
k_vec1 = get_multihot_vector(key_seq[0])
m_vec1 = get_mouse_vector(mouse_seq[0])
keyboard_arr[:mid_point] = k_vec1
mouse_arr[:mid_point] = m_vec1
# Second Half
k_vec2 = get_multihot_vector(key_seq[1])
m_vec2 = get_mouse_vector(mouse_seq[1])
keyboard_arr[mid_point:] = k_vec2
mouse_arr[mid_point:] = m_vec2
return keyboard_arr, mouse_arr
def save_action(index, keyboard_arr, mouse_arr):
filename = f"{index:06d}_action.npy"
filepath = os.path.join(VIDEO_OUTPUT_DIR, filename)
action_dict = {
'keyboard': keyboard_arr,
'mouse': mouse_arr
}
np.save(filepath, action_dict)
return filename
def generate_description(key_seq, mouse_seq):
"""Generates a human-readable string for the combination."""
k1, k2 = key_seq
m1, m2 = mouse_seq
# Format Keyboard Description
if not k1 and not k2:
k_desc = "No Key"
elif k1 == k2:
k_desc = f"Hold [{k1}]"
else:
k_desc = f"Switch [{k1}]->[{k2}]"
# Format Mouse Description
if m1 == "stop" and m2 == "stop":
m_desc = "Static"
elif m1 == m2:
m_desc = f"Hold [{m1}]"
else:
m_desc = f"Switch [{m1}]->[{m2}]"
return f"{k_desc} + {m_desc}"
# ==========================================
# Main Generation Logic
# ==========================================
configs = []
readme_content = []
# Group 1: Constant Keyboard, No Mouse (0-7)
keys_basic = ['W', 'S', 'A', 'D', 'WA', 'WD', 'SA', 'SD']
for k in keys_basic:
configs.append(((k, k), ("stop", "stop")))
# Group 2: No Keyboard, Constant Mouse (8-15)
mouse_basic = ['up', 'down', 'left', 'right', 'up_right', 'up_left', 'down_right', 'down_left']
for m in mouse_basic:
configs.append((("", ""), (m, m)))
# Group 3: Split Keyboard, No Mouse (16-23)
split_keys = [
('W', 'S'), ('S', 'W'),
('A', 'D'), ('D', 'A'),
('W', 'A'), ('W', 'D'),
('S', 'A'), ('S', 'D')
]
for k1, k2 in split_keys:
configs.append(((k1, k2), ("stop", "stop")))
# Group 4: No Keyboard, Split Mouse (24-31)
split_mouse = [
('left', 'right'), ('right', 'left'),
('up', 'down'), ('down', 'up'),
('up_left', 'up_right'), ('up_right', 'up_left'),
('left', 'up'), ('right', 'down')
]
for m1, m2 in split_mouse:
configs.append((("", ""), (m1, m2)))
# Group 5: Constant Keyboard + Constant Mouse (32-47)
combo_keys = ['W', 'S', 'W', 'S', 'A', 'D', 'WA', 'WD', 'W', 'S', 'W', 'S', 'A', 'D', 'WA', 'WD']
combo_mice = ['left', 'left', 'right', 'right', 'up', 'up', 'down', 'down', 'up_left', 'up_left', 'up_right', 'up_right', 'down_left', 'down_right', 'right', 'left']
for i in range(16):
configs.append(((combo_keys[i], combo_keys[i]), (combo_mice[i], combo_mice[i])))
# Group 6: Constant Keyboard, Split Mouse (48-55)
complex_1_keys = ['W'] * 8
complex_1_mice = [
('left', 'right'), ('right', 'left'),
('up', 'down'), ('down', 'up'),
('left', 'up'), ('right', 'up'),
('left', 'down'), ('right', 'down')
]
for i in range(8):
configs.append(((complex_1_keys[i], complex_1_keys[i]), complex_1_mice[i]))
# Group 7: Split Keyboard, Constant Mouse (56-63)
complex_2_keys = [
('W', 'S'), ('S', 'W'),
('A', 'D'), ('D', 'A'),
('W', 'A'), ('W', 'D'),
('S', 'A'), ('S', 'D')
]
complex_2_mouse = 'up'
for k1, k2 in complex_2_keys:
configs.append(((k1, k2), (complex_2_mouse, complex_2_mouse)))
# Execution
print(f"Preparing to generate {len(configs)} action files...")
for i, (key_seq, mouse_seq) in enumerate(configs):
if i >= 16: break
# Generate Data
kb_arr, ms_arr = generate_sequence(key_seq, mouse_seq)
filename = save_action(i, kb_arr, ms_arr)
# Generate Description for README
description = generate_description(key_seq, mouse_seq)
readme_entry = f"{i:02d}. {description}"
readme_content.append(readme_entry)
print(f"Generated {filename} -> {description}")
# Write README
readme_path = os.path.join(BASE_OUTPUT_DIR, 'README.md')
with open(readme_path, 'w', encoding='utf-8') as f:
f.write(f"Total Files: {len(readme_content)}\n\n")
for line in readme_content:
f.write(line + '\n')
print(f"\nProcessing complete.")
print(f"64 .npy files generated in {VIDEO_OUTPUT_DIR}")
print(f"Manifest saved to {readme_path}")
@@ -0,0 +1,27 @@
#!/bin/bash
GPU_NUM=1 # 2,4,8
MODEL_PATH="weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers"
DATA_MERGE_PATH="mc_wasd_10/merge.txt"
OUTPUT_DIR="mc_wasd_10/preprocessed/"
# export CUDA_VISIBLE_DEVICES=0
export MASTER_ADDR=localhost
export MASTER_PORT=29500
export RANK=0
export WORLD_SIZE=1
python fastvideo/pipelines/preprocess/v1_preprocess.py \
--model_path $MODEL_PATH \
--data_merge_path $DATA_MERGE_PATH \
--preprocess_video_batch_size 10 \
--seed 42 \
--max_height 352 \
--max_width 640 \
--num_frames 77 \
--dataloader_num_workers 0 \
--output_dir=$OUTPUT_DIR \
--samples_per_file 10 \
--train_fps 25 \
--flush_frequency 10 \
--preprocess_task wangame
@@ -0,0 +1,404 @@
{
"data": [
{
"caption": "0",
"image_path": "../../../../mc_wasd_10/validate/000000.jpg",
"action_path": "../../../../mc_wasd_10/videos/000000_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "1",
"image_path": "../../../../mc_wasd_10/validate/000001.jpg",
"action_path": "../../../../mc_wasd_10/videos/000001_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "2",
"image_path": "../../../../mc_wasd_10/validate/000002.jpg",
"action_path": "../../../../mc_wasd_10/videos/000002_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "3",
"image_path": "../../../../mc_wasd_10/validate/000003.jpg",
"action_path": "../../../../mc_wasd_10/videos/000003_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "4",
"image_path": "../../../../mc_wasd_10/validate/000004.jpg",
"action_path": "../../../../mc_wasd_10/videos/000004_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "5",
"image_path": "../../../../mc_wasd_10/validate/000005.jpg",
"action_path": "../../../../mc_wasd_10/videos/000005_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "6",
"image_path": "../../../../mc_wasd_10/validate/000006.jpg",
"action_path": "../../../../mc_wasd_10/videos/000006_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "7",
"image_path": "../../../../mc_wasd_10/validate/000007.jpg",
"action_path": "../../../../mc_wasd_10/videos/000007_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "00. Hold [W] + Static",
"image_path": "../../../../mc_wasd_10/validate/000000.jpg",
"action_path": "action/000000_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "01. Hold [S] + Static",
"image_path": "../../../../mc_wasd_10/validate/000001.jpg",
"action_path": "action/000001_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "02. Hold [A] + Static",
"image_path": "../../../../mc_wasd_10/validate/000002.jpg",
"action_path": "action/000002_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "03. Hold [D] + Static",
"image_path": "../../../../mc_wasd_10/validate/000003.jpg",
"action_path": "action/000003_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "04. Hold [WA] + Static",
"image_path": "../../../../mc_wasd_10/validate/000004.jpg",
"action_path": "action/000004_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "05. Hold [WD] + Static",
"image_path": "../../../../mc_wasd_10/validate/000005.jpg",
"action_path": "action/000005_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "06. Hold [SA] + Static",
"image_path": "../../../../mc_wasd_10/validate/000006.jpg",
"action_path": "action/000006_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "07. Hold [SD] + Static",
"image_path": "../../../../mc_wasd_10/validate/000007.jpg",
"action_path": "action/000007_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "08. No Key + Hold [up]",
"image_path": "../../../../mc_wasd_10/validate/000000.jpg",
"action_path": "action/000008_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "09. No Key + Hold [down]",
"image_path": "../../../../mc_wasd_10/validate/000001.jpg",
"action_path": "action/000009_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "10. No Key + Hold [left]",
"image_path": "../../../../mc_wasd_10/validate/000002.jpg",
"action_path": "action/000010_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "11. No Key + Hold [right]",
"image_path": "../../../../mc_wasd_10/validate/000003.jpg",
"action_path": "action/000011_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "12. No Key + Hold [up_right]",
"image_path": "../../../../mc_wasd_10/validate/000004.jpg",
"action_path": "action/000012_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "13. No Key + Hold [up_left]",
"image_path": "../../../../mc_wasd_10/validate/000005.jpg",
"action_path": "action/000013_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "14. No Key + Hold [down_right]",
"image_path": "../../../../mc_wasd_10/validate/000006.jpg",
"action_path": "action/000014_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "15. No Key + Hold [down_left]",
"image_path": "../../../../mc_wasd_10/validate/000007.jpg",
"action_path": "action/000015_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "00. Hold [W] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000000.jpg",
"action_path": "action/000000_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "01. Hold [S] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000001.jpg",
"action_path": "action/000001_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "02. Hold [A] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000002.jpg",
"action_path": "action/000002_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "03. Hold [D] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000003.jpg",
"action_path": "action/000003_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "04. Hold [WA] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000004.jpg",
"action_path": "action/000004_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "05. Hold [WD] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000005.jpg",
"action_path": "action/000005_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "06. Hold [SA] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000006.jpg",
"action_path": "action/000006_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "07. Hold [SD] + Static",
"image_path": "../../../../mc_wasd_10/validate/gen_000007.jpg",
"action_path": "action/000007_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "08. No Key + Hold [up]",
"image_path": "../../../../mc_wasd_10/validate/gen_000000.jpg",
"action_path": "action/000008_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "09. No Key + Hold [down]",
"image_path": "../../../../mc_wasd_10/validate/gen_000001.jpg",
"action_path": "action/000009_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "10. No Key + Hold [left]",
"image_path": "../../../../mc_wasd_10/validate/gen_000002.jpg",
"action_path": "action/000010_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "11. No Key + Hold [right]",
"image_path": "../../../../mc_wasd_10/validate/gen_000003.jpg",
"action_path": "action/000011_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "12. No Key + Hold [up_right]",
"image_path": "../../../../mc_wasd_10/validate/gen_000004.jpg",
"action_path": "action/000012_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "13. No Key + Hold [up_left]",
"image_path": "../../../../mc_wasd_10/validate/gen_000005.jpg",
"action_path": "action/000013_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "14. No Key + Hold [down_right]",
"image_path": "../../../../mc_wasd_10/validate/gen_000006.jpg",
"action_path": "action/000014_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
},
{
"caption": "15. No Key + Hold [down_left]",
"image_path": "../../../../mc_wasd_10/validate/gen_000007.jpg",
"action_path": "action/000015_action.npy",
"video_path": null,
"num_inference_steps": 40,
"height": 352,
"width": 640,
"num_frames": 77
}
]
}
+2 -2
View File
@@ -1,7 +1,7 @@
# SPDX-License-Identifier: Apache-2.0
import torch
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell
from sageattn3 import sageattn3_blackwell
from fastvideo.attention.backends.abstract import (AttentionBackend,
AttentionImpl,
@@ -67,6 +67,6 @@ class SageAttention3Impl(AttentionImpl):
query = query.transpose(1, 2)
key = key.transpose(1, 2)
value = value.transpose(1, 2)
output = sageattn_blackwell(query, key, value, is_causal=self.causal)
output = sageattn3_blackwell(query, key, value, is_causal=self.causal)
output = output.transpose(1, 2)
return output
+2 -2
View File
@@ -8,8 +8,8 @@ class MatrixGameWanVideoArchConfig(WanVideoArchConfig):
# because MatrixGame checkpoints already have patch_embedding.proj format
param_names_mapping: dict = field(
default_factory=lambda: {
# Removed: r"^patch_embedding\.(.*)$": r"patch_embedding.proj.\1"
# because checkpoint already has correct format
r"^patch_embedding\.(?!proj\.)(.*)$":
r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
@@ -0,0 +1,115 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_blocks(n: str, m) -> bool:
return "blocks" in n and str.isdigit(n.split(".")[-1])
@dataclass
class WanGameVideoArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(default_factory=lambda: [is_blocks])
param_names_mapping: dict = field(
default_factory=lambda: {
r"^patch_embedding\.(.*)$":
r"patch_embedding.proj.\1",
r"^condition_embedder\.text_embedder\.linear_1\.(.*)$":
r"condition_embedder.text_embedder.fc_in.\1",
r"^condition_embedder\.text_embedder\.linear_2\.(.*)$":
r"condition_embedder.text_embedder.fc_out.\1",
r"^condition_embedder\.time_embedder\.linear_1\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_in.\1",
r"^condition_embedder\.time_embedder\.linear_2\.(.*)$":
r"condition_embedder.time_embedder.mlp.fc_out.\1",
r"^condition_embedder\.time_proj\.(.*)$":
r"condition_embedder.time_modulation.linear.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.0\.proj\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_in.\1",
r"^condition_embedder\.image_embedder\.ff\.net\.2\.(.*)$":
r"condition_embedder.image_embedder.ff.fc_out.\1",
r"^blocks\.(\d+)\.attn1\.to_q\.(.*)$":
r"blocks.\1.to_q.\2",
r"^blocks\.(\d+)\.attn1\.to_k\.(.*)$":
r"blocks.\1.to_k.\2",
r"^blocks\.(\d+)\.attn1\.to_v\.(.*)$":
r"blocks.\1.to_v.\2",
r"^blocks\.(\d+)\.attn1\.to_out\.0\.(.*)$":
r"blocks.\1.to_out.\2",
r"^blocks\.(\d+)\.attn1\.norm_q\.(.*)$":
r"blocks.\1.norm_q.\2",
r"^blocks\.(\d+)\.attn1\.norm_k\.(.*)$":
r"blocks.\1.norm_k.\2",
r"^blocks\.(\d+)\.attn2\.to_out\.0\.(.*)$":
r"blocks.\1.attn2.to_out.\2",
r"^blocks\.(\d+)\.ffn\.net\.0\.proj\.(.*)$":
r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.net\.2\.(.*)$":
r"blocks.\1.ffn.fc_out.\2",
r"^blocks\.(\d+)\.norm2\.(.*)$":
r"blocks.\1.self_attn_residual_norm.norm.\2",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Some LoRA adapters use the original official layer names instead of hf layer names,
# so apply this before the param_names_mapping
lora_param_names_mapping: dict = field(
default_factory=lambda: {
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.attn1.to_q.\2",
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.attn1.to_k.\2",
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.attn1.to_v.\2",
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$":
r"blocks.\1.attn1.to_out.0.\2",
r"^blocks\.(\d+)\.cross_attn\.q\.(.*)$": r"blocks.\1.attn2.to_q.\2",
r"^blocks\.(\d+)\.cross_attn\.k\.(.*)$": r"blocks.\1.attn2.to_k.\2",
r"^blocks\.(\d+)\.cross_attn\.v\.(.*)$": r"blocks.\1.attn2.to_v.\2",
r"^blocks\.(\d+)\.cross_attn\.o\.(.*)$":
r"blocks.\1.attn2.to_out.0.\2",
r"^blocks\.(\d+)\.ffn\.0\.(.*)$": r"blocks.\1.ffn.fc_in.\2",
r"^blocks\.(\d+)\.ffn\.2\.(.*)$": r"blocks.\1.ffn.fc_out.\2",
})
patch_size: tuple[int, int, int] = (1, 2, 2)
text_len = 512
num_attention_heads: int = 40
attention_head_dim: int = 128
in_channels: int = 16
out_channels: int = 16
text_dim: int = 4096
freq_dim: int = 256
ffn_dim: int = 13824
num_layers: int = 40
cross_attn_norm: bool = True
qk_norm: str = "rms_norm_across_heads"
eps: float = 1e-6
image_dim: int | None = None
added_kv_proj_dim: int | None = None
rope_max_seq_len: int = 1024
pos_embed_seq_len: int | None = None
exclude_lora_layers: list[str] = field(default_factory=lambda: ["embedder"])
# Wan MoE
boundary_ratio: float | None = None
# Causal Wan
local_attn_size: int = -1 # Window size for temporal local attention (-1 indicates global attention)
sink_size: int = 0 # Size of the attention sink, we keep the first `sink_size` frames unchanged when rolling the KV cache
num_frames_per_block: int = 3
sliding_window_num_frames: int = 21
def __post_init__(self):
super().__post_init__()
self.out_channels = self.out_channels or self.in_channels
self.hidden_size = self.num_attention_heads * self.attention_head_dim
self.num_channels_latents = self.out_channels
@dataclass
class WanGameVideoConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=WanGameVideoArchConfig)
prefix: str = "WanGame"
+1
View File
@@ -42,6 +42,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers": HYWorldConfig,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
+6
View File
@@ -196,6 +196,12 @@ class SelfForcingWan2_2_T2V480PConfig(Wan2_2_T2V_A14B_Config):
# =============================================
# ============= Matrix Game ===================
# =============================================
@dataclass
class MatrixGameBaseI2V480PConfig(WanI2V480PConfig):
dit_config: DiTConfig = field(default_factory=MatrixGameWanVideoConfig)
flow_shift: float | None = 5.0
@dataclass
class MatrixGameI2V480PConfig(WanI2V480PConfig):
dit_config: DiTConfig = field(default_factory=MatrixGameWanVideoConfig)
+2
View File
@@ -58,6 +58,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
Wan2_1_Fun_1_3B_Control_SamplingParam,
+6 -6
View File
@@ -113,13 +113,13 @@ class FastWanT2V480P_SamplingParam(WanT2V_1_3B_SamplingParam):
@dataclass
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
height: int = 352
width: int = 640
num_frames: int = 77
fps: int = 25
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale: float = 6.0
num_inference_steps: int = 50
guidance_scale: float = 1.0
num_inference_steps: int = 40
@dataclass
+81
View File
@@ -116,3 +116,84 @@ pyarrow_schema_text_only = pa.schema([
# --- Metadata ---
pa.field("caption", pa.string()),
])
pyarrow_schema_matrixgame = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("first_frame_latent_bytes", pa.binary()),
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
pa.field("first_frame_latent_dtype", pa.string()),
# --- Action ---
pa.field("mouse_cond_bytes", pa.binary()),
pa.field("mouse_cond_shape", pa.list_(pa.int64())), # [T, 2]
pa.field("mouse_cond_dtype", pa.string()),
pa.field("keyboard_cond_bytes", pa.binary()),
pa.field("keyboard_cond_shape", pa.list_(pa.int64())), # [T, 4]
pa.field("keyboard_cond_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
pyarrow_schema_wangame = pa.schema([
pa.field("id", pa.string()),
# --- Image/Video VAE latents ---
# Tensors are stored as raw bytes with shape and dtype info for loading
pa.field("vae_latent_bytes", pa.binary()),
# e.g., [C, T, H, W] or [C, H, W]
pa.field("vae_latent_shape", pa.list_(pa.int64())),
# e.g., 'float32'
pa.field("vae_latent_dtype", pa.string()),
#I2V
pa.field("clip_feature_bytes", pa.binary()),
pa.field("clip_feature_shape", pa.list_(pa.int64())),
pa.field("clip_feature_dtype", pa.string()),
pa.field("first_frame_latent_bytes", pa.binary()),
pa.field("first_frame_latent_shape", pa.list_(pa.int64())),
pa.field("first_frame_latent_dtype", pa.string()),
# --- Action ---
pa.field("mouse_cond_bytes", pa.binary()),
pa.field("mouse_cond_shape", pa.list_(pa.int64())), # [T, 2]
pa.field("mouse_cond_dtype", pa.string()),
pa.field("keyboard_cond_bytes", pa.binary()),
pa.field("keyboard_cond_shape", pa.list_(pa.int64())), # [T, 4]
pa.field("keyboard_cond_dtype", pa.string()),
# I2V Validation
pa.field("pil_image_bytes", pa.binary()),
pa.field("pil_image_shape", pa.list_(pa.int64())),
pa.field("pil_image_dtype", pa.string()),
# --- Metadata ---
pa.field("file_name", pa.string()),
pa.field("caption", pa.string()),
pa.field("media_type", pa.string()), # 'image' or 'video'
pa.field("width", pa.int64()),
pa.field("height", pa.int64()),
# -- Video-specific (can be null/default for images) ---
# Number of frames processed (e.g., 1 for image, N for video)
pa.field("num_frames", pa.int64()),
pa.field("duration_sec", pa.float64()),
pa.field("fps", pa.float64()),
])
+32 -13
View File
@@ -40,6 +40,7 @@ class PreprocessBatch:
num_frames: int | None = None
sample_frame_index: list[int] | None = None
sample_num_frames: int | None = None
action_path: str | None = None
# Processed data
pixel_values: torch.Tensor | None = None
@@ -365,9 +366,11 @@ class ImageTransformStage(DatasetStage):
image = self.transform_topcrop(image)
elif self.transform is not None:
image = self.transform(image)
image = image.float() / 127.5 - 1.0
else:
image = image.float() / 127.5 - 1.0
image = image.transpose(0, 1) # [1 C H W] -> [C 1 H W]
image = image.float() / 127.5 - 1.0
batch.pixel_values = image
return batch
@@ -470,8 +473,11 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
# Initialize tokenizer
tokenizer_path = os.path.join(args.model_path, "tokenizer")
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
if os.path.exists(tokenizer_path):
tokenizer = AutoTokenizer.from_pretrained(tokenizer_path,
cache_dir=args.cache_dir)
else:
tokenizer = None
# Initialize processing stages
self._init_stages(args, transform, transform_topcrop, tokenizer)
@@ -495,11 +501,14 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
self.video_transform_stage = VideoTransformStage(transform)
self.image_transform_stage = ImageTransformStage(
transform, transform_topcrop)
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=args.training_cfg_rate,
seed=self.seed)
if tokenizer is not None:
self.text_encoding_stage = TextEncodingStage(
tokenizer=tokenizer,
text_max_length=args.text_max_length,
cfg_rate=args.training_cfg_rate,
seed=self.seed)
else:
self.text_encoding_stage = None
def _load_raw_data(self) -> list[dict]:
"""Load raw data from JSON files."""
@@ -521,6 +530,8 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
# Update paths with folder prefix
for item in data_items:
item["path"] = opj(folder, item["path"])
if "action_path" in item and item["action_path"]:
item["action_path"] = opj(folder, item["action_path"])
return data_items
@@ -542,7 +553,8 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
cap=item["cap"],
resolution=item.get("resolution"),
fps=item.get("fps"),
duration=item.get("duration"))
duration=item.get("duration"),
action_path=item.get("action_path"))
# Apply filtering stages
if not self._apply_filter_stages(batch, filter_counts):
@@ -604,20 +616,27 @@ class VideoCaptionMergedDataset(torch.utils.data.IterableDataset,
# Apply transformation stages
batch = self.video_transform_stage.process(batch)
batch = self.image_transform_stage.process(batch)
batch = self.text_encoding_stage.process(batch)
if self.text_encoding_stage is not None:
batch = self.text_encoding_stage.process(batch)
# Build result dictionary
result = {
"pixel_values": batch.pixel_values,
"text": batch.text,
"input_ids": batch.input_ids,
"cond_mask": batch.cond_mask,
"path": batch.path,
}
if batch.text is not None:
result["text"] = batch.text
result["input_ids"] = batch.input_ids
result["cond_mask"] = batch.cond_mask
# Add video-specific fields
if batch.is_video:
result.update({"fps": batch.fps, "duration": batch.duration})
# Add action_path
if batch.action_path:
result["action_path"] = batch.action_path
return result
+21
View File
@@ -4,6 +4,7 @@ import os
import pathlib
import datasets
import numpy as np
from torch.utils.data import IterableDataset
from fastvideo.distributed import (get_sp_world_size, get_world_rank,
@@ -160,5 +161,25 @@ class ValidationDataset(IterableDataset):
else:
sample["control_video"] = load_video(control_video_path)
if sample.get("action_path", None) is not None:
action_path = sample["action_path"]
action_path = os.path.join(self.dir, action_path)
sample["action_path"] = action_path
if not pathlib.Path(action_path).is_file():
logger.warning("Action file %s does not exist.", action_path)
else:
try:
action_data = np.load(action_path, allow_pickle=True)
num_frames = sample["num_frames"]
if action_data.dtype == object: action_data = action_data.item()
if isinstance(action_data, dict):
sample["keyboard_cond"] = action_data["keyboard"][:num_frames]
sample["mouse_cond"] = action_data["mouse"][:num_frames]
else:
sample["keyboard_cond"] = action_data[:num_frames]
except Exception as e:
logger.error("Error loading action file %s: %s",
action_path, e)
sample = {k: v for k, v in sample.items() if v is not None}
yield sample
+267 -1
View File
@@ -15,7 +15,7 @@ import torch
from scipy.spatial.transform import Rotation as R
from typing import Union, Optional
from .trajectory import generate_camera_trajectory_local
from fastvideo.models.dits.hyworld.trajectory import generate_camera_trajectory_local
# Mapping from one-hot action encoding to single label
@@ -411,3 +411,269 @@ def compute_num_frames(latent_num: int) -> int:
Number of video frames
"""
return (latent_num - 1) * 4 + 1
def reformat_keyboard_and_mouse_tensors(keyboard_tensor, mouse_tensor):
"""
Reformat the keyboard and mouse tensors to the format compatible with HyWorld.
"""
num_frames = keyboard_tensor.shape[0]
assert (num_frames - 1) % 4 == 0, "num_frames must be a multiple of 4"
assert mouse_tensor.shape[0] == num_frames, "mouse_tensor must have the same number of frames as keyboard_tensor"
keyboard_tensor = keyboard_tensor[1:, :]
mouse_tensor = mouse_tensor[1:, :]
groups = keyboard_tensor.view(-1, 4, keyboard_tensor.shape[1])
assert (groups == groups[:, 0:1]).all(dim=1).all(), "keyboard_tensor must have the same value for each group"
groups = mouse_tensor.view(-1, 4, mouse_tensor.shape[1])
assert (groups == groups[:, 0:1]).all(dim=1).all(), "mouse_tensor must have the same value for each group"
return keyboard_tensor[::4], mouse_tensor[::4]
def process_custom_actions(keyboard_tensor, mouse_tensor, forward_speed=DEFAULT_FORWARD_SPEED):
"""
Process custom keyboard and mouse tensors into model inputs (viewmats, intrinsics, action_labels).
Assumes inputs correspond to each LATENT frame.
"""
if keyboard_tensor.ndim == 3:
keyboard_tensor = keyboard_tensor.squeeze(0)
if mouse_tensor.ndim == 3:
mouse_tensor = mouse_tensor.squeeze(0)
keyboard_tensor, mouse_tensor = reformat_keyboard_and_mouse_tensors(keyboard_tensor, mouse_tensor)
motions = []
# 1. Translate tensors to motions for trajectory generation
for t in range(keyboard_tensor.shape[0]):
frame_motion = {}
# --- Translation ---
# MatrixGame convention: 0:W, 1:S, 2:A, 3:D
fwd = 0.0
if keyboard_tensor[t, 0] > 0.5: fwd += forward_speed # W
if keyboard_tensor[t, 1] > 0.5: fwd -= forward_speed # S
if fwd != 0: frame_motion["forward"] = fwd
rgt = 0.0
if keyboard_tensor[t, 2] > 0.5: rgt -= forward_speed # A (Left is negative Right)
if keyboard_tensor[t, 3] > 0.5: rgt += forward_speed # D (Right)
if rgt != 0: frame_motion["right"] = rgt
# --- Rotation ---
# MatrixGame convention: mouse is [Pitch, Yaw] (or Y, X)
# Apply scaling (e.g. to match HyWorld distribution)
pitch = mouse_tensor[t, 0].item()
yaw = mouse_tensor[t, 1].item()
if abs(pitch) > 1e-4: frame_motion["pitch"] = pitch
if abs(yaw) > 1e-4: frame_motion["yaw"] = yaw
motions.append(frame_motion)
# 2. Generate Camera Trajectory
# generate_camera_trajectory_local returns T+1 poses (starting at Identity)
# We take the first T poses to match the latent count.
# Pose 0 is Identity. Pose 1 is Identity + Motion[0].
poses = generate_camera_trajectory_local(motions)
# poses = np.array(poses[:T])
# 3. Compute Viewmats (w2c) and Intrinsics
w2c_list = []
intrinsic_list = []
# Setup default intrinsic (normalized)
K = np.array(DEFAULT_INTRINSIC)
K[0, 0] /= K[0, 2] * 2
K[1, 1] /= K[1, 2] * 2
K[0, 2] = 0.5
K[1, 2] = 0.5
for i in range(len(poses)):
c2w = np.array(poses[i])
w2c = np.linalg.inv(c2w)
w2c_list.append(w2c)
intrinsic_list.append(K)
viewmats = torch.as_tensor(np.array(w2c_list))
intrinsics = torch.as_tensor(np.array(intrinsic_list))
# 4. Generate Action Labels by analyzing the generated trajectory
# This ensures consistency with complex simultaneous movements, exactly as pose_to_input does.
# Calculate relative camera-to-world transforms
# c2ws = inverse(viewmats)
c2ws = np.linalg.inv(np.array(w2c_list))
# Calculate relative movement between frames
# relative_c2w[i] = inv(c2ws[i-1]) @ c2ws[i]
C_inv = np.linalg.inv(c2ws[:-1])
relative_c2w = np.zeros_like(c2ws)
relative_c2w[0, ...] = c2ws[0, ...] # First is anchor
relative_c2w[1:, ...] = C_inv @ c2ws[1:, ...]
# Initialize one-hot action encodings
trans_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
rotate_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
move_norm_valid = 0.0001
# Skip index 0 (anchor/identity)
for i in range(1, relative_c2w.shape[0]):
move_dirs = relative_c2w[i, :3, 3] # direction vector
move_norms = np.linalg.norm(move_dirs)
if move_norms > move_norm_valid: # threshold for movement
move_norm_dirs = move_dirs / move_norms
angles_rad = np.arccos(move_norm_dirs.clip(-1.0, 1.0))
trans_angles_deg = angles_rad * (180.0 / np.pi) # convert to degrees
else:
trans_angles_deg = np.zeros(3)
R_rel = relative_c2w[i, :3, :3]
r = R.from_matrix(R_rel)
rot_angles_deg = r.as_euler("xyz", degrees=True)
# Determine movement actions based on trajectory
# Note: HyWorld logic checks if rotation is small before assigning translation labels
# to avoid ambiguity in TPS mode, but here we generally want to capture the dominant movement.
tps = False # Default assumption, can be made an arg if needed
if move_norms > move_norm_valid:
if (not tps) or (
tps and abs(rot_angles_deg[1]) < 5e-2 and abs(rot_angles_deg[0]) < 5e-2
):
# Z-axis (Forward/Back)
if trans_angles_deg[2] < 60:
trans_one_hot[i, 0] = 1 # forward
elif trans_angles_deg[2] > 120:
trans_one_hot[i, 1] = 1 # backward
# X-axis (Right/Left)
if trans_angles_deg[0] < 60:
trans_one_hot[i, 2] = 1 # right
elif trans_angles_deg[0] > 120:
trans_one_hot[i, 3] = 1 # left
# Determine rotation actions
# Y-axis (Yaw)
if rot_angles_deg[1] > 5e-2:
rotate_one_hot[i, 0] = 1 # right
elif rot_angles_deg[1] < -5e-2:
rotate_one_hot[i, 1] = 1 # left
# X-axis (Pitch)
if rot_angles_deg[0] > 5e-2:
rotate_one_hot[i, 2] = 1 # up
elif rot_angles_deg[0] < -5e-2:
rotate_one_hot[i, 3] = 1 # down
trans_one_hot = torch.tensor(trans_one_hot)
rotate_one_hot = torch.tensor(rotate_one_hot)
# Convert to single labels
trans_label = one_hot_to_one_dimension(trans_one_hot)
rotate_label = one_hot_to_one_dimension(rotate_one_hot)
action_labels = trans_label * 9 + rotate_label
return viewmats, intrinsics, action_labels
if __name__ == "__main__":
print("Running comparison test between process_custom_actions and pose_to_input...")
def test_process_custom_actions(pose_string: str, keyboard: torch.Tensor, mouse: torch.Tensor, latent_num: int):
# Run process_custom_actions
# Note: We need to pass float tensors
print("Running process_custom_actions...")
viewmats_1, intrinsics_1, labels_1 = process_custom_actions(
keyboard, mouse
)
print(f"Running pose_to_input with string: '{pose_string}'...")
viewmats_2, intrinsics_2, labels_2 = pose_to_input(
pose_string, latent_num=latent_num
)
# print(f"Viewmats: {viewmats_1} vs \n {viewmats_2}")
# print(f"Intrinsics: {intrinsics_1} vs \n {intrinsics_2}")
# print(f"Labels: {labels_1} vs \n {labels_2}")
# 3. Compare Results
print("\nComparison Results:")
# Check Shapes
print(f"Shapes (Viewmats): {viewmats_1.shape} vs {viewmats_2.shape}")
assert viewmats_1.shape == viewmats_2.shape, "Shape mismatch for viewmats"
# Check Values
# Viewmats
diff_viewmats = (viewmats_1 - viewmats_2).abs().max().item()
print(f"Max difference in Viewmats: {diff_viewmats}")
if diff_viewmats < 1e-5:
print("✅ Viewmats match.")
else:
print("❌ Viewmats mismatch.")
# Check intrinsics
diff_intrinsics = (intrinsics_1 - intrinsics_2).abs().max().item()
print(f"Max difference in Intrinsics: {diff_intrinsics}")
if diff_intrinsics < 1e-5:
print("✅ Intrinsics match.")
else:
print("❌ Intrinsics mismatch.")
# Check labels
diff_labels = (labels_1 - labels_2).abs().max().item()
print(f"Max difference in Labels: {diff_labels}")
if diff_labels < 1e-5:
print("✅ Labels match.")
else:
print("❌ Labels mismatch.")
print("All checks passed.")
# Define shared parameters
latent_num = 13
pose_string = "w-2, a-3, s-1, d-6"
num_frames = 4 * (latent_num - 1) + 1
keyboard = torch.zeros((num_frames, 6))
mouse = torch.zeros((num_frames, 2))
# Frame 0 is ignored/start
# Frames 1-8: Press W (index 0)
keyboard[1:9, 0] = 1.0
# Frames 9-20: Press A (index 2)
keyboard[9:21, 2] = 1.0
# Frames 21-24: Press S (index 1)
keyboard[21:25, 1] = 1.0
# Frames 25-48: Press D (index 3)
keyboard[25:49, 3] = 1.0
test_process_custom_actions(pose_string, keyboard, mouse, latent_num)
# Test keyboard AND mouse
latent_num = 25
pose_string = "w-2, up-2, a-3, down-4, s-1, left-2, d-6, right-4"
num_frames = 4 * (latent_num - 1) + 1
keyboard = torch.zeros((num_frames, 6))
mouse = torch.zeros((num_frames, 2))
# Frame 0 is ignored/start
# Frames 1-8: Press W (index 0)
keyboard[1:9, 0] = 1.0
# Frames 17-28: Press A (index 2)
keyboard[17:29, 2] = 1.0
# Frames 45-48: Press S (index 1)
keyboard[45:49, 1] = 1.0
# Frames 57-80: Press D (index 3)
keyboard[57:81, 3] = 1.0
# Frames 9-16: Press Up (index 4)
mouse[9:17, 0] = DEFAULT_PITCH_SPEED
# Frames 25-32: Press Down (index 5)
mouse[29:45, 0] = -DEFAULT_PITCH_SPEED
# Frames 41-48: Press Left (index 6)
mouse[49:57, 1] = -DEFAULT_YAW_SPEED
# Frames 57-64: Press Right (index 7)
mouse[81:, 1] = DEFAULT_YAW_SPEED
test_process_custom_actions(pose_string, keyboard, mouse, latent_num)
@@ -257,6 +257,17 @@ class ActionModule(nn.Module):
'''
assert use_rope_keyboard
target_device = x.device
target_dtype = x.dtype
if mouse_condition is not None:
mouse_condition = mouse_condition.to(device=target_device,
dtype=target_dtype)
if keyboard_condition is not None:
keyboard_condition = keyboard_condition.to(
device=target_device, dtype=target_dtype)
else:
return x
B, N_frames, C = keyboard_condition.shape
assert tt*th*tw == x.shape[1]
assert ((N_frames - 1) + self.vae_time_compression_ratio) % self.vae_time_compression_ratio == 0
@@ -272,7 +283,9 @@ class ActionModule(nn.Module):
# Defined freqs_cis early so it's available for both mouse and keyboard
freqs_cis = (self._freqs_cos, self._freqs_sin)
assert (N_feats == tt and ((is_causal and kv_cache_mouse is None) or not is_causal)) or ((N_frames - 1) // self.vae_time_compression_ratio + 1 == start_frame + num_frame_per_block and is_causal)
if is_causal:
assert (N_feats == tt and kv_cache_mouse is None) or ((N_frames - 1) // self.vae_time_compression_ratio + 1 == start_frame + num_frame_per_block)
# For non-causal (training), we trust that the caller provides correctly shaped inputs
if self.enable_mouse and mouse_condition is not None:
hidden_states = rearrange(x, "B (T S) C -> (B S) T C", T=tt, S=th*tw) # 65*272*480 -> 17*(272//16)*(480//16) -> 8670
@@ -289,13 +302,15 @@ class ActionModule(nn.Module):
mouse_condition = mouse_condition[:, self.vae_time_compression_ratio*(N_feats - num_frame_per_block - self.windows_size) + pad_t:, :]
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(num_frame_per_block)]
else:
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(N_feats)]
local_num_frames = tt
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(local_num_frames)]
group_mouse = torch.stack(group_mouse, dim = 1)
actual_num_frames = group_mouse.shape[1] # Use actual stacked frame count
S = th * tw
group_mouse = group_mouse.unsqueeze(-1).expand(B, num_frame_per_block, pad_t, C, S)
group_mouse = group_mouse.permute(0, 4, 1, 2, 3).reshape(B * S, num_frame_per_block, pad_t * C)
group_mouse = group_mouse.unsqueeze(-1).expand(B, actual_num_frames, pad_t, C, S)
group_mouse = group_mouse.permute(0, 4, 1, 2, 3).reshape(B * S, actual_num_frames, pad_t * C)
group_mouse = torch.cat([hidden_states, group_mouse], dim = -1)
group_mouse = self.mouse_mlp(group_mouse)
@@ -314,7 +329,7 @@ class ActionModule(nn.Module):
## TODO: adding cache here
if is_causal:
if kv_cache_mouse is None:
assert q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0 # == 880, f"{q.shape[0]},{k.shape[0]}"
assert q.shape[0] == k.shape[0] and q.shape[0] % S == 0 # == 880, f"{q.shape[0]},{k.shape[0]}"
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
padded_q = torch.cat(
[q,
@@ -395,7 +410,8 @@ class ActionModule(nn.Module):
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(num_frame_per_block)]
else:
keyboard_condition = self.keyboard_embed(keyboard_condition)
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(N_feats)]
local_num_frames = tt
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(local_num_frames)]
group_keyboard = torch.stack(group_keyboard, dim = 1) # B F RW C
group_keyboard = group_keyboard.reshape(shape=(group_keyboard.shape[0],group_keyboard.shape[1],-1))
# apply cross attn
@@ -415,7 +431,7 @@ class ActionModule(nn.Module):
q = self.key_attn_q_norm(q).to(v)
k = self.key_attn_k_norm(k).to(v)
S = th * tw
assert S == 880
# assert S == 880
# position embed
if use_rope_keyboard:
B, TS, H, D = q.shape
@@ -424,13 +440,13 @@ class ActionModule(nn.Module):
q, k = _apply_rotary_emb_qk(q, k, freqs_cis[0], freqs_cis[1], start_offset=start_frame)
k1, k2, k3, k4 = k.shape
k = k.expand(S, k2, k3, k4)
v = v.expand(S, k2, k3, k4)
k = k.repeat_interleave(S, dim=0)
v = v.repeat_interleave(S, dim=0)
if is_causal:
if kv_cache_keyboard is None:
assert q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0
assert q.shape[0] == k.shape[0] and q.shape[0] % S == 0
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
padded_q = torch.cat(
@@ -480,7 +496,7 @@ class ActionModule(nn.Module):
else:
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - kv_cache_keyboard["global_end_index"].item()
local_start_index = local_end_index - num_new_tokens
assert k.shape[0] == 880 # BS == 1 or the cache should not be saved/ load method should be modified
assert k.shape[0] == S # BS == 1 or the cache should not be saved/ load method should be modified
kv_cache_keyboard["k"][:, local_start_index:local_end_index] = k[:1]
kv_cache_keyboard["v"][:, local_start_index:local_end_index] = v[:1]
@@ -1096,4 +1096,4 @@ class CausalMatrixGameWanModel(BaseDiT):
if block.action_model.proj_keyboard.bias is not None:
nn.init.zeros_(block.action_model.proj_keyboard.bias)
except AttributeError:
pass
pass
@@ -61,14 +61,17 @@ class MatrixGameTimeImageEmbedding(nn.Module):
timestep_proj = self.time_modulation(temb)
# MatrixGame does not use text embeddings, so we ignore encoder_hidden_states
# and return None for the text embedding part
if encoder_hidden_states_image is not None:
assert self.image_embedder is not None
encoder_hidden_states_image = self.image_embedder(
encoder_hidden_states_image)
return temb, timestep_proj, None, encoder_hidden_states_image
encoder_hidden_states = torch.zeros((timestep.shape[0], 0, temb.shape[-1]),
device=temb.device,
dtype=temb.dtype)
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
class MatrixGameCrossAttention(WanSelfAttention):
@@ -313,6 +316,7 @@ class MatrixGameWanModel(BaseDiT):
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.num_channels_latents = config.out_channels
self.patch_size = config.patch_size
# 1. Patch & position embedding
@@ -395,7 +399,8 @@ class MatrixGameWanModel(BaseDiT):
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000)
rope_theta=10000,
do_sp_sharding=True)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos.float(),
@@ -157,8 +157,8 @@ def load_initial_image(image_path: str = None) -> Image.Image:
return Image.new("RGB", (640, 352), (128, 128, 128))
def create_action_presets(num_frames: int, keyboard_dim: int = 4, seed: int = None):
if keyboard_dim not in (2, 4, 7):
raise ValueError(f"keyboard_dim must be 2, 4, or 7, got {keyboard_dim}")
if keyboard_dim not in (2, 4, 6, 7):
raise ValueError(f"keyboard_dim must be 2, 4, 6, or 7, got {keyboard_dim}")
if num_frames % 4 != 1:
raise ValueError("Matrix-Game conditioning expects num_frames to be 4k+1.")
@@ -181,6 +181,11 @@ def create_action_presets(num_frames: int, keyboard_dim: int = 4, seed: int = No
actions_double_action = []
actions_single_camera = ["camera_l", "camera_r"]
keyboard_idx = {"forward": 0, "back": 1}
elif keyboard_dim == 6:
actions_single_action = ["forward", "back", "left", "right"]
actions_double_action = ["forward_left", "forward_right"]
actions_single_camera = ["camera_l", "camera_r"]
keyboard_idx = {"forward": 0, "back": 1, "left": 2, "right": 3, "t1": 4, "t2": 5}
else: # keyboard_dim == 7
# Temple Run model: still, w, s, left, right, a, d (no mouse)
actions_single_action = ["forward", "back", "left", "right"]
@@ -0,0 +1,5 @@
from .model import WanGameActionTransformer3DModel
__all__ = [
"WanGameActionTransformer3DModel",
]
@@ -0,0 +1,231 @@
import torch
import torch.nn as nn
from fastvideo.layers.visual_embedding import TimestepEmbedder, ModulateProjection, timestep_embedding
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.attention import DistributedAttention
from fastvideo.forward_context import set_forward_context
from fastvideo.models.dits.wanvideo import WanImageEmbedding
from fastvideo.models.dits.hyworld.camera_rope import prope_qkv
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
from fastvideo.layers.mlp import MLP
class WanGameActionTimeImageEmbedding(nn.Module):
def __init__(
self,
dim: int,
time_freq_dim: int,
image_embed_dim: int | None = None,
):
super().__init__()
self.time_freq_dim = time_freq_dim
self.time_embedder = TimestepEmbedder(
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
self.time_modulation = ModulateProjection(dim,
factor=6,
act_layer="silu")
self.image_embedder = None
if image_embed_dim is not None:
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
self.action_embedder = MLP(
time_freq_dim,
dim,
dim,
bias=True,
act_type="silu"
)
# Initialize with zeros for residual-like behavior
nn.init.zeros_(self.action_embedder.fc_out.weight)
if self.action_embedder.fc_out.bias is not None:
nn.init.zeros_(self.action_embedder.fc_out.bias)
def forward(
self,
timestep: torch.Tensor,
action: torch.Tensor,
encoder_hidden_states: torch.Tensor, # Kept for interface compatibility
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
):
temb = self.time_embedder(timestep, timestep_seq_len)
action_emb = timestep_embedding(action.flatten(), self.time_freq_dim)
action_embedder_dtype = next(iter(self.action_embedder.parameters())).dtype
if (
action_emb.dtype != action_embedder_dtype
and action_embedder_dtype != torch.int8
):
action_emb = action_emb.to(action_embedder_dtype)
action_emb = self.action_embedder(action_emb).type_as(temb)
temb = temb + action_emb
timestep_proj = self.time_modulation(temb)
# MatrixGame does not use text embeddings, so we ignore encoder_hidden_states
if encoder_hidden_states_image is not None:
assert self.image_embedder is not None
encoder_hidden_states_image = self.image_embedder(
encoder_hidden_states_image)
encoder_hidden_states = torch.zeros((timestep.shape[0], 0, temb.shape[-1]),
device=temb.device,
dtype=temb.dtype)
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
class WanGameActionSelfAttention(nn.Module):
"""
Self-attention module with support for:
- Standard RoPE-based attention
- Camera PRoPE-based attention (when viewmats and Ks are provided)
- KV caching for autoregressive generation
"""
def __init__(self,
dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm=True,
eps=1e-6) -> None:
assert dim % num_heads == 0
super().__init__()
self.dim = dim
self.num_heads = num_heads
self.head_dim = dim // num_heads
self.local_attn_size = local_attn_size
self.sink_size = sink_size
self.qk_norm = qk_norm
self.eps = eps
self.max_attention_size = 32760 if local_attn_size == -1 else local_attn_size * 1560
# Scaled dot product attention (using DistributedAttention for SP support)
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=self.head_dim,
softmax_scale=None,
causal=False,
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA))
def forward(self,
q: torch.Tensor,
k: torch.Tensor,
v: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
kv_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
is_cache: bool = False,
attention_mask: torch.Tensor | None = None):
"""
Forward pass with camera PRoPE attention combining standard RoPE and projective positional encoding.
Args:
q, k, v: Query, key, value tensors [B, L, num_heads, head_dim]
freqs_cis: RoPE frequency cos/sin tensors
kv_cache: KV cache dict (may have None values for training)
current_start: Current position for KV cache
cache_start: Cache start position
viewmats: Camera view matrices for PRoPE [B, cameras, 4, 4]
Ks: Camera intrinsics for PRoPE [B, cameras, 3, 3]
is_cache: Whether to store to KV cache (for inference)
attention_mask: Attention mask [B, L] (1 = attend, 0 = mask)
"""
if cache_start is None:
cache_start = current_start
# Apply RoPE manually
cos, sin = freqs_cis
query_rope = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
key_rope = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
value_rope = v
# Get PRoPE transformed q, k, v
query_prope, key_prope, value_prope, apply_fn_o = prope_qkv(
q.transpose(1, 2), # [B, num_heads, L, head_dim]
k.transpose(1, 2),
v.transpose(1, 2),
viewmats=viewmats,
Ks=Ks,
patches_x=40, # hardcoded for now
patches_y=22,
)
# PRoPE returns [B, num_heads, L, head_dim], convert to [B, L, num_heads, head_dim]
query_prope = query_prope.transpose(1, 2)
key_prope = key_prope.transpose(1, 2)
value_prope = value_prope.transpose(1, 2)
# KV cache handling
if kv_cache is not None:
cache_key = kv_cache.get("k", None)
cache_value = kv_cache.get("v", None)
if cache_value is not None and not is_cache:
cache_key_rope, cache_key_prope = cache_key.chunk(2, dim=-1)
cache_value_rope, cache_value_prope = cache_value.chunk(2, dim=-1)
key_rope = torch.cat([cache_key_rope, key_rope], dim=1)
value_rope = torch.cat([cache_value_rope, value_rope], dim=1)
key_prope = torch.cat([cache_key_prope, key_prope], dim=1)
value_prope = torch.cat([cache_value_prope, value_prope], dim=1)
if is_cache:
# Store to cache (update input dict directly)
kv_cache["k"] = torch.cat([key_rope, key_prope], dim=-1)
kv_cache["v"] = torch.cat([value_rope, value_prope], dim=-1)
# Concatenate rope and prope paths (matching original)
query_all = torch.cat([query_rope, query_prope], dim=0)
key_all = torch.cat([key_rope, key_prope], dim=0)
value_all = torch.cat([value_rope, value_prope], dim=0)
# Check if Q and KV have different sequence lengths (KV cache mode)
# In this case, use LocalAttention (supports different Q/KV lengths)
if query_all.shape[1] != key_all.shape[1]:
raise ValueError("Q and KV have different sequence lengths")
# KV cache mode: Q has new tokens only, KV has cached + new tokens
# Use LocalAttention which supports different Q/KV lengths
# LocalAttention will use the appropriate backend (SageAttn, FlashAttn, etc.)
if not hasattr(self, '_kv_cache_attn'):
from fastvideo.attention import LocalAttention
self._kv_cache_attn = LocalAttention(
num_heads=self.num_heads,
head_size=self.head_dim,
causal=False,
supported_attention_backends=(AttentionBackendEnum.SAGE_ATTN,
AttentionBackendEnum.FLASH_ATTN,
AttentionBackendEnum.TORCH_SDPA)
)
hidden_states_all = self._kv_cache_attn(query_all, key_all, value_all)
else:
# Same sequence length: use DistributedAttention (supports SP)
# Create default attention mask if not provided
if attention_mask is None:
batch_size, seq_len = q.shape[0], q.shape[1]
attention_mask = torch.ones(batch_size, seq_len, device=q.device, dtype=q.dtype)
if q.dtype == torch.float32:
from fastvideo.attention.backends.sdpa import SDPAMetadataBuilder
attn_metadata_builder = SDPAMetadataBuilder
else:
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata_builder = FlashAttnMetadataBuilder
attn_metadata = attn_metadata_builder().build(
current_timestep=0,
attn_mask=attention_mask,
)
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
hidden_states_all, _ = self.attn(query_all, key_all, value_all, attention_mask=attention_mask)
hidden_states_rope, hidden_states_prope = hidden_states_all.chunk(2, dim=0)
hidden_states_prope = apply_fn_o(hidden_states_prope.transpose(1, 2)).transpose(1, 2)
return hidden_states_rope, hidden_states_prope
+424
View File
@@ -0,0 +1,424 @@
# SPDX-License-Identifier: Apache-2.0
import math
from typing import Any
import torch
import torch.nn as nn
from fastvideo.configs.models.dits.wangamevideo import WanGameVideoConfig
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
RMSNorm, ScaleResidual,
ScaleResidualLayerNormScaleShift)
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.mlp import MLP
from fastvideo.layers.rotary_embedding import (_apply_rotary_emb,
get_rotary_pos_embed)
from fastvideo.layers.visual_embedding import PatchEmbed
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanI2VCrossAttention
from fastvideo.platforms import AttentionBackendEnum, current_platform
# Import ActionModule
from fastvideo.models.dits.wangame.hyworld_action_module import WanGameActionTimeImageEmbedding, WanGameActionSelfAttention
logger = init_logger(__name__)
class WanGameCrossAttention(WanI2VCrossAttention):
def forward(self, x, context, context_lens=None):
r"""
Args:
x(Tensor): Shape [B, L1, C]
context(Tensor): Shape [B, L2, C]
context_lens(Tensor): Shape [B]
"""
context_img = context
b, n, d = x.size(0), self.num_heads, self.head_dim
# compute query, key, value
q = self.norm_q(self.to_q(x)[0]).view(b, -1, n, d)
k_img = self.norm_added_k(self.add_k_proj(context_img)[0]).view(
b, -1, n, d)
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
# output
x = img_x.flatten(2)
x, _ = self.to_out(x)
return x
class WanGameActionTransformerBlock(nn.Module):
"""
Transformer block for WAN Action model with support for:
- Self-attention with RoPE and camera PRoPE
- Cross-attention with text/image context
- Feed-forward network with AdaLN modulation
"""
def __init__(self,
dim: int,
ffn_dim: int,
num_heads: int,
local_attn_size: int = -1,
sink_size: int = 0,
qk_norm: str = "rms_norm_across_heads",
cross_attn_norm: bool = False,
eps: float = 1e-6,
added_kv_proj_dim: int | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
prefix: str = ""):
super().__init__()
# 1. Self-attention
self.norm1 = FP32LayerNorm(dim, eps, elementwise_affine=False)
self.to_q = ReplicatedLinear(dim, dim, bias=True)
self.to_k = ReplicatedLinear(dim, dim, bias=True)
self.to_v = ReplicatedLinear(dim, dim, bias=True)
self.to_out = ReplicatedLinear(dim, dim, bias=True)
self.attn1 = WanGameActionSelfAttention(
dim,
num_heads,
local_attn_size=local_attn_size,
sink_size=sink_size,
qk_norm=qk_norm,
eps=eps)
self.hidden_dim = dim
self.num_attention_heads = num_heads
self.local_attn_size = local_attn_size
dim_head = dim // num_heads
if qk_norm == "rms_norm":
self.norm_q = RMSNorm(dim_head, eps=eps)
self.norm_k = RMSNorm(dim_head, eps=eps)
elif qk_norm == "rms_norm_across_heads":
self.norm_q = RMSNorm(dim, eps=eps)
self.norm_k = RMSNorm(dim, eps=eps)
else:
raise ValueError(f"QK Norm type {qk_norm} not supported")
assert cross_attn_norm is True
self.self_attn_residual_norm = ScaleResidualLayerNormScaleShift(
dim,
norm_type="layer",
eps=eps,
elementwise_affine=True,
compute_dtype=torch.float32)
# 2. Cross-attention (I2V only for now)
self.attn2 = WanGameCrossAttention(dim,
num_heads,
qk_norm=qk_norm,
eps=eps)
# norm3 for FFN input
self.norm3 = LayerNormScaleShift(dim, norm_type="layer", eps=eps,
elementwise_affine=False)
# 3. Feed-forward
self.ffn = MLP(dim, ffn_dim, act_type="gelu_pytorch_tanh")
self.mlp_residual = ScaleResidual()
self.scale_shift_table = nn.Parameter(torch.randn(1, 6, dim) / dim**0.5)
# PRoPE output projection (initialized via add_discrete_action_parameters on the model)
self.to_out_prope = nn.ModuleList([
nn.Linear(dim, dim, bias=True),
])
nn.init.zeros_(self.to_out_prope[0].weight)
if self.to_out_prope[0].bias is not None:
nn.init.zeros_(self.to_out_prope[0].bias)
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor,
temb: torch.Tensor,
freqs_cis: tuple[torch.Tensor, torch.Tensor],
kv_cache: dict | None = None,
crossattn_cache: dict | None = None,
current_start: int = 0,
cache_start: int | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
is_cache: bool = False,
) -> torch.Tensor:
if hidden_states.dim() == 4:
hidden_states = hidden_states.squeeze(1)
num_frames = temb.shape[1]
frame_seqlen = hidden_states.shape[1] // num_frames
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# Cast temb to float32 for scale/shift computation
e = self.scale_shift_table + temb.float()
assert e.shape == (bs, num_frames, 6, self.hidden_dim)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(6, dim=2)
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen)) *
(1 + scale_msa) + shift_msa).to(orig_dtype).flatten(1, 2)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
if self.norm_q is not None:
query = self.norm_q.forward_native(query)
if self.norm_k is not None:
key = self.norm_k.forward_native(key)
query = query.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
key = key.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
value = value.squeeze(1).unflatten(2, (self.num_attention_heads, -1))
# Self-attention with optional camera PRoPE
attn_output_rope, attn_output_prope = self.attn1(
query, key, value, freqs_cis,
kv_cache, current_start, cache_start, viewmats, Ks,
is_cache=is_cache
)
# Combine rope and prope outputs
attn_output_rope = attn_output_rope.flatten(2)
attn_output_rope, _ = self.to_out(attn_output_rope)
attn_output_prope = attn_output_prope.flatten(2)
attn_output_prope = self.to_out_prope[0](attn_output_prope)
attn_output = attn_output_rope.squeeze(1) + attn_output_prope.squeeze(1)
# Self-attention residual + norm in float32
null_shift = null_scale = torch.zeros(1, device=hidden_states.device, dtype=torch.float32)
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
hidden_states.float(), attn_output.float(), gate_msa, null_shift, null_scale)
hidden_states = hidden_states.type_as(attn_output)
norm_hidden_states = norm_hidden_states.type_as(attn_output)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states.to(orig_dtype),
context=encoder_hidden_states,
context_lens=None)
# Cross-attention residual in bfloat16
hidden_states = hidden_states + attn_output
# norm3 for FFN input in float32
norm_hidden_states = self.norm3(
hidden_states.float(), c_shift_msa, c_scale_msa
).type_as(hidden_states)
# 3. Feed-forward
ff_output = self.ffn(norm_hidden_states.to(orig_dtype))
hidden_states = self.mlp_residual(hidden_states.float(), ff_output.float(), c_gate_msa)
hidden_states = hidden_states.to(orig_dtype) # Cast back to original dtype
return hidden_states
class WanGameActionTransformer3DModel(BaseDiT):
"""
WAN Action Transformer 3D Model for video generation with action conditioning.
Extends the base WAN video model with:
- Action embedding support for controllable generation
- camera PRoPE attention for 3D-aware generation
- KV caching for autoregressive inference
"""
_fsdp_shard_conditions = WanGameVideoConfig()._fsdp_shard_conditions
_compile_conditions = WanGameVideoConfig()._compile_conditions
_supported_attention_backends = WanGameVideoConfig()._supported_attention_backends
param_names_mapping = WanGameVideoConfig().param_names_mapping
reverse_param_names_mapping = WanGameVideoConfig().reverse_param_names_mapping
lora_param_names_mapping = WanGameVideoConfig().lora_param_names_mapping
def __init__(self, config: WanGameVideoConfig, hf_config: dict[str, Any]) -> None:
super().__init__(config=config, hf_config=hf_config)
inner_dim = config.num_attention_heads * config.attention_head_dim
self.hidden_size = config.hidden_size
self.num_attention_heads = config.num_attention_heads
self.attention_head_dim = config.attention_head_dim
self.in_channels = config.in_channels
self.out_channels = config.out_channels
self.num_channels_latents = config.num_channels_latents
self.patch_size = config.patch_size
self.local_attn_size = config.local_attn_size
self.inner_dim = inner_dim
# 1. Patch & position embedding
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
embed_dim=inner_dim,
patch_size=config.patch_size,
flatten=False)
# 2. Condition embeddings (with action support)
self.condition_embedder = WanGameActionTimeImageEmbedding(
dim=inner_dim,
time_freq_dim=config.freq_dim,
image_embed_dim=config.image_dim,
)
# 3. Transformer blocks
self.blocks = nn.ModuleList([
WanGameActionTransformerBlock(
inner_dim,
config.ffn_dim,
config.num_attention_heads,
config.local_attn_size,
config.sink_size,
config.qk_norm,
config.cross_attn_norm,
config.eps,
config.added_kv_proj_dim,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.blocks.{i}")
for i in range(config.num_layers)
])
# 4. Output norm & projection
self.norm_out = LayerNormScaleShift(inner_dim,
norm_type="layer",
eps=config.eps,
elementwise_affine=False,
dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
self.scale_shift_table = nn.Parameter(torch.randn(1, 2, inner_dim) / inner_dim**0.5)
self.gradient_checkpointing = False
# Causal-specific
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
self.__post_init__()
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: torch.Tensor | list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: torch.Tensor | list[torch.Tensor],
guidance=None,
action: torch.Tensor | None = None,
viewmats: torch.Tensor | None = None,
Ks: torch.Tensor | None = None,
kv_cache: list[dict] | None = None,
crossattn_cache: list[dict] | None = None,
current_start: int = 0,
cache_start: int = 0,
start_frame: int = 0,
is_cache: bool = False,
**kwargs
) -> torch.Tensor:
"""
Forward pass for both training and inference with KV caching.
Args:
hidden_states: Video latents [B, C, T, H, W]
encoder_hidden_states: Text embeddings [B, L, D]
timestep: Timestep tensor
encoder_hidden_states_image: Optional image embeddings
action: Action tensor [B, T] for per-frame conditioning
viewmats: Camera view matrices for PRoPE [B, T, 4, 4]
Ks: Camera intrinsics for PRoPE [B, T, 3, 3]
kv_cache: KV cache for autoregressive inference (list of dicts per layer)
crossattn_cache: Cross-attention cache for inference
current_start: Current position for KV cache
cache_start: Cache start position
start_frame: RoPE offset for new frames in autoregressive mode
is_cache: If True, populate KV cache and return early (cache-only mode)
"""
orig_dtype = hidden_states.dtype
# if not isinstance(encoder_hidden_states, torch.Tensor):
# encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image, list) and len(encoder_hidden_states_image) > 0:
encoder_hidden_states_image = encoder_hidden_states_image[0]
# else:
# encoder_hidden_states_image = None
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height, post_patch_width),
self.hidden_size,
self.num_attention_heads,
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000,
start_frame=start_frame
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
if timestep.dim() == 2:
timestep = timestep.flatten()
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, action, encoder_hidden_states, encoder_hidden_states_image=encoder_hidden_states_image)
# Reshape timestep_proj: [T, 6*dim] -> [B, T, 6, dim]
# For training: batch_size=1, T=num_frames (diffusion forcing)
# For inference: batch_size can vary
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size))
if timestep_proj.shape[0] == post_patch_num_frames and batch_size == 1:
# Training mode: timestep_proj is [T, 6, dim], add batch dim -> [1, T, 6, dim]
timestep_proj = timestep_proj.unsqueeze(0)
else:
# Inference mode: reshape based on timestep shape
timestep_proj = timestep_proj.unflatten(dim=0, sizes=timestep.shape)
encoder_hidden_states = encoder_hidden_states_image
# Transformer blocks
for block_idx, block in enumerate(self.blocks):
if torch.is_grad_enabled() and self.gradient_checkpointing:
hidden_states = self._gradient_checkpointing_func(
block, hidden_states, encoder_hidden_states, timestep_proj, freqs_cis,
kv_cache[block_idx] if kv_cache else None,
crossattn_cache[block_idx] if crossattn_cache else None,
current_start, cache_start,
viewmats, Ks, is_cache)
else:
hidden_states = block(
hidden_states, encoder_hidden_states, timestep_proj, freqs_cis,
kv_cache[block_idx] if kv_cache else None,
crossattn_cache[block_idx] if crossattn_cache else None,
current_start, cache_start,
viewmats, Ks, is_cache)
# If cache-only mode, return early
if is_cache:
return kv_cache
# Output norm, projection & unpatchify
# Reshape temb to match timestep_proj shape: [T, dim] -> [B, T, 1, dim]
if temb.shape[0] == post_patch_num_frames and batch_size == 1:
# Training mode: temb is [T, dim] -> [1, T, 1, dim]
temb = temb.unsqueeze(0).unsqueeze(2)
else:
# Inference mode: reshape based on timestep shape
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2, dim=2)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
post_patch_width, p_t, p_h, p_w,
-1)
hidden_states = hidden_states.permute(0, 7, 1, 4, 2, 5, 3, 6)
output = hidden_states.flatten(6, 7).flatten(4, 5).flatten(2, 3)
return output
+19 -7
View File
@@ -76,7 +76,7 @@ class WanTimeTextImageEmbedding(nn.Module):
dim,
dim,
bias=True,
act_type="gelu_pytorch_tanh")
act_type="gelu_pytorch_tanh") if text_embed_dim > 0 else None
self.image_embedder = None
if image_embed_dim is not None:
@@ -92,7 +92,10 @@ class WanTimeTextImageEmbedding(nn.Module):
temb = self.time_embedder(timestep, timestep_seq_len)
timestep_proj = self.time_modulation(temb)
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
if self.text_embedder is not None:
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
else:
encoder_hidden_states = torch.zeros((timestep.shape[0], 0, temb.shape[-1]), device=temb.device, dtype=temb.dtype)
if encoder_hidden_states_image is not None:
assert self.image_embedder is not None
encoder_hidden_states_image = self.image_embedder(
@@ -179,7 +182,10 @@ class WanT2VCrossAttention(WanSelfAttention):
v = self.to_v(context)[0].view(b, -1, n, d)
# compute attention
x = self.attn(q, k, v)
if k.size(1) > 0:
x = self.attn(q, k, v)
else:
x = torch.zeros_like(q)
# output
x = x.flatten(2)
@@ -227,7 +233,10 @@ class WanI2VCrossAttention(WanSelfAttention):
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
img_x = self.attn(q, k_img, v_img)
# compute attention
x = self.attn(q, k, v)
if k.size(1) > 0:
x = self.attn(q, k, v)
else:
x = torch.zeros_like(q)
# output
x = x.flatten(2)
@@ -633,7 +642,7 @@ class WanTransformer3DModel(CachableDiT):
enable_teacache = forward_batch is not None and forward_batch.enable_teacache
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
if encoder_hidden_states is not None and not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
if isinstance(encoder_hidden_states_image,
list) and len(encoder_hidden_states_image) > 0:
@@ -705,8 +714,11 @@ class WanTransformer3DModel(CachableDiT):
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
if encoder_hidden_states is not None:
encoder_hidden_states = torch.concat(
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
else:
encoder_hidden_states = encoder_hidden_states_image
if current_platform.is_mps() or current_platform.is_npu():
encoder_hidden_states = encoder_hidden_states.to(orig_dtype)
@@ -806,6 +806,10 @@ class TransformerLoader(ComponentLoader):
cls_name.startswith("Cosmos25")
or cls_name == "Cosmos25Transformer3DModel"
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "Cosmos25"
) and not (
cls_name.startswith("WanGame")
or cls_name == "WanGameActionTransformer3DModel"
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "WanGame"
)
model = maybe_load_fsdp_model(
model_cls=model_cls,
+5 -2
View File
@@ -138,7 +138,7 @@ def maybe_load_fsdp_model(
weight_iterator = safetensors_weights_iterator(weight_dir_list)
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
load_model_from_full_model_state_dict(
incompatible_keys, unexpected_keys = load_model_from_full_model_state_dict(
model,
weight_iterator,
device,
@@ -147,6 +147,9 @@ def maybe_load_fsdp_model(
cpu_offload=cpu_offload,
param_names_mapping=param_names_mapping_fn,
)
if incompatible_keys or unexpected_keys:
logger.warning("Incompatible keys: %s", incompatible_keys)
logger.warning("Unexpected keys: %s", unexpected_keys)
for n, p in chain(model.named_parameters(), model.named_buffers()):
if p.is_meta:
raise RuntimeError(
@@ -340,7 +343,7 @@ def load_model_from_full_model_state_dict(
unused_keys)
# List of allowed parameter name patterns
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # Can be extended as needed
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l", "to_out_prope", "action_embedder"] # Can be extended as needed
for new_param_name in unused_keys:
if not any(pattern in new_param_name
for pattern in ALLOWED_NEW_PARAM_PATTERNS):
+3 -2
View File
@@ -42,8 +42,9 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
"MatrixGameWanModel": ("dits", "matrix_game", "MatrixGameWanModel"),
"CausalMatrixGameWanModel": ("dits", "matrix_game", "CausalMatrixGameWanModel"),
"WanGameActionTransformer3DModel": ("dits", "wangame", "WanGameActionTransformer3DModel"),
"MatrixGameWanModel": ("dits", "matrixgame", "MatrixGameWanModel"),
"CausalMatrixGameWanModel": ("dits", "matrixgame", "CausalMatrixGameWanModel"),
}
_TEXT_ENCODER_MODELS = {
@@ -0,0 +1,74 @@
# SPDX-License-Identifier: Apache-2.0
"""
Wan video diffusion pipeline implementation.
This module contains an implementation of the Wan video diffusion pipeline
using the modular pipeline architecture.
"""
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.logger import init_logger
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
# isort: off
from fastvideo.pipelines.stages import (
ImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
TimestepPreparationStage)
# isort: on
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
logger = init_logger(__name__)
class WanGameActionImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
_required_config_modules = [
"vae", "transformer", "scheduler", \
"image_encoder", "image_processor"
]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
"""Set up pipeline stages with proper dependency injection."""
self.add_stage(stage_name="input_validation_stage",
stage=InputValidationStage())
self.add_stage(
stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
self.add_stage(stage_name="conditioning_stage",
stage=ConditioningStage())
self.add_stage(stage_name="timestep_preparation_stage",
stage=TimestepPreparationStage(
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="latent_preparation_stage",
stage=LatentPreparationStage(
scheduler=self.get_module("scheduler"),
transformer=self.get_module("transformer")))
self.add_stage(stage_name="image_latent_preparation_stage",
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
self.add_stage(stage_name="denoising_stage",
stage=DenoisingStage(
transformer=self.get_module("transformer"),
scheduler=self.get_module("scheduler")))
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
EntryClass = WanGameActionImageToVideoPipeline
+1
View File
@@ -21,6 +21,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
"WanPipeline": "wan",
"WanDMDPipeline": "wan",
"WanImageToVideoPipeline": "wan",
"WanGameActionImageToVideoPipeline": "wan",
"WanVideoToVideoPipeline": "wan",
"WanCausalDMDPipeline": "wan",
"TurboDiffusionPipeline": "turbodiffusion",
@@ -0,0 +1,303 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Any
import numpy as np
import torch
from PIL import Image
from fastvideo.dataset.dataloader.schema import pyarrow_schema_matrixgame
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import ImageEncodingStage
class PreprocessPipeline_MatrixGame(BasePreprocessPipeline):
"""I2V preprocessing pipeline implementation."""
_required_config_modules = ["vae", "image_encoder", "image_processor"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
def get_pyarrow_schema(self):
"""Return the PyArrow schema for I2V pipeline."""
return pyarrow_schema_matrixgame
def get_extra_features(self, valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
# TODO(will): move these to cpu at some point
self.get_module("image_encoder").to(get_local_torch_device())
self.get_module("vae").to(get_local_torch_device())
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
_, _, num_frames, height, width = valid_data["pixel_values"].shape
# latent_height = height // self.get_module(
# "vae").spatial_compression_ratio
# latent_width = width // self.get_module("vae").spatial_compression_ratio
processed_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
processed_img = self.get_module("image_processor")(
images=frame_pil, return_tensors="pt")
processed_images.append(processed_img)
# Get CLIP features
pixel_values = torch.cat(
[img['pixel_values'] for img in processed_images],
dim=0).to(get_local_torch_device())
with torch.no_grad():
image_inputs = {'pixel_values': pixel_values}
with set_forward_context(current_timestep=0, attn_metadata=None):
clip_features = self.get_module("image_encoder")(**image_inputs)
clip_features = clip_features.last_hidden_state
features["clip_feature"] = clip_features
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in first_frame:
processed_img = frame.to(device="cpu", dtype=torch.float32)
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
2).unsqueeze(2)
# (B, H, W, C) -> (B, C, 1, H, W)
video_condition = torch.cat([
processed_img,
processed_img.new_zeros(processed_img.shape[0],
processed_img.shape[1], num_frames - 1,
height, width)
],
dim=2)
video_condition = video_condition.to(
device=get_local_torch_device(), dtype=torch.float32)
video_conditions.append(video_condition)
video_conditions = torch.cat(video_conditions, dim=0)
with torch.autocast(device_type="cuda",
dtype=torch.float32,
enabled=True):
encoder_outputs = self.get_module("vae").encode(video_conditions)
latent_condition = encoder_outputs.mean
if (hasattr(self.get_module("vae"), "shift_factor")
and self.get_module("vae").shift_factor is not None):
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
latent_condition -= self.get_module("vae").shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.get_module("vae").shift_factor
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor.to(latent_condition.device,
latent_condition.dtype)
else:
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor
# mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
# latent_width)
# mask_lat_size[:, :, list(range(1, num_frames))] = 0
# first_frame_mask = mask_lat_size[:, :, 0:1]
# first_frame_mask = torch.repeat_interleave(
# first_frame_mask,
# dim=2,
# repeats=self.get_module("vae").temporal_compression_ratio)
# mask_lat_size = torch.concat(
# [first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
# mask_lat_size = mask_lat_size.view(
# batch_size, -1,
# self.get_module("vae").temporal_compression_ratio, latent_height,
# latent_width)
# mask_lat_size = mask_lat_size.transpose(1, 2)
# mask_lat_size = mask_lat_size.to(latent_condition.device)
# image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
features["first_frame_latent"] = latent_condition
if "action_path" in valid_data and valid_data["action_path"]:
keyboard_cond_list = []
mouse_cond_list = []
num_bits = 6
for action_path in valid_data["action_path"]:
if action_path:
action_data = np.load(action_path, allow_pickle=True)
if isinstance(
action_data,
np.ndarray) and action_data.dtype == np.dtype('O'):
action_dict = action_data.item()
if "keyboard" in action_dict:
keyboard_raw = action_dict["keyboard"]
# Convert 1D bit-flag values to 2D multi-hot encoding
if isinstance(keyboard_raw, np.ndarray):
if keyboard_raw.ndim == 1:
# [T] -> [T, num_bits]
T = len(keyboard_raw)
multi_hot = np.zeros((T, num_bits),
dtype=np.float32)
action_values = keyboard_raw.astype(int)
for bit_idx in range(num_bits):
target_idx = (
2 -
(bit_idx % 3)) + 3 * (bit_idx // 3)
if target_idx < num_bits:
multi_hot[:, target_idx] = (
(action_values >> bit_idx)
& 1).astype(np.float32)
keyboard_cond_list.append(multi_hot)
else:
# If already 2D, pad to num_bits if necessary
k_data = keyboard_raw.astype(np.float32)
if k_data.ndim == 2 and k_data.shape[
-1] < num_bits:
padding = np.zeros(
(k_data.shape[0],
num_bits - k_data.shape[-1]),
dtype=np.float32)
k_data = np.concatenate(
[k_data, padding], axis=-1)
keyboard_cond_list.append(k_data)
else:
keyboard_cond_list.append(keyboard_raw)
if "mouse" in action_dict:
mouse_cond_list.append(action_dict["mouse"])
else:
if isinstance(action_data,
np.ndarray) and action_data.ndim == 1:
T = len(action_data)
multi_hot = np.zeros((T, num_bits),
dtype=np.float32)
action_values = action_data.astype(int)
for bit_idx in range(num_bits):
target_idx = (
2 - (bit_idx % 3)) + 3 * (bit_idx // 3)
if target_idx < num_bits:
multi_hot[:, target_idx] = (
(action_values >> bit_idx) & 1).astype(
np.float32)
keyboard_cond_list.append(multi_hot)
else:
# If already 2D, pad to num_bits if necessary
k_data = action_data.astype(np.float32)
if k_data.ndim == 2 and k_data.shape[-1] < num_bits:
padding = np.zeros(
(k_data.shape[0],
num_bits - k_data.shape[-1]),
dtype=np.float32)
k_data = np.concatenate([k_data, padding],
axis=-1)
keyboard_cond_list.append(k_data)
if keyboard_cond_list:
features["keyboard_cond"] = keyboard_cond_list
if mouse_cond_list:
features["mouse_cond"] = mouse_cond_list
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "clip_feature" in extra_features:
clip_feature = extra_features["clip_feature"]
record.update({
"clip_feature_bytes": clip_feature.tobytes(),
"clip_feature_shape": list(clip_feature.shape),
"clip_feature_dtype": str(clip_feature.dtype),
})
else:
record.update({
"clip_feature_bytes": b"",
"clip_feature_shape": [],
"clip_feature_dtype": "",
})
if extra_features and "first_frame_latent" in extra_features:
first_frame_latent = extra_features["first_frame_latent"]
record.update({
"first_frame_latent_bytes":
first_frame_latent.tobytes(),
"first_frame_latent_shape":
list(first_frame_latent.shape),
"first_frame_latent_dtype":
str(first_frame_latent.dtype),
})
else:
record.update({
"first_frame_latent_bytes": b"",
"first_frame_latent_shape": [],
"first_frame_latent_dtype": "",
})
if extra_features and "pil_image" in extra_features:
pil_image = extra_features["pil_image"]
record.update({
"pil_image_bytes": pil_image.tobytes(),
"pil_image_shape": list(pil_image.shape),
"pil_image_dtype": str(pil_image.dtype),
})
else:
record.update({
"pil_image_bytes": b"",
"pil_image_shape": [],
"pil_image_dtype": "",
})
if extra_features and "keyboard_cond" in extra_features:
keyboard_cond = extra_features["keyboard_cond"]
record.update({
"keyboard_cond_bytes": keyboard_cond.tobytes(),
"keyboard_cond_shape": list(keyboard_cond.shape),
"keyboard_cond_dtype": str(keyboard_cond.dtype),
})
else:
record.update({
"keyboard_cond_bytes": b"",
"keyboard_cond_shape": [],
"keyboard_cond_dtype": "",
})
if extra_features and "mouse_cond" in extra_features:
mouse_cond = extra_features["mouse_cond"]
record.update({
"mouse_cond_bytes": mouse_cond.tobytes(),
"mouse_cond_shape": list(mouse_cond.shape),
"mouse_cond_dtype": str(mouse_cond.dtype),
})
else:
record.update({
"mouse_cond_bytes": b"",
"mouse_cond_shape": [],
"mouse_cond_dtype": "",
})
return record
EntryClass = PreprocessPipeline_MatrixGame
@@ -320,12 +320,18 @@ class BasePreprocessPipeline(ComposedPipelineBase):
"pixel_values":
torch.stack(
[data["pixel_values"][i] for i in valid_indices]),
"text": [data["text"][i] for i in valid_indices],
"text": [data["text"][i] for i in valid_indices]
if "text" in data else ["" for _ in valid_indices],
"path": [data["path"][i] for i in valid_indices],
"fps": [data["fps"][i] for i in valid_indices],
"duration": [data["duration"][i] for i in valid_indices],
}
if "action_path" in data:
valid_data["action_path"] = [
data["action_path"][i] for i in valid_indices
]
# VAE
with torch.autocast("cuda", dtype=torch.float32):
latents = self.get_module("vae").encode(
@@ -343,28 +349,35 @@ class BasePreprocessPipeline(ComposedPipelineBase):
prompt_embeds=[],
prompt_attention_mask=[],
)
assert hasattr(self, "prompt_encoding_stage")
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
prompt_embeds, prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
assert prompt_embeds.shape[0] == prompt_attention_mask.shape[0]
# Get sequence lengths from attention masks (number of 1s)
seq_lens = prompt_attention_mask.sum(dim=1)
if hasattr(self, "prompt_encoding_stage"):
result_batch = self.prompt_encoding_stage(
batch, fastvideo_args)
prompt_embeds, prompt_attention_mask = result_batch.prompt_embeds[
0], result_batch.prompt_attention_mask[0]
assert prompt_embeds.shape[
0] == prompt_attention_mask.shape[0]
non_padded_embeds = []
non_padded_masks = []
# Get sequence lengths from attention masks (number of 1s)
seq_lens = prompt_attention_mask.sum(dim=1)
# Process each item in the batch
for i in range(prompt_embeds.size(0)):
seq_len = seq_lens[i].item()
# Slice the embeddings and masks to keep only non-padding parts
non_padded_embeds.append(prompt_embeds[i, :seq_len])
non_padded_masks.append(prompt_attention_mask[i, :seq_len])
non_padded_embeds = []
non_padded_masks = []
# Update the tensors with non-padded versions
prompt_embeds = non_padded_embeds
prompt_attention_mask = non_padded_masks
# Process each item in the batch
for i in range(prompt_embeds.size(0)):
seq_len = seq_lens[i].item()
# Slice the embeddings and masks to keep only non-padding parts
non_padded_embeds.append(prompt_embeds[i, :seq_len])
non_padded_masks.append(
prompt_attention_mask[i, :seq_len])
# Update the tensors with non-padded versions
prompt_embeds = non_padded_embeds
prompt_attention_mask = non_padded_masks
else:
bs = len(valid_indices)
prompt_embeds = [torch.zeros(0) for _ in range(bs)]
# Prepare batch data for Parquet dataset
batch_data = []
@@ -16,6 +16,10 @@ from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
PreprocessPipeline_T2V)
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
PreprocessPipeline_Text)
from fastvideo.pipelines.preprocess.matrixgame.matrixgame_preprocess_pipeline import (
PreprocessPipeline_MatrixGame)
from fastvideo.pipelines.preprocess.wangame.wangame_preprocess_pipeline import (
PreprocessPipeline_WanGame)
from fastvideo.utils import maybe_download_model
logger = init_logger(__name__)
@@ -60,9 +64,14 @@ def main(args) -> None:
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
elif args.preprocess_task == "matrixgame":
PreprocessPipeline = PreprocessPipeline_MatrixGame
elif args.preprocess_task == "wangame":
PreprocessPipeline = PreprocessPipeline_WanGame
else:
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only")
raise ValueError(
f"Invalid preprocess task: {args.preprocess_task}. "
f"Valid options: t2v, i2v, ode_trajectory, text_only, matrixgame, wangame")
logger.info("Preprocess task: %s using %s", args.preprocess_task,
PreprocessPipeline.__name__)
@@ -106,11 +115,12 @@ if __name__ == "__main__":
parser.add_argument("--group_frame", action="store_true") # TODO
parser.add_argument("--group_resolution", action="store_true") # TODO
parser.add_argument("--flow_shift", type=float, default=None)
parser.add_argument("--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "text_only", "ode_trajectory"],
help="Type of preprocessing task to run")
parser.add_argument(
"--preprocess_task",
type=str,
default="t2v",
choices=["t2v", "i2v", "text_only", "ode_trajectory", "matrixgame", "wangame"],
help="Type of preprocessing task to run")
parser.add_argument("--train_fps", type=int, default=30)
parser.add_argument("--use_image_num", type=int, default=0)
parser.add_argument("--text_max_length", type=int, default=256)
@@ -0,0 +1,303 @@
# SPDX-License-Identifier: Apache-2.0
from typing import Any
import numpy as np
import torch
from PIL import Image
from fastvideo.dataset.dataloader.schema import pyarrow_schema_wangame
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
BasePreprocessPipeline)
from fastvideo.pipelines.stages import ImageEncodingStage
class PreprocessPipeline_WanGame(BasePreprocessPipeline):
"""I2V preprocessing pipeline implementation."""
_required_config_modules = ["vae", "image_encoder", "image_processor"]
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
self.add_stage(stage_name="image_encoding_stage",
stage=ImageEncodingStage(
image_encoder=self.get_module("image_encoder"),
image_processor=self.get_module("image_processor"),
))
def get_pyarrow_schema(self):
"""Return the PyArrow schema for I2V pipeline."""
return pyarrow_schema_wangame
def get_extra_features(self, valid_data: dict[str, Any],
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
# TODO(will): move these to cpu at some point
self.get_module("image_encoder").to(get_local_torch_device())
self.get_module("vae").to(get_local_torch_device())
features = {}
"""Get CLIP features from the first frame of each video."""
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
_, _, num_frames, height, width = valid_data["pixel_values"].shape
# latent_height = height // self.get_module(
# "vae").spatial_compression_ratio
# latent_width = width // self.get_module("vae").spatial_compression_ratio
processed_images = []
# Frame has values between -1 and 1
for frame in first_frame:
frame = (frame + 1) * 127.5
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
processed_img = self.get_module("image_processor")(
images=frame_pil, return_tensors="pt")
processed_images.append(processed_img)
# Get CLIP features
pixel_values = torch.cat(
[img['pixel_values'] for img in processed_images],
dim=0).to(get_local_torch_device())
with torch.no_grad():
image_inputs = {'pixel_values': pixel_values}
with set_forward_context(current_timestep=0, attn_metadata=None):
clip_features = self.get_module("image_encoder")(**image_inputs)
clip_features = clip_features.last_hidden_state
features["clip_feature"] = clip_features
"""Get VAE features from the first frame of each video"""
video_conditions = []
for frame in first_frame:
processed_img = frame.to(device="cpu", dtype=torch.float32)
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
2).unsqueeze(2)
# (B, H, W, C) -> (B, C, 1, H, W)
video_condition = torch.cat([
processed_img,
processed_img.new_zeros(processed_img.shape[0],
processed_img.shape[1], num_frames - 1,
height, width)
],
dim=2)
video_condition = video_condition.to(
device=get_local_torch_device(), dtype=torch.float32)
video_conditions.append(video_condition)
video_conditions = torch.cat(video_conditions, dim=0)
with torch.autocast(device_type="cuda",
dtype=torch.float32,
enabled=True):
encoder_outputs = self.get_module("vae").encode(video_conditions)
latent_condition = encoder_outputs.mean
if (hasattr(self.get_module("vae"), "shift_factor")
and self.get_module("vae").shift_factor is not None):
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
latent_condition -= self.get_module("vae").shift_factor.to(
latent_condition.device, latent_condition.dtype)
else:
latent_condition -= self.get_module("vae").shift_factor
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor.to(latent_condition.device,
latent_condition.dtype)
else:
latent_condition = latent_condition * self.get_module(
"vae").scaling_factor
# mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
# latent_width)
# mask_lat_size[:, :, list(range(1, num_frames))] = 0
# first_frame_mask = mask_lat_size[:, :, 0:1]
# first_frame_mask = torch.repeat_interleave(
# first_frame_mask,
# dim=2,
# repeats=self.get_module("vae").temporal_compression_ratio)
# mask_lat_size = torch.concat(
# [first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
# mask_lat_size = mask_lat_size.view(
# batch_size, -1,
# self.get_module("vae").temporal_compression_ratio, latent_height,
# latent_width)
# mask_lat_size = mask_lat_size.transpose(1, 2)
# mask_lat_size = mask_lat_size.to(latent_condition.device)
# image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
features["first_frame_latent"] = latent_condition
if "action_path" in valid_data and valid_data["action_path"]:
keyboard_cond_list = []
mouse_cond_list = []
num_bits = 6
for action_path in valid_data["action_path"]:
if action_path:
action_data = np.load(action_path, allow_pickle=True)
if isinstance(
action_data,
np.ndarray) and action_data.dtype == np.dtype('O'):
action_dict = action_data.item()
if "keyboard" in action_dict:
keyboard_raw = action_dict["keyboard"]
# Convert 1D bit-flag values to 2D multi-hot encoding
if isinstance(keyboard_raw, np.ndarray):
if keyboard_raw.ndim == 1:
# [T] -> [T, num_bits]
T = len(keyboard_raw)
multi_hot = np.zeros((T, num_bits),
dtype=np.float32)
action_values = keyboard_raw.astype(int)
for bit_idx in range(num_bits):
target_idx = (
2 -
(bit_idx % 3)) + 3 * (bit_idx // 3)
if target_idx < num_bits:
multi_hot[:, target_idx] = (
(action_values >> bit_idx)
& 1).astype(np.float32)
keyboard_cond_list.append(multi_hot)
else:
# If already 2D, pad to num_bits if necessary
k_data = keyboard_raw.astype(np.float32)
if k_data.ndim == 2 and k_data.shape[
-1] < num_bits:
padding = np.zeros(
(k_data.shape[0],
num_bits - k_data.shape[-1]),
dtype=np.float32)
k_data = np.concatenate(
[k_data, padding], axis=-1)
keyboard_cond_list.append(k_data)
else:
keyboard_cond_list.append(keyboard_raw)
if "mouse" in action_dict:
mouse_cond_list.append(action_dict["mouse"])
else:
if isinstance(action_data,
np.ndarray) and action_data.ndim == 1:
T = len(action_data)
multi_hot = np.zeros((T, num_bits),
dtype=np.float32)
action_values = action_data.astype(int)
for bit_idx in range(num_bits):
target_idx = (
2 - (bit_idx % 3)) + 3 * (bit_idx // 3)
if target_idx < num_bits:
multi_hot[:, target_idx] = (
(action_values >> bit_idx) & 1).astype(
np.float32)
keyboard_cond_list.append(multi_hot)
else:
# If already 2D, pad to num_bits if necessary
k_data = action_data.astype(np.float32)
if k_data.ndim == 2 and k_data.shape[-1] < num_bits:
padding = np.zeros(
(k_data.shape[0],
num_bits - k_data.shape[-1]),
dtype=np.float32)
k_data = np.concatenate([k_data, padding],
axis=-1)
keyboard_cond_list.append(k_data)
if keyboard_cond_list:
features["keyboard_cond"] = keyboard_cond_list
if mouse_cond_list:
features["mouse_cond"] = mouse_cond_list
return features
def create_record(
self,
video_name: str,
vae_latent: np.ndarray,
text_embedding: np.ndarray,
valid_data: dict[str, Any],
idx: int,
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
"""Create a record for the Parquet dataset with CLIP features."""
record = super().create_record(video_name=video_name,
vae_latent=vae_latent,
text_embedding=text_embedding,
valid_data=valid_data,
idx=idx,
extra_features=extra_features)
if extra_features and "clip_feature" in extra_features:
clip_feature = extra_features["clip_feature"]
record.update({
"clip_feature_bytes": clip_feature.tobytes(),
"clip_feature_shape": list(clip_feature.shape),
"clip_feature_dtype": str(clip_feature.dtype),
})
else:
record.update({
"clip_feature_bytes": b"",
"clip_feature_shape": [],
"clip_feature_dtype": "",
})
if extra_features and "first_frame_latent" in extra_features:
first_frame_latent = extra_features["first_frame_latent"]
record.update({
"first_frame_latent_bytes":
first_frame_latent.tobytes(),
"first_frame_latent_shape":
list(first_frame_latent.shape),
"first_frame_latent_dtype":
str(first_frame_latent.dtype),
})
else:
record.update({
"first_frame_latent_bytes": b"",
"first_frame_latent_shape": [],
"first_frame_latent_dtype": "",
})
if extra_features and "pil_image" in extra_features:
pil_image = extra_features["pil_image"]
record.update({
"pil_image_bytes": pil_image.tobytes(),
"pil_image_shape": list(pil_image.shape),
"pil_image_dtype": str(pil_image.dtype),
})
else:
record.update({
"pil_image_bytes": b"",
"pil_image_shape": [],
"pil_image_dtype": "",
})
if extra_features and "keyboard_cond" in extra_features:
keyboard_cond = extra_features["keyboard_cond"]
record.update({
"keyboard_cond_bytes": keyboard_cond.tobytes(),
"keyboard_cond_shape": list(keyboard_cond.shape),
"keyboard_cond_dtype": str(keyboard_cond.dtype),
})
else:
record.update({
"keyboard_cond_bytes": b"",
"keyboard_cond_shape": [],
"keyboard_cond_dtype": "",
})
if extra_features and "mouse_cond" in extra_features:
mouse_cond = extra_features["mouse_cond"]
record.update({
"mouse_cond_bytes": mouse_cond.tobytes(),
"mouse_cond_shape": list(mouse_cond.shape),
"mouse_cond_dtype": str(mouse_cond.dtype),
})
else:
record.update({
"mouse_cond_bytes": b"",
"mouse_cond_shape": [],
"mouse_cond_dtype": "",
})
return record
EntryClass = PreprocessPipeline_WanGame
+16
View File
@@ -168,6 +168,20 @@ class DenoisingStage(PipelineStage):
},
)
if batch.mouse_cond is not None and batch.keyboard_cond is not None:
from fastvideo.models.dits.hyworld.pose import process_custom_actions
viewmats, intrinsics, action_labels = process_custom_actions(batch.keyboard_cond, batch.mouse_cond)
camera_action_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
"viewmats": viewmats.unsqueeze(0).to(get_local_torch_device(), dtype=target_dtype),
"Ks": intrinsics.unsqueeze(0).to(get_local_torch_device(), dtype=target_dtype),
"action": action_labels.unsqueeze(0).to(get_local_torch_device(), dtype=target_dtype),
},
)
else:
camera_action_kwargs = {}
action_kwargs = self.prepare_extra_func_kwargs(
self.transformer.forward,
{
@@ -406,6 +420,7 @@ class DenoisingStage(PipelineStage):
**image_kwargs,
**pos_cond_kwargs,
**action_kwargs,
**camera_action_kwargs,
)
if batch.do_classifier_free_guidance:
@@ -423,6 +438,7 @@ class DenoisingStage(PipelineStage):
**image_kwargs,
**neg_cond_kwargs,
**action_kwargs,
**camera_action_kwargs,
)
noise_pred_text = noise_pred
@@ -210,9 +210,9 @@ class InputValidationStage(PipelineStage):
f"keyboard_cond must have 3 dimensions (B, T, K), but got {batch.keyboard_cond.dim()}"
)
keyboard_dim = batch.keyboard_cond.shape[-1]
if keyboard_dim not in {2, 4, 6, 7}:
if keyboard_dim not in {2, 3, 4, 6, 7}:
raise ValueError(
f"keyboard_cond last dimension must be 2, 4, 6, or 7, but got {keyboard_dim}"
f"keyboard_cond last dimension must be 2, 3, 4, 6, or 7, but got {keyboard_dim}"
)
logger.info(
"Action control: keyboard_cond validated - shape %s (dim=%d)",
+1 -1
View File
@@ -155,7 +155,7 @@ class CudaPlatformBase(Platform):
)
elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE:
try:
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell # noqa: F401
from sageattn3 import sageattn3_blackwell # noqa: F401
from fastvideo.attention.backends.sage_attn3 import ( # noqa: F401
SageAttention3Backend)
@@ -0,0 +1,138 @@
# SPDX-License-Identifier: Apache-2.0
"""
Mismatch compatibility test (direction B):
- WORLD_SIZE=4
- SP: sp_size=2
- HSDP/FSDP2 mesh: replicate_dim=1, shard_dim=4 (mesh_shape=(1,4))
We verify gradients from:
(A) full MSE: mean((pred-target)^2)
(B) SP-sharded loss: sp_size * local_sse / total_numel
match under this mismatch.
Run:
source /home/hao_lab/miniconda3/etc/profile.d/conda.sh
conda activate alexfv
torchrun --standalone --nproc_per_node=4 -m pytest -q fastvideo/tests/distributed/test_hsdp_sp_mismatch_sp2_hsdp_shard4.py
"""
import os
import pytest
import torch
import torch.distributed as dist
from fastvideo.distributed import cleanup_dist_env_and_memory
from fastvideo.distributed.parallel_state import (
get_sp_group,
maybe_init_distributed_environment_and_model_parallel,
)
from fastvideo.training.training_utils import shard_latents_across_sp
def _world_size() -> int:
return int(os.environ.get("WORLD_SIZE", "1"))
def _rank() -> int:
return int(os.environ.get("RANK", "0"))
@pytest.fixture(scope="module")
def dist_setup():
if _world_size() != 4:
pytest.skip("Designed for torchrun WORLD_SIZE=4")
if not torch.cuda.is_available():
pytest.skip("Requires CUDA")
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
torch.cuda.set_device(local_rank)
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=2)
assert get_sp_group().world_size == 2
yield
if dist.is_available() and dist.is_initialized():
dist.barrier()
cleanup_dist_env_and_memory()
class ScaleModule(torch.nn.Module):
def __init__(self, c: int, init: torch.Tensor):
super().__init__()
assert init.shape == (c,)
self.scale = torch.nn.Parameter(init.clone())
def forward(self, x: torch.Tensor) -> torch.Tensor:
return x * self.scale.view(1, -1, 1, 1, 1)
def _broadcast_(t: torch.Tensor, src: int = 0) -> None:
dist.broadcast(t, src=src)
def _gather_full_vec_allranks(local_shard: torch.Tensor) -> torch.Tensor:
gathered = [torch.empty_like(local_shard) for _ in range(dist.get_world_size())]
dist.all_gather(gathered, local_shard)
return torch.cat(gathered, dim=0)
def test_sp_sharded_loss_matches_full_mse_under_mismatch(dist_setup):
from torch.distributed.device_mesh import init_device_mesh
from torch.distributed.fsdp import MixedPrecisionPolicy, fully_shard
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
torch.manual_seed(2027)
b, c, t, h, w = 2, 8, 1, 1, 8 # thw=8 divisible by sp=2
x = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
if _rank() == 0:
x.normal_()
_broadcast_(x)
target = (0.25 * x).contiguous()
init = torch.empty((c,), device=device, dtype=torch.float32)
if _rank() == 0:
init.normal_()
_broadcast_(init)
mesh = init_device_mesh(
"cuda",
mesh_shape=(1, 4),
mesh_dim_names=("replicate", "shard"),
)
mp = MixedPrecisionPolicy(
param_dtype=torch.float32,
reduce_dtype=torch.float32,
output_dtype=torch.float32,
cast_forward_inputs=False,
)
# (A) Full loss.
m_full = ScaleModule(c=c, init=init).to(device)
fully_shard(m_full, mesh=mesh, reshard_after_forward=True, mp_policy=mp)
pred_full = m_full(x)
loss_full = ((pred_full - target) ** 2).mean()
loss_full.backward()
g_full = m_full.scale.grad
g_full_local = g_full.to_local() if hasattr(g_full, "to_local") else g_full
g_full_vec = _gather_full_vec_allranks(g_full_local.flatten())
# (B) SP-sharded loss.
m_sp = ScaleModule(c=c, init=init).to(device)
fully_shard(m_sp, mesh=mesh, reshard_after_forward=True, mp_policy=mp)
pred = m_sp(x)
sp = get_sp_group().world_size # 2
sharded_pred = shard_latents_across_sp(pred)
sharded_target = shard_latents_across_sp(target)
local_sse = ((sharded_pred - sharded_target) ** 2).sum()
loss_sp = (sp * local_sse) / pred.numel()
loss_sp.backward()
g_sp = m_sp.scale.grad
g_sp_local = g_sp.to_local() if hasattr(g_sp, "to_local") else g_sp
g_sp_vec = _gather_full_vec_allranks(g_sp_local.flatten())
torch.testing.assert_close(g_sp_vec, g_full_vec, rtol=1e-4, atol=1e-5)
@@ -0,0 +1,152 @@
# SPDX-License-Identifier: Apache-2.0
"""
Mismatch compatibility test (direction A):
- WORLD_SIZE=4
- SP: sp_size=4
- HSDP/FSDP2 mesh: replicate_dim=2, shard_dim=2 (mesh_shape=(2,2))
We verify gradients from:
(A) full MSE: mean((pred-target)^2)
(B) SP-sharded loss: sp_size * local_sse / total_numel
match under this mismatch.
Run:
source /home/hao_lab/miniconda3/etc/profile.d/conda.sh
conda activate alexfv
torchrun --standalone --nproc_per_node=4 -m pytest -q fastvideo/tests/distributed/test_hsdp_sp_mismatch_sp4_hsdp_shard2.py
"""
import os
import pytest
import torch
import torch.distributed as dist
from fastvideo.distributed import cleanup_dist_env_and_memory
from fastvideo.distributed.parallel_state import (
get_sp_group,
maybe_init_distributed_environment_and_model_parallel,
)
from fastvideo.training.training_utils import shard_latents_across_sp
def _world_size() -> int:
return int(os.environ.get("WORLD_SIZE", "1"))
def _rank() -> int:
return int(os.environ.get("RANK", "0"))
@pytest.fixture(scope="module")
def dist_setup():
if _world_size() != 4:
pytest.skip("Designed for torchrun WORLD_SIZE=4")
if not torch.cuda.is_available():
pytest.skip("Requires CUDA")
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
torch.cuda.set_device(local_rank)
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=4)
assert get_sp_group().world_size == 4
yield
if dist.is_available() and dist.is_initialized():
dist.barrier()
cleanup_dist_env_and_memory()
class ScaleModule(torch.nn.Module):
def __init__(self, c: int, init: torch.Tensor):
super().__init__()
assert init.shape == (c,)
self.scale = torch.nn.Parameter(init.clone())
def forward(self, x: torch.Tensor) -> torch.Tensor:
return x * self.scale.view(1, -1, 1, 1, 1)
def _broadcast_(t: torch.Tensor, src: int = 0) -> None:
dist.broadcast(t, src=src)
def _reconstruct_full_vec_from_replicate0(local_shard: torch.Tensor) -> torch.Tensor:
"""
Mesh layout: [[0,2],[1,3]] with dim_names=("replicate","shard").
replicate0 group is ranks [0,2], and shard order inside replicate0 is [0,2].
Reconstruct a single full vector from replicate0 only, then broadcast to all ranks.
"""
group_ranks = [0, 2]
pg = dist.new_group(ranks=group_ranks)
if dist.get_rank() in group_ranks:
gathered = [torch.empty_like(local_shard) for _ in range(len(group_ranks))]
dist.all_gather(gathered, local_shard, group=pg)
full = torch.cat(gathered, dim=0)
else:
full = torch.empty((local_shard.numel() * len(group_ranks),),
device=local_shard.device,
dtype=local_shard.dtype)
dist.broadcast(full, src=0)
return full
def test_sp_sharded_loss_matches_full_mse_under_mismatch(dist_setup):
from torch.distributed.device_mesh import DeviceMesh
from torch.distributed.fsdp import MixedPrecisionPolicy, fully_shard
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
torch.manual_seed(2026)
b, c, t, h, w = 2, 8, 1, 1, 8 # thw=8 divisible by sp=4
x = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
if _rank() == 0:
x.normal_()
_broadcast_(x)
target = (0.5 * x).contiguous()
init = torch.empty((c,), device=device, dtype=torch.float32)
if _rank() == 0:
init.normal_()
_broadcast_(init)
mesh = DeviceMesh(
"cuda",
[[0, 2], [1, 3]],
mesh_dim_names=("replicate", "shard"),
)
mp = MixedPrecisionPolicy(
param_dtype=torch.float32,
reduce_dtype=torch.float32,
output_dtype=torch.float32,
cast_forward_inputs=False,
)
# (A) Full loss.
m_full = ScaleModule(c=c, init=init).to(device)
fully_shard(m_full, mesh=mesh, reshard_after_forward=True, mp_policy=mp)
pred_full = m_full(x)
loss_full = ((pred_full - target) ** 2).mean()
loss_full.backward()
g_full = m_full.scale.grad
g_full_local = g_full.to_local() if hasattr(g_full, "to_local") else g_full
g_full_vec = _reconstruct_full_vec_from_replicate0(g_full_local.flatten())
# (B) SP-sharded loss.
m_sp = ScaleModule(c=c, init=init).to(device)
fully_shard(m_sp, mesh=mesh, reshard_after_forward=True, mp_policy=mp)
pred = m_sp(x)
sp = get_sp_group().world_size # 4
sharded_pred = shard_latents_across_sp(pred)
sharded_target = shard_latents_across_sp(target)
local_sse = ((sharded_pred - sharded_target) ** 2).sum()
loss_sp = (sp * local_sse) / pred.numel()
loss_sp.backward()
g_sp = m_sp.scale.grad
g_sp_local = g_sp.to_local() if hasattr(g_sp, "to_local") else g_sp
g_sp_vec = _reconstruct_full_vec_from_replicate0(g_sp_local.flatten())
torch.testing.assert_close(g_sp_vec, g_full_vec, rtol=1e-4, atol=1e-5)
@@ -0,0 +1,218 @@
# SPDX-License-Identifier: Apache-2.0
import os
import pytest
import torch
import torch.distributed as dist
from fastvideo.distributed import (cleanup_dist_env_and_memory,
get_sp_group,
maybe_init_distributed_environment_and_model_parallel)
from fastvideo.training.training_utils import shard_latents_across_sp
def _world_size() -> int:
return int(os.environ.get("WORLD_SIZE", "1"))
def _rank() -> int:
return int(os.environ.get("RANK", "0"))
@pytest.fixture(scope="module")
def sp_dist():
"""
These tests are intended to be run under torchrun with multiple GPUs, e.g.
torchrun --nproc_per_node=2 -m pytest -q fastvideo/tests/distributed/test_sp_shard_latents_and_loss.py
torchrun --nproc_per_node=4 -m pytest -q fastvideo/tests/distributed/test_sp_shard_latents_and_loss.py
torchrun --nproc_per_node=8 -m pytest -q fastvideo/tests/distributed/test_sp_shard_latents_and_loss.py
"""
ws = _world_size()
if ws <= 1:
pytest.skip("Requires torchrun with WORLD_SIZE>1")
if not torch.cuda.is_available():
pytest.skip("Requires CUDA")
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
torch.cuda.set_device(local_rank)
torch.manual_seed(1234)
# Initialize FastVideo dist + model-parallel groups.
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=ws)
assert get_sp_group().world_size == ws
yield
if dist.is_available() and dist.is_initialized():
dist.barrier()
cleanup_dist_env_and_memory()
def _broadcast_tensor_(t: torch.Tensor, src: int = 0) -> torch.Tensor:
if dist.is_available() and dist.is_initialized():
dist.broadcast(t, src=src)
return t
def _all_gather_cat(t: torch.Tensor, dim: int) -> torch.Tensor:
ws = _world_size()
gathered = [torch.empty_like(t) for _ in range(ws)]
dist.all_gather(gathered, t)
return torch.cat(gathered, dim=dim)
def test_shard_latents_t_not_divisible_but_thw_divisible_no_error(sp_dist):
"""
Regression test for the commit "fix splitting on t":
- t is NOT divisible by sp_size
- (t*h*w) IS divisible by sp_size
- sharding should NOT raise, and shards should round-trip to the original.
"""
ws = _world_size()
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
b, c, t, h, w = 2, 3, 3, 2, 4 # t=3 (not divisible by 8), thw=24 (divisible by 8)
assert (t * h * w) % ws == 0
assert t % ws != 0
latents = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
if _rank() == 0:
latents.normal_()
_broadcast_tensor_(latents)
shard = shard_latents_across_sp(latents)
assert shard.shape == (b, c, (t * h * w) // ws)
gathered = _all_gather_cat(shard, dim=2)
torch.testing.assert_close(gathered, latents.reshape(b, c, t * h * w))
dist.barrier()
@pytest.mark.parametrize(
"shape",
[
# No padding needed: thw divisible by 8, but t is not.
(2, 1, 3, 2, 4), # thw=24
# Padding needed: thw not divisible by 8.
(2, 1, 3, 2, 5), # thw=30 -> pad to 32
],
)
def test_sharded_loss_gradient_matches_full_mse_after_rank_avg(sp_dist, shape):
"""
Validates the math behind the new loss computation:
If each rank computes:
loss_rank = sp_world_size * local_sse / total_numel
and the training stack averages gradients across ranks, then the resulting
gradient should match the full MSE gradient.
"""
ws = _world_size()
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
b, c, t, h, w = shape
# Make inputs identical across ranks for deterministic comparison.
init = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
x = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
y = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
if _rank() == 0:
init.normal_()
x.normal_()
y.normal_()
_broadcast_tensor_(init)
_broadcast_tensor_(x)
_broadcast_tensor_(y)
# Full loss baseline (redundant compute; correct reference).
w_full = torch.nn.Parameter(init.clone())
pred_full = w_full * x
full_loss = ((pred_full - y)**2).mean()
full_loss.backward()
grad_full = w_full.grad.detach().clone()
# Sharded loss (compute on a shard only).
w_shard = torch.nn.Parameter(init.clone())
pred = w_shard * x
sharded_pred = shard_latents_across_sp(pred)
sharded_target = shard_latents_across_sp(y)
local_sse = ((sharded_pred - sharded_target)**2).sum()
shard_loss = ws * local_sse / pred.numel()
shard_loss.backward()
grad_shard = w_shard.grad.detach().clone()
# Simulate "gradient averaging across ranks" (DDP/FSDP replicate-dim behavior).
dist.all_reduce(grad_shard, op=dist.ReduceOp.SUM)
grad_shard /= ws
# (Optional) also average grad_full, for symmetry.
dist.all_reduce(grad_full, op=dist.ReduceOp.SUM)
grad_full /= ws
torch.testing.assert_close(grad_shard, grad_full, rtol=1e-5, atol=1e-6)
dist.barrier()
def test_padding_does_not_change_global_sse(sp_dist):
"""
Padding-specific correctness test.
When (t*h*w) is NOT divisible by sp_size, shard_latents_across_sp pads with
zeros on the flattened axis. This test validates that:
1) Summed SSE across SP ranks equals the full (unpadded) SSE.
2) Any padded tokens that land on a rank are exactly zero for both pred/target.
"""
from fastvideo.distributed.utils import compute_padding_for_sp
ws = _world_size()
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
b, c, t, h = 2, 2, 1, 1
# Force padding for any ws>1: seq_len = 2*ws + 1 -> remainder 1
w = 2 * ws + 1
seq_len = t * h * w
assert seq_len % ws != 0
pred = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
target = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
if _rank() == 0:
pred.normal_()
target.normal_()
_broadcast_tensor_(pred)
_broadcast_tensor_(target)
# Full (unpadded) SSE reference.
full_sse = ((pred - target)**2).sum()
# Sharded local SSE (includes padded region, which should contribute 0).
sharded_pred = shard_latents_across_sp(pred)
sharded_target = shard_latents_across_sp(target)
local_sse = ((sharded_pred - sharded_target)**2).sum()
global_sse = local_sse.clone()
dist.all_reduce(global_sse, op=dist.ReduceOp.SUM)
torch.testing.assert_close(global_sse, full_sse, rtol=1e-5, atol=1e-6)
# Explicitly validate padded tokens are zeros on ranks that cover them.
padded_seq_len, padding_amount = compute_padding_for_sp(seq_len, ws)
assert padding_amount > 0
elements_per_rank = padded_seq_len // ws
start = _rank() * elements_per_rank
end = (_rank() + 1) * elements_per_rank
if end > seq_len:
# There is padding on this rank: positions [max(seq_len, start), end)
pad_start_global = max(seq_len, start)
pad_end_global = end
pad_start_local = pad_start_global - start
pad_end_local = pad_end_global - start
assert pad_start_local < pad_end_local
pad_slice_pred = sharded_pred[:, :, pad_start_local:pad_end_local]
pad_slice_target = sharded_target[:, :, pad_start_local:pad_end_local]
torch.testing.assert_close(pad_slice_pred, torch.zeros_like(pad_slice_pred))
torch.testing.assert_close(pad_slice_target, torch.zeros_like(pad_slice_target))
dist.barrier()
@@ -5,7 +5,7 @@ import torch
import pytest
from fastvideo import VideoGenerator
from fastvideo.models.dits.matrix_game.utils import create_action_presets
from fastvideo.models.dits.matrixgame.utils import create_action_presets
from fastvideo.logger import init_logger
from fastvideo.tests.utils import (
compute_video_ssim_torchvision,
@@ -0,0 +1,250 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any
import torch
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset.dataloader.schema import pyarrow_schema_matrixgame
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.pipelines.basic.matrixgame.matrixgame_i2v_pipeline import (
MatrixGamePipeline)
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class MatrixGameTrainingPipeline(TrainingPipeline):
"""
A training pipeline for Matrix-Game-2.0.
"""
_required_config_modules = ["scheduler", "transformer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_matrixgame
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
args_copy = deepcopy(training_args)
args_copy.inference_mode = True
args_copy.dit_cpu_offload = True
# args_copy.pipeline_config.vae_config.load_encoder = False
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
self.validation_pipeline = MatrixGamePipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
dit_cpu_offload=True)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
# encoder_hidden_states = batch['text_embedding']
# encoder_attention_mask = batch['text_attention_mask']
clip_features = batch['clip_feature']
image_latents = batch['first_frame_latent']
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
pil_image = batch['pil_image']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = None
training_batch.encoder_attention_mask = None
# MatrixGame doesn't use text encoder
training_batch.preprocessed_image = pil_image.to(
get_local_torch_device())
training_batch.image_embeds = clip_features.to(get_local_torch_device())
training_batch.image_latents = image_latents.to(
get_local_torch_device())
training_batch.infos = infos
# Action conditioning
if 'mouse_cond' in batch and batch['mouse_cond'].numel() > 0:
training_batch.mouse_cond = batch['mouse_cond'].to(
get_local_torch_device(), dtype=torch.bfloat16)
else:
training_batch.mouse_cond = None
if 'keyboard_cond' in batch and batch['keyboard_cond'].numel() > 0:
training_batch.keyboard_cond = batch['keyboard_cond'].to(
get_local_torch_device(), dtype=torch.bfloat16)
else:
training_batch.keyboard_cond = None
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
# First, call parent method to prepare noise, timesteps, etc. for video latents
training_batch = super()._prepare_dit_inputs(training_batch)
assert isinstance(training_batch.image_latents, torch.Tensor)
image_latents = training_batch.image_latents.to(
get_local_torch_device(), dtype=torch.bfloat16)
temporal_compression_ratio = self.training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (self.training_args.num_latent_t -
1) * temporal_compression_ratio + 1
batch_size, num_channels, _, latent_height, latent_width = image_latents.shape
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
latent_width)
mask_lat_size[:, :, 1:] = 0
first_frame_mask = mask_lat_size[:, :, :1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask, dim=2, repeats=temporal_compression_ratio)
mask_lat_size = torch.cat([first_frame_mask, mask_lat_size[:, :, 1:]],
dim=2)
mask_lat_size = mask_lat_size.view(batch_size, -1,
temporal_compression_ratio,
latent_height, latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(
image_latents.device).to(dtype=torch.bfloat16)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, mask_lat_size, image_latents],
dim=1)
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
# Image Embeds for conditioning
image_embeds = training_batch.image_embeds
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_local_torch_device(),
dtype=torch.bfloat16)
encoder_hidden_states_image = image_embeds
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states, # None for MatrixGame
"timestep":
training_batch.timesteps.to(get_local_torch_device(),
dtype=torch.bfloat16),
# "encoder_attention_mask":
# training_batch.encoder_attention_mask,
"encoder_hidden_states_image":
encoder_hidden_states_image,
# Action conditioning
"mouse_cond":
training_batch.mouse_cond,
"keyboard_cond":
training_batch.keyboard_cond,
"return_dict":
False,
}
return training_batch
def _prepare_validation_batch(self, sampling_param: SamplingParam,
training_args: TrainingArgs,
validation_batch: dict[str, Any],
num_inference_steps: int) -> ForwardBatch:
sampling_param.prompt = validation_batch['prompt']
sampling_param.height = training_args.num_height
sampling_param.width = training_args.num_width
sampling_param.image_path = validation_batch.get(
'image_path') or validation_batch.get('video_path')
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
assert self.seed is not None
sampling_param.seed = self.seed
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
sampling_param.num_frames = num_frames
batch = ForwardBatch(
**shallow_asdict(sampling_param),
latents=None,
generator=torch.Generator(device="cpu").manual_seed(self.seed),
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
if "image" in validation_batch and validation_batch["image"] is not None:
batch.pil_image = validation_batch["image"]
if "keyboard_cond" in validation_batch and validation_batch[
"keyboard_cond"] is not None:
keyboard_cond = validation_batch["keyboard_cond"]
keyboard_cond = torch.tensor(keyboard_cond, dtype=torch.bfloat16)
keyboard_cond = keyboard_cond.unsqueeze(0)
batch.keyboard_cond = keyboard_cond
if "mouse_cond" in validation_batch and validation_batch[
"mouse_cond"] is not None:
mouse_cond = validation_batch["mouse_cond"]
mouse_cond = torch.tensor(mouse_cond, dtype=torch.bfloat16)
mouse_cond = mouse_cond.unsqueeze(0)
batch.mouse_cond = mouse_cond
return batch
def main(args) -> None:
logger.info("Starting training pipeline...")
pipeline = MatrixGameTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
main(args)
+25 -6
View File
@@ -443,11 +443,26 @@ class TrainingPipeline(LoRAPipeline, ABC):
# make sure no implicit broadcasting happens
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
sharded_pred = shard_latents_across_sp(model_pred)
sharded_target = shard_latents_across_sp(target)
loss = (torch.mean(
(sharded_pred.float() - sharded_target.float())**2) /
self.training_args.gradient_accumulation_steps)
# Defensive: avoid NaNs/div0 if an upstream bug ever produces empty tensors.
# mean(empty) -> NaN; division by 0 -> inf/NaN. Keep a 0 loss with grad history.
if model_pred.numel() == 0:
loss = (model_pred.sum() *
0.0) / self.training_args.gradient_accumulation_steps
else:
# Compute MSE on one SP shard. We shard along the flattened token axis
# (t*h*w) with optional padding so (t*h*w) need not be divisible by sp_size.
sp_world_size = get_sp_group().world_size
if sp_world_size > 1:
sharded_pred = shard_latents_across_sp(model_pred)
sharded_target = shard_latents_across_sp(target)
local_sse = ((sharded_pred.float() -
sharded_target.float())**2).sum()
loss = (sp_world_size * local_sse / model_pred.numel()
) / self.training_args.gradient_accumulation_steps
else:
loss = (torch.mean(
(model_pred.float() - target.float())**2) /
self.training_args.gradient_accumulation_steps)
loss.backward()
@@ -654,7 +669,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] //
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
context_len = int(training_batch.encoder_hidden_states.shape[1])
if training_batch.encoder_hidden_states is not None:
context_len = int(
training_batch.encoder_hidden_states.shape[1])
else:
context_len = 0
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
+21 -5
View File
@@ -20,9 +20,10 @@ from fastvideo.training.checkpointing_utils import (ModelWrapper,
RandomStateWrapper,
SchedulerWrapper)
from einops import rearrange
from fastvideo.distributed.parallel_state import (get_sp_parallel_rank,
get_sp_world_size)
from fastvideo.distributed.utils import (compute_padding_for_sp,
pad_sequence_tensor)
logger = init_logger(__name__)
@@ -867,10 +868,25 @@ def shard_latents_across_sp(latents: torch.Tensor) -> torch.Tensor:
sp_world_size = get_sp_world_size()
rank_in_sp_group = get_sp_parallel_rank()
if sp_world_size > 1:
latents = rearrange(latents,
"b c (n s) h w -> b c n s h w",
n=sp_world_size).contiguous()
latents = latents[:, :, rank_in_sp_group, :, :, :]
# Shard on the flattened token axis (t*h*w) rather than raw `t`, so we
# don't require t % sp_world_size == 0. Pad on the flattened axis if needed.
assert latents.ndim == 5, f"Expected latents [b,c,t,h,w], got {latents.shape}"
b, c, t, h, w = latents.shape
latents = latents.reshape(b, c, t * h * w)
original_seq_len = latents.shape[2]
padded_seq_len, padding_amount = compute_padding_for_sp(
original_seq_len, sp_world_size)
if padding_amount > 0:
latents = pad_sequence_tensor(latents,
padded_seq_len,
seq_dim=2,
pad_value=0.0)
elements_per_rank = padded_seq_len // sp_world_size
start = rank_in_sp_group * elements_per_rank
end = (rank_in_sp_group + 1) * elements_per_rank
latents = latents[:, :, start:end].contiguous()
return latents
@@ -0,0 +1,250 @@
# SPDX-License-Identifier: Apache-2.0
import sys
from copy import deepcopy
from typing import Any
import torch
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset.dataloader.schema import pyarrow_schema_wangame
from fastvideo.distributed import get_local_torch_device
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
from fastvideo.logger import init_logger
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
FlowUniPCMultistepScheduler)
from fastvideo.pipelines.basic.wan.wangame_i2v_pipeline import WanGameActionImageToVideoPipeline
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
vsa_available = is_vsa_available()
logger = init_logger(__name__)
class WanGameTrainingPipeline(TrainingPipeline):
"""
A training pipeline for WanGame-2.1-Fun-1.3B-InP.
"""
_required_config_modules = ["scheduler", "transformer", "vae"]
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
shift=fastvideo_args.pipeline_config.flow_shift)
def create_training_stages(self, training_args: TrainingArgs):
"""
May be used in future refactors.
"""
pass
def set_schemas(self):
self.train_dataset_schema = pyarrow_schema_wangame
def initialize_validation_pipeline(self, training_args: TrainingArgs):
logger.info("Initializing validation pipeline...")
# args_copy.pipeline_config.vae_config.load_encoder = False
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
self.validation_pipeline = WanGameActionImageToVideoPipeline.from_pretrained(
training_args.model_path,
args=None,
inference_mode=True,
loaded_modules={
"transformer": self.get_module("transformer"),
},
tp_size=training_args.tp_size,
sp_size=training_args.sp_size,
num_gpus=training_args.num_gpus,
dit_cpu_offload=False)
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
batch = next(self.train_loader_iter, None) # type: ignore
if batch is None:
self.current_epoch += 1
logger.info("Starting epoch %s", self.current_epoch)
# Reset iterator for next epoch
self.train_loader_iter = iter(self.train_dataloader)
# Get first batch of new epoch
batch = next(self.train_loader_iter)
latents = batch['vae_latent']
latents = latents[:, :, :self.training_args.num_latent_t]
# encoder_hidden_states = batch['text_embedding']
# encoder_attention_mask = batch['text_attention_mask']
clip_features = batch['clip_feature']
image_latents = batch['first_frame_latent']
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
pil_image = batch['pil_image']
infos = batch['info_list']
training_batch.latents = latents.to(get_local_torch_device(),
dtype=torch.bfloat16)
training_batch.encoder_hidden_states = None
training_batch.encoder_attention_mask = None
# MatrixGame doesn't use text encoder
training_batch.preprocessed_image = pil_image.to(
get_local_torch_device())
training_batch.image_embeds = clip_features.to(get_local_torch_device())
training_batch.image_latents = image_latents.to(
get_local_torch_device())
training_batch.infos = infos
# Action conditioning
if 'mouse_cond' in batch and batch['mouse_cond'].numel() > 0:
training_batch.mouse_cond = batch['mouse_cond'].to(
get_local_torch_device(), dtype=torch.bfloat16)
else:
training_batch.mouse_cond = None
if 'keyboard_cond' in batch and batch['keyboard_cond'].numel() > 0:
training_batch.keyboard_cond = batch['keyboard_cond'].to(
get_local_torch_device(), dtype=torch.bfloat16)
else:
training_batch.keyboard_cond = None
return training_batch
def _prepare_dit_inputs(self,
training_batch: TrainingBatch) -> TrainingBatch:
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
# First, call parent method to prepare noise, timesteps, etc. for video latents
training_batch = super()._prepare_dit_inputs(training_batch)
assert isinstance(training_batch.image_latents, torch.Tensor)
image_latents = training_batch.image_latents.to(
get_local_torch_device(), dtype=torch.bfloat16)
temporal_compression_ratio = self.training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (self.training_args.num_latent_t -
1) * temporal_compression_ratio + 1
batch_size, num_channels, _, latent_height, latent_width = image_latents.shape
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
latent_width)
mask_lat_size[:, :, 1:] = 0
first_frame_mask = mask_lat_size[:, :, :1]
first_frame_mask = torch.repeat_interleave(
first_frame_mask, dim=2, repeats=temporal_compression_ratio)
mask_lat_size = torch.cat([first_frame_mask, mask_lat_size[:, :, 1:]],
dim=2)
mask_lat_size = mask_lat_size.view(batch_size, -1,
temporal_compression_ratio,
latent_height, latent_width)
mask_lat_size = mask_lat_size.transpose(1, 2)
mask_lat_size = mask_lat_size.to(
image_latents.device).to(dtype=torch.bfloat16)
training_batch.noisy_model_input = torch.cat(
[training_batch.noisy_model_input, mask_lat_size, image_latents],
dim=1)
return training_batch
def _build_input_kwargs(self,
training_batch: TrainingBatch) -> TrainingBatch:
# Image Embeds for conditioning
image_embeds = training_batch.image_embeds
assert torch.isnan(image_embeds).sum() == 0
image_embeds = image_embeds.to(get_local_torch_device(),
dtype=torch.bfloat16)
encoder_hidden_states_image = image_embeds
from fastvideo.models.dits.hyworld.pose import process_custom_actions
viewmats, intrinsics, action_labels = process_custom_actions(training_batch.keyboard_cond, training_batch.mouse_cond)
viewmats = viewmats.unsqueeze(0).to(get_local_torch_device(), dtype=torch.bfloat16)
intrinsics = intrinsics.unsqueeze(0).to(get_local_torch_device(), dtype=torch.bfloat16)
action_labels = action_labels.unsqueeze(0).to(get_local_torch_device(), dtype=torch.bfloat16)
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
training_batch.input_kwargs = {
"hidden_states":
training_batch.noisy_model_input,
"encoder_hidden_states":
training_batch.encoder_hidden_states, # None for MatrixGame
"timestep":
training_batch.timesteps.to(get_local_torch_device(),
dtype=torch.bfloat16),
# "encoder_attention_mask":
# training_batch.encoder_attention_mask,
"encoder_hidden_states_image":
encoder_hidden_states_image,
# Action conditioning
"viewmats": viewmats,
"Ks": intrinsics,
"action": action_labels,
"return_dict":
False,
}
return training_batch
def _prepare_validation_batch(self, sampling_param: SamplingParam,
training_args: TrainingArgs,
validation_batch: dict[str, Any],
num_inference_steps: int) -> ForwardBatch:
sampling_param.prompt = validation_batch['prompt']
sampling_param.height = training_args.num_height
sampling_param.width = training_args.num_width
sampling_param.image_path = validation_batch.get(
'image_path') or validation_batch.get('video_path')
sampling_param.num_inference_steps = num_inference_steps
sampling_param.data_type = "video"
assert self.seed is not None
sampling_param.seed = self.seed
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
sampling_param.height // 8, sampling_param.width // 8]
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = (training_args.num_latent_t -
1) * temporal_compression_factor + 1
sampling_param.num_frames = num_frames
batch = ForwardBatch(
**shallow_asdict(sampling_param),
latents=None,
generator=torch.Generator(device="cpu").manual_seed(self.seed),
n_tokens=n_tokens,
eta=0.0,
VSA_sparsity=training_args.VSA_sparsity,
)
if "image" in validation_batch and validation_batch["image"] is not None:
batch.pil_image = validation_batch["image"]
if "keyboard_cond" in validation_batch and validation_batch[
"keyboard_cond"] is not None:
keyboard_cond = validation_batch["keyboard_cond"]
keyboard_cond = torch.tensor(keyboard_cond, dtype=torch.bfloat16)
keyboard_cond = keyboard_cond.unsqueeze(0)
batch.keyboard_cond = keyboard_cond
if "mouse_cond" in validation_batch and validation_batch[
"mouse_cond"] is not None:
mouse_cond = validation_batch["mouse_cond"]
mouse_cond = torch.tensor(mouse_cond, dtype=torch.bfloat16)
mouse_cond = mouse_cond.unsqueeze(0)
batch.mouse_cond = mouse_cond
return batch
def main(args) -> None:
logger.info("Starting training pipeline...")
pipeline = WanGameTrainingPipeline.from_pretrained(
args.pretrained_model_name_or_path, args=args)
args = pipeline.training_args
pipeline.train()
logger.info("Training pipeline done")
if __name__ == "__main__":
argv = sys.argv
from fastvideo.fastvideo_args import TrainingArgs
from fastvideo.utils import FlexibleArgumentParser
parser = FlexibleArgumentParser()
parser = TrainingArgs.add_cli_args(parser)
parser = FastVideoArgs.add_cli_args(parser)
args = parser.parse_args()
args.dit_cpu_offload = False
main(args)
+4 -2
View File
@@ -152,9 +152,10 @@ nav:
- Overview: attention/index.md
- Video Sparse Attention: attention/vsa/index.md
- Sliding Tile Attention: attention/sta/index.md
- Adding a New Attention Backend: attention/developer/index.md
- Backend Development: contributing/attention_backend.md
- Utilities:
- LoRA: utilities/lora.md
- Debugging: utilities/debugging.md
- Design:
- Overview: design/overview.md
- Developer Guide:
@@ -165,7 +166,8 @@ nav:
- RunPod: contributing/developer_env/runpod.md
- Testing: contributing/testing.md
- Profiling: contributing/profiling.md
- Adding a New Attention Backend: attention/developer/index.md
- Coding Agents: contributing/coding_agents.md
- Attention Backend Development: contributing/attention_backend.md
- API Reference:
- FastVideo: api/fastvideo.md