Compare commits

...
21 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
Matthew Noto 1b503554d1 [bugfix] fix torchvision import (#1039) 2026-01-24 22:37:14 -08:00
Shreejith SGandgemini-code-assist[bot] 351ceb7c59 [bugfix]: handle architectural differences while lora extraction (#1035)
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-01-24 15:42:28 -08:00
KyleShao 10875e0d7b [bugfix] Fix NCCL all_gather contiguity + correct ParallelTiledVAE decode tiling threshold (#1037) 2026-01-24 15:37:35 -08:00
alexzms 1eaae8a10b [ci] Increase ci test error threshold (#1038) 2026-01-24 15:36:10 -08:00
Mingjia Huo 59e00f6164 [feat] Add HY-World1.5-Bidirectional-480P-I2V (#1027)
VAE requires further improvement, will raise PR in near future.
2026-01-23 14:18:04 -08:00
745cc05b10 [bugfix] Allow update timesteps for hy1.5 model. (#1033)
Co-authored-by: Davids048 <jundasu@ucsd.edu>
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-01-22 22:04:53 -08:00
William Lin c5dc244871 [bugfix] add omegaconf as dep. (#1032) 2026-01-22 11:59:28 -08:00
alexzms dbf3917bf4 [fastvideo-kernel] replace map to index with Triton implementation + add vsa benchmark (#1029) 2026-01-22 11:35:02 -08:00
XOR-op 050f189c95 fix: SP for hunyuanvideo 1.5 (#1026) 2026-01-21 14:40:06 -08:00
Shao DuanandWill Lin 029216029f Added LTX-2 Distilled T2V Generation (#1016)
Co-authored-by: Will Lin <wlsaidhi@gmail.com>
2026-01-21 14:11:39 -08:00
169 changed files with 23844 additions and 676 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
+224
View File
@@ -0,0 +1,224 @@
# SPDX-License-Identifier: Apache-2.0
"""
Basic example for HYWorld (HY-WorldPlay) video generation using FastVideo.
This example replicates the same functionality as HY-WorldPlay/run.sh,
demonstrating image-to-video generation with camera trajectory control.
"""
import time
import math
import numpy as np
import imageio
import torchvision
from einops import rearrange
from fastvideo import VideoGenerator
from fastvideo.pipelines import ForwardBatch
from fastvideo.utils import shallow_asdict, align_to
from fastvideo.logger import init_logger
from fastvideo.models.dits.hyworld.pose import pose_to_input, compute_latent_num
from fastvideo.models.dits.hyworld.resolution_utils import get_resolution_from_image
logger = init_logger(__name__)
class HYWorldVideoGenerator(VideoGenerator):
"""Extended VideoGenerator that adds HYWorld-specific parameters to batch.extra."""
def _generate_single_video(self, prompt: str, sampling_param=None, **kwargs):
"""Override to add viewmats, Ks, and action to batch.extra."""
fastvideo_args = self.fastvideo_args
pipeline_config = fastvideo_args.pipeline_config
if sampling_param is None:
from fastvideo.configs.sample import SamplingParam
sampling_param = SamplingParam.from_pretrained(fastvideo_args.model_path)
# Update sampling param with kwargs
if kwargs:
for key, value in kwargs.items():
if hasattr(sampling_param, key):
setattr(sampling_param, key, value)
# Get pose string from sampling_param or kwargs
pose = kwargs.get('pose', getattr(sampling_param, 'POSE', 'w-31'))
num_frames = kwargs.get('num_frames', getattr(sampling_param, 'num_frames', 125))
# Calculate number of latents
latent_num = compute_latent_num(num_frames)
# Convert pose to viewmats, Ks, and action
viewmats, Ks, action = pose_to_input(pose, latent_num)
# Convert to tensors and add batch dimension
viewmats = viewmats.unsqueeze(0) # (1, T, 4, 4)
Ks = Ks.unsqueeze(0) # (1, T, 3, 3)
action = action.unsqueeze(0) # (1, T)
# Validate inputs
prompt = prompt.strip()
sampling_param = sampling_param.__class__(**shallow_asdict(sampling_param))
output_path = kwargs.get("output_path", sampling_param.output_path)
sampling_param.prompt = prompt
if sampling_param.negative_prompt is not None:
sampling_param.negative_prompt = sampling_param.negative_prompt.strip()
# Validate dimensions
if (sampling_param.height <= 0 or sampling_param.width <= 0 or
sampling_param.num_frames <= 0):
raise ValueError(
f"Height, width, and num_frames must be positive integers")
temporal_scale_factor = pipeline_config.vae_config.arch_config.temporal_compression_ratio
num_frames = sampling_param.num_frames
num_gpus = fastvideo_args.num_gpus
use_temporal_scaling_frames = pipeline_config.vae_config.use_temporal_scaling_frames
# Adjust number of frames based on number of GPUs
if use_temporal_scaling_frames:
orig_latent_num_frames = (num_frames - 1) // temporal_scale_factor + 1
else:
orig_latent_num_frames = sampling_param.num_frames // 17 * 3
if orig_latent_num_frames % fastvideo_args.num_gpus != 0:
if use_temporal_scaling_frames:
new_num_frames = (orig_latent_num_frames - 1) * temporal_scale_factor + 1
else:
divisor = math.lcm(3, num_gpus)
orig_latent_num_frames = (
(orig_latent_num_frames + divisor - 1) // divisor) * divisor
new_num_frames = orig_latent_num_frames // 3 * 17
logger.info(
"Adjusting number of frames from %s to %s based on number of GPUs (%s)",
sampling_param.num_frames, new_num_frames, fastvideo_args.num_gpus)
sampling_param.num_frames = new_num_frames
# Calculate sizes
target_height = align_to(sampling_param.height, 16)
target_width = align_to(sampling_param.width, 16)
# Calculate latent sizes
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]
# Prepare batch
batch = ForwardBatch(
**shallow_asdict(sampling_param),
eta=0.0,
n_tokens=n_tokens,
VSA_sparsity=fastvideo_args.VSA_sparsity,
)
# Add HYWorld-specific parameters to batch.extra
batch.extra['viewmats'] = viewmats
batch.extra['Ks'] = Ks
batch.extra['action'] = action
batch.extra['chunk_latent_frames'] = 16 # For bidirectional model
# Run inference
start_time = time.perf_counter()
output_batch = self.executor.execute_forward(batch, fastvideo_args)
samples = output_batch.output
logging_info = output_batch.logging_info
gen_time = time.perf_counter() - start_time
logger.info("Generated successfully in %.2f seconds", gen_time)
# Process outputs
videos = rearrange(samples, "b c t h w -> t b c h w")
frames = []
for x in videos:
x = torchvision.utils.make_grid(x, nrow=6)
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
frames.append((x * 255).numpy().astype(np.uint8))
# Save video if requested
if batch.save_video:
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
if batch.return_frames:
return frames
else:
return {
"samples": samples,
"frames": frames,
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
"logging_info": logging_info,
"trajectory": output_batch.trajectory_latents,
"trajectory_timesteps": output_batch.trajectory_timesteps,
"trajectory_decoded": output_batch.trajectory_decoded,
}
# Default prompt from HY-WorldPlay run.sh
DEFAULT_PROMPT = 'A paved pathway leads towards a stone arch bridge spanning a calm body of water. Lush green trees and foliage line the path and the far bank of the water. A traditional-style pavilion with a tiered, reddish-brown roof sits on the far shore. The water reflects the surrounding greenery and the sky. The scene is bathed in soft, natural light, creating a tranquil and serene atmosphere. The pathway is composed of large, rectangular stones, and the bridge is constructed of light gray stone. The overall composition emphasizes the peaceful and harmonious nature of the landscape.'
DEFAULT_IMAGE = 'https://raw.githubusercontent.com/Tencent-Hunyuan/HY-WorldPlay/main/assets/img/test.png'
def main():
import argparse
# pose: (a, w, s, d) - (15, 31)
# num_frames: (61, 125)
parser = argparse.ArgumentParser(description="HYWorld video generation with FastVideo")
parser.add_argument("--prompt", type=str, default=DEFAULT_PROMPT, help="Text prompt for video generation")
parser.add_argument("--image", type=str, default=DEFAULT_IMAGE, help="Path or URL to input image")
parser.add_argument("--pose", type=str, default='w-31', help="Pose string (e.g., 'a-31', 'w-31', 's-31', 'd-31')")
parser.add_argument("--output_path", type=str, default='video_samples_hyworld', help="Output video path")
parser.add_argument("--num-frames", type=int, default=125, help="Number of frames")
parser.add_argument("--seed", type=int, default=1, help="Random seed")
parser.add_argument("--resolution", type=str, default="480p", help="Only support 480p for now")
args = parser.parse_args()
# Automatically determine resolution from input image
HEIGHT, WIDTH = get_resolution_from_image(args.image, args.resolution)
print(f"Image: {args.image}")
print(f"Pose: {args.pose}")
print(f"Resolution: {HEIGHT}x{WIDTH} (from {args.resolution} buckets)")
print(f"Num frames: {args.num_frames}")
print(f"Output path: {args.output_path}")
# Initialize generator
print("\nInitializing VideoGenerator for HYWorld...")
generator = HYWorldVideoGenerator.from_pretrained(
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers",
num_gpus=1,
use_fsdp_inference=True,
dit_cpu_offload=True,
vae_cpu_offload=True,
text_encoder_cpu_offload=True,
pin_cpu_memory=True,
image_encoder_cpu_offload=True,
)
# Generate video
print("\nGenerating video...")
start_time = time.time()
video = generator.generate_video(
prompt=args.prompt,
image_path=args.image,
output_path=args.output_path,
save_video=True,
negative_prompt="",
num_frames=args.num_frames,
fps=24,
height=HEIGHT,
width=WIDTH,
seed=args.seed,
pose=args.pose,
)
elapsed = time.time() - start_time
print(f"\nVideo generated successfully!")
print(f"Saved to: {args.output_path}")
print(f"Time: {elapsed:.2f}s")
if __name__ == "__main__":
main()
+34
View File
@@ -0,0 +1,34 @@
from fastvideo import VideoGenerator
PROMPT = (
"A warm sunny backyard. The camera starts in a tight cinematic close-up "
"of a woman and a man in their 30s, facing each other with serious "
"expressions. The woman, emotional and dramatic, says softly, \"That's "
"it... Dad's lost it. And we've lost Dad.\" The man exhales, slightly "
"annoyed: \"Stop being so dramatic, Jess.\" A beat. He glances aside, "
"then mutters defensively, \"He's just having fun.\" The camera slowly "
"pans right, revealing the grandfather in the garden wearing enormous "
"butterfly wings, waving his arms in the air like he's trying to take "
"off. He shouts, \"Wheeeew!\" as he flaps his wings with full commitment. "
"The woman covers her face, on the verge of tears. The tone is deadpan, "
"absurd, and quietly tragic."
)
def main() -> None:
generator = VideoGenerator.from_pretrained(
"FastVideo/LTX2-Distilled-Diffusers",
num_gpus=1,
)
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
generator.generate_video(
prompt=PROMPT,
output_path=output_path,
save_video=True,
)
generator.shutdown()
if __name__ == "__main__":
main()
+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
}
]
}
+11
View File
@@ -40,6 +40,17 @@ out = video_sparse_attn(q, k, v, block_sizes, block_sizes, topk=5)
out = moba_attn_varlen(q, k, v, cu_seqlens_q, cu_seqlens_k, ...)
```
## Benchmark
### VSA (block-sparse) TFLOPs
After building/installing `fastvideo-kernel`, run:
```bash
cd fastvideo-kernel
python benchmarks/bench_vsa.py --batch_size 1 --num_heads 16 --head_dim 128 --q_seq_lens 49152 --topk 64
```
### TurboDiffusion Kernels
This package also includes kernels from [TurboDiffusion](https://github.com/thu-ml/TurboDiffusion), including INT8 GEMM, Quantization, RMSNorm and LayerNorm.
+166
View File
@@ -0,0 +1,166 @@
#!/usr/bin/env python3
"""
Benchmark VSA *wrapper* performance (forward + backward) and report TFLOPs.
This script benchmarks the autograd-enabled wrapper:
- fastvideo_kernel.block_sparse_attn.block_sparse_attn
So measured time includes wrapper overhead (map->index conversion, dispatch) plus kernel time.
"""
from __future__ import annotations
import argparse
import os
import random
from typing import Tuple, Callable
import numpy as np
import torch
try:
from triton.testing import do_bench
except Exception as e: # pragma: no cover
raise ImportError("This benchmark requires triton (for triton.testing.do_bench).") from e
BLOCK_M = 64
BLOCK_N = 64
def set_seed(seed: int = 42) -> None:
random.seed(seed)
np.random.seed(seed)
torch.manual_seed(seed)
torch.cuda.manual_seed(seed)
torch.cuda.manual_seed_all(seed)
def parse_arguments() -> argparse.Namespace:
p = argparse.ArgumentParser(description="Benchmark FastVideo VSA block-sparse attention")
p.add_argument("--batch_size", type=int, default=1)
p.add_argument("--num_heads", type=int, default=12)
p.add_argument("--head_dim", type=int, default=128, choices=[64, 128])
p.add_argument("--topk", type=int, default=None, help="KV blocks per Q block (default: ~90%% sparsity)")
p.add_argument("--q_seq_lens", type=int, nargs="+", default=[49152], help="Q sequence lengths (must be /64)")
p.add_argument("--kv_seq_lens", type=int, nargs="+", default=None, help="KV sequence lengths (defaults to q_seq_len)")
p.add_argument("--warmup", type=int, default=5)
p.add_argument("--rep", type=int, default=20)
p.add_argument("--seed", type=int, default=42)
p.add_argument("--dtype", type=str, default="bf16", choices=["bf16", "fp16"])
p.add_argument("--force_triton", action="store_true", help="Force wrapper to use Triton path (if supported by shapes).")
return p.parse_args()
def create_qkv(batch: int, heads: int, q_len: int, kv_len: int, d: int, dtype: torch.dtype) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
q = torch.randn(batch, heads, q_len, d, dtype=dtype, device="cuda")
k = torch.randn(batch, heads, kv_len, d, dtype=dtype, device="cuda")
v = torch.randn(batch, heads, kv_len, d, dtype=dtype, device="cuda")
return q, k, v
def make_block_map(bs: int, h: int, num_q_blocks: int, num_kv_blocks: int, topk: int) -> torch.Tensor:
# block_map: [bs, h, num_q_blocks, num_kv_blocks] bool
scores = torch.rand(bs, h, num_q_blocks, num_kv_blocks, device="cuda")
topk = min(max(1, topk), num_kv_blocks)
idx = torch.topk(scores, topk, dim=-1).indices
block_map = torch.zeros(bs, h, num_q_blocks, num_kv_blocks, dtype=torch.bool, device="cuda")
block_map.scatter_(-1, idx, True)
return block_map
def flops_sparse_attention(bs: int, h: int, d: int, q_len: int, topk_blocks: int, block_n: int) -> float:
# Approx: QK^T + PV, each is ~2*bs*h*q_len*(topk_blocks*block_n)*d
return 4.0 * bs * h * d * q_len * (topk_blocks * block_n)
def bench_ms(fn: Callable[[], object], warmup: int, rep: int) -> float:
return do_bench(fn, warmup=warmup, rep=rep, quantiles=None)
def main() -> None:
args = parse_arguments()
set_seed(args.seed)
dtype = torch.bfloat16 if args.dtype == "bf16" else torch.float16
if args.force_triton:
os.environ["FASTVIDEO_KERNEL_VSA_FORCE_TRITON"] = "1"
from fastvideo_kernel.block_sparse_attn import block_sparse_attn
bs, h, d = args.batch_size, args.num_heads, args.head_dim
kv_seq_lens = args.kv_seq_lens
if kv_seq_lens is None:
kv_seq_lens = args.q_seq_lens
if len(kv_seq_lens) != len(args.q_seq_lens):
raise ValueError("kv_seq_lens must have the same number of entries as q_seq_lens (or be omitted).")
print("VSA Block-Sparse Attention Benchmark (WRAPPER)")
print(f"device: {torch.cuda.get_device_name(0)}")
print(f"batch={bs}, heads={h}, head_dim={d}, dtype={args.dtype}")
print(f"BLOCK_M={BLOCK_M}, BLOCK_N={BLOCK_N}")
print("NOTE: timings include wrapper overhead (map->index + dispatch).")
if args.force_triton:
print("dispatch: forced Triton (FASTVIDEO_KERNEL_VSA_FORCE_TRITON=1)")
else:
print("dispatch: SM90 if available, else Triton")
for q_len, kv_len in zip(args.q_seq_lens, kv_seq_lens):
if q_len % BLOCK_M != 0 or kv_len % BLOCK_N != 0:
print(f"[skip] q_len={q_len}, kv_len={kv_len} must be divisible by 64")
continue
num_q_blocks = q_len // BLOCK_M
num_kv_blocks = kv_len // BLOCK_N
topk = args.topk if args.topk is not None else max(1, num_kv_blocks // 10)
topk = min(topk, num_kv_blocks)
print("\n" + "=" * 80)
print(f"q_len={q_len}, kv_len={kv_len}, num_q_blocks={num_q_blocks}, num_kv_blocks={num_kv_blocks}, topk={topk}")
q, k, v = create_qkv(bs, h, q_len, kv_len, d, dtype)
block_map = make_block_map(bs, h, num_q_blocks, num_kv_blocks, topk)
# Variable block sizes: default full blocks (64 tokens per KV block)
variable_block_sizes = torch.full((num_kv_blocks,), BLOCK_N, dtype=torch.int32, device="cuda")
def _fwd():
return block_sparse_attn(q, k, v, block_map, variable_block_sizes)
fwd_ms = bench_ms(_fwd, warmup=args.warmup, rep=args.rep)
# Backward benchmark (wrapper autograd). We build the graph once, then repeatedly run backward
# on the retained graph so bwd timing excludes the forward compute.
q_ = q.detach().requires_grad_(True)
k_ = k.detach().requires_grad_(True)
v_ = v.detach().requires_grad_(True)
o_, _aux_ = block_sparse_attn(q_, k_, v_, block_map, variable_block_sizes)
og = torch.randn_like(o_)
loss = (o_ * og).sum()
for _ in range(max(1, args.warmup // 2)):
torch.autograd.grad(loss, (q_, k_, v_), retain_graph=True)
torch.cuda.synchronize()
bwd_ms = bench_ms(
lambda: torch.autograd.grad(loss, (q_, k_, v_), retain_graph=True),
warmup=0,
rep=max(5, args.rep // 2),
)
flops = flops_sparse_attention(bs, h, d, q_len, topk, BLOCK_N)
fwd_tflops = flops / fwd_ms * 1e-12 * 1e3
# Rough backward multiplier (attention backward typically ~2-3x forward)
bwd_tflops = (2.5 * flops) / bwd_ms * 1e-12 * 1e3
print(f"fwd(wrapper): {fwd_ms:.3f} ms | {fwd_tflops:.2f} TFLOPs (approx)")
print(f"bwd(wrapper): {bwd_ms:.3f} ms | {bwd_tflops:.2f} TFLOPs (approx)")
if __name__ == "__main__":
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for this benchmark.")
main()
+1 -1
View File
@@ -9,7 +9,7 @@ build-backend = "scikit_build_core.build"
[project]
name = "fastvideo-kernel"
version = "0.2.4"
version = "0.2.5"
description = "Unified CUDA kernels for FastVideo"
readme = "README.md"
requires-python = ">=3.10"
@@ -30,13 +30,12 @@ def _force_triton() -> bool:
return os.environ.get("FASTVIDEO_KERNEL_VSA_FORCE_TRITON", "0") == "1"
def _map_to_index_torch(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
def _map_to_index(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Pure-torch (no triton) conversion:
block_map: [B, H, Q, KV] bool (or [H, Q, KV] which will be treated as B=1)
returns:
index: [B, H, Q, KV] int32 (packed KV indices, -1 padding)
num: [B, H, Q] int32 (#kv blocks per q block)
Preferred map->index conversion used by the wrapper.
This wrapper **requires** the Triton implementation.
If Triton (or the Triton map_to_index module) is not available, it raises.
"""
if block_map.dim() == 3:
block_map = block_map.unsqueeze(0)
@@ -45,20 +44,17 @@ def _map_to_index_torch(block_map: torch.Tensor) -> Tuple[torch.Tensor, torch.Te
if block_map.dtype != torch.bool:
block_map = block_map.to(torch.bool)
B, H, Q, KV = block_map.shape
index = torch.full((B, H, Q, KV), -1, dtype=torch.int32, device=block_map.device)
num = torch.zeros((B, H, Q), dtype=torch.int32, device=block_map.device)
if not block_map.is_cuda:
raise RuntimeError("block_map must be a CUDA tensor (Triton map_to_index required).")
# Small sizes in practice (B=1, H<=16, Q/KV<=64), so a Python loop is fine.
for b in range(B):
for h in range(H):
for q in range(Q):
kv_idx = torch.nonzero(block_map[b, h, q], as_tuple=False).flatten().to(torch.int32)
n = int(kv_idx.numel())
if n:
index[b, h, q, :n] = kv_idx
num[b, h, q] = n
return index, num
try:
from fastvideo_kernel.triton_kernels.index import map_to_index as triton_map_to_index # local import
except Exception as e:
raise ImportError(
"Triton map_to_index is required but not available. "
"Ensure Triton is installed and fastvideo_kernel.triton_kernels.index is importable."
) from e
return triton_map_to_index(block_map)
@torch.library.custom_op(
@@ -77,7 +73,7 @@ def block_sparse_attn_triton(
k = k.contiguous()
v = v.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
q2k_idx, q2k_num = _map_to_index(block_map)
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
triton_block_sparse_attn_forward,
@@ -87,6 +83,7 @@ def block_sparse_attn_triton(
return o, M
@torch.library.register_fake("fastvideo_kernel::block_sparse_attn_triton")
def _block_sparse_attn_triton_fake(
q: torch.Tensor,
@@ -117,8 +114,8 @@ def block_sparse_attn_backward_triton(
) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
grad_output = grad_output.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
q2k_idx, q2k_num = _map_to_index(block_map)
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
from fastvideo_kernel.triton_kernels.block_sparse_attn_triton import ( # local import
triton_block_sparse_attn_backward,
@@ -182,7 +179,7 @@ def block_sparse_attn_sm90(
k_padded = k_padded.contiguous()
v_padded = v_padded.contiguous()
block_map = block_map.to(torch.bool)
q2k_idx, q2k_num = _map_to_index_torch(block_map)
q2k_idx, q2k_num = _map_to_index(block_map)
o_padded, lse_padded = block_sparse_fwd(
q_padded, k_padded, v_padded, q2k_idx, q2k_num, variable_block_sizes.int()
@@ -224,7 +221,7 @@ def block_sparse_attn_backward_sm90(
grad_output_padded = grad_output_padded.contiguous()
block_map = block_map.to(torch.bool)
k2q_idx, k2q_num = _map_to_index_torch(block_map.transpose(-1, -2).contiguous())
k2q_idx, k2q_num = _map_to_index(block_map.transpose(-1, -2).contiguous())
dq, dk, dv = block_sparse_bwd(
q_padded,
@@ -1 +1 @@
__version__ = "0.2.4"
__version__ = "0.2.5"
+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
+12 -1
View File
@@ -2,5 +2,16 @@ from fastvideo.configs.models.base import ModelConfig
from fastvideo.configs.models.dits.base import DiTConfig
from fastvideo.configs.models.encoders.base import EncoderConfig
from fastvideo.configs.models.vaes.base import VAEConfig
from fastvideo.configs.models.audio import (LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig)
__all__ = ["ModelConfig", "VAEConfig", "DiTConfig", "EncoderConfig"]
__all__ = [
"ModelConfig",
"VAEConfig",
"DiTConfig",
"EncoderConfig",
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
]
@@ -0,0 +1,13 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.configs.models.audio.ltx2_audio_vae import (
LTX2AudioDecoderConfig,
LTX2AudioEncoderConfig,
LTX2VocoderConfig,
)
__all__ = [
"LTX2AudioEncoderConfig",
"LTX2AudioDecoderConfig",
"LTX2VocoderConfig",
]
@@ -0,0 +1,31 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 audio VAE and vocoder configuration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.base import ArchConfig, ModelConfig
@dataclass
class LTX2AudioArchConfig(ArchConfig):
architectures: list[str] = field(default_factory=list)
@dataclass
class LTX2AudioEncoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2AudioEncoder"]))
@dataclass
class LTX2AudioDecoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2AudioDecoder"]))
@dataclass
class LTX2VocoderConfig(ModelConfig):
arch_config: ArchConfig = field(default_factory=lambda: LTX2AudioArchConfig(
architectures=["LTX2Vocoder"]))
+3 -1
View File
@@ -3,11 +3,13 @@ from fastvideo.configs.models.dits.cosmos2_5 import Cosmos25VideoConfig
from fastvideo.configs.models.dits.hunyuanvideo import HunyuanVideoConfig
from fastvideo.configs.models.dits.hunyuanvideo15 import HunyuanVideo15Config
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
from fastvideo.configs.models.dits.wanvideo import WanVideoConfig
from fastvideo.configs.models.dits.hyworld import HYWorldConfig
__all__ = [
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
"LongCatVideoConfig"
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
]
+207
View File
@@ -0,0 +1,207 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_double_block(n: str, m) -> bool:
return "double" in n and str.isdigit(n.split(".")[-1])
def is_single_block(n: str, m) -> bool:
return "single" in n and str.isdigit(n.split(".")[-1])
# def is_refiner_block(n: str, m) -> bool:
# return "refiner" in n and str.isdigit(n.split(".")[-1])
def is_txt_in(n: str, m) -> bool:
return n.split(".")[-1] == "txt_in"
@dataclass
class HYWorldArchConfig(DiTArchConfig):
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block])
_compile_conditions: list = field(
default_factory=lambda: [is_double_block, is_single_block, is_txt_in])
param_names_mapping: dict = field(
default_factory=lambda: {
# 1. txt_in submodules (text embedder, refiner blocks):
r"^txt_in\.t_embedder\.mlp\.0\.(.*)$":
r"txt_in.t_embedder.mlp.fc_in.\1",
r"^txt_in\.t_embedder\.mlp\.2\.(.*)$":
r"txt_in.t_embedder.mlp.fc_out.\1",
r"^txt_in\.c_embedder\.linear_1\.(.*)$":
r"txt_in.c_embedder.fc_in.\1",
r"^txt_in\.c_embedder\.linear_2\.(.*)$":
r"txt_in.c_embedder.fc_out.\1",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm1\.(.*)$":
r"txt_in.refiner_blocks.\1.norm1.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.norm2\.(.*)$":
r"txt_in.refiner_blocks.\1.norm2.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_qkv\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_qkv.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.self_attn_proj\.(.*)$":
r"txt_in.refiner_blocks.\1.self_attn_proj.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.mlp\.fc1\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_in.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.mlp\.fc2\.(.*)$":
r"txt_in.refiner_blocks.\1.mlp.fc_out.\2",
r"^txt_in\.individual_token_refiner\.blocks\.(\d+)\.adaLN_modulation\.1\.(.*)$":
r"txt_in.refiner_blocks.\1.adaLN_modulation.linear.\2",
# 2. time_in mappings (HYWorld uses TimestepEmbedder directly,
# but FastVideo model inherits HunyuanVideo15TimeEmbedding with timestep_embedder):
r"^time_in\.mlp\.0\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_in.\1",
r"^time_in\.mlp\.2\.(.*)$":
r"time_in.timestep_embedder.mlp.fc_out.\1",
# 3. action_in mappings:
r"^action_in\.mlp\.0\.(.*)$":
r"action_in.mlp.fc_in.\1",
r"^action_in\.mlp\.2\.(.*)$":
r"action_in.mlp.fc_out.\1",
# 4. byt5_in -> txt_in_2 mappings:
r"^byt5_in\.layernorm\.(.*)$":
r"txt_in_2.norm.\1",
r"^byt5_in\.fc1\.(.*)$":
r"txt_in_2.linear_1.\1",
r"^byt5_in\.fc2\.(.*)$":
r"txt_in_2.linear_2.\1",
r"^byt5_in\.fc3\.(.*)$":
r"txt_in_2.linear_3.\1",
# 5. cond_type_embedding -> cond_type_embed:
r"^cond_type_embedding\.(.*)$":
r"cond_type_embed.\1",
# 6. vision_in -> image_embedder mappings:
r"^vision_in\.proj\.0\.(.*)$":
r"image_embedder.norm_in.\1",
r"^vision_in\.proj\.1\.(.*)$":
r"image_embedder.linear_1.\1",
r"^vision_in\.proj\.3\.(.*)$":
r"image_embedder.linear_2.\1",
r"^vision_in\.proj\.4\.(.*)$":
r"image_embedder.norm_out.\1",
# 7. double_blocks mapping:
r"^double_blocks\.(\d+)\.img_attn_q\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 0, 3),
r"^double_blocks\.(\d+)\.img_attn_k\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 1, 3),
r"^double_blocks\.(\d+)\.img_attn_v\.(.*)$":
(r"double_blocks.\1.img_attn_qkv.\2", 2, 3),
r"^double_blocks\.(\d+)\.txt_attn_q\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 0, 3),
r"^double_blocks\.(\d+)\.txt_attn_k\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 1, 3),
r"^double_blocks\.(\d+)\.txt_attn_v\.(.*)$":
(r"double_blocks.\1.txt_attn_qkv.\2", 2, 3),
r"^double_blocks\.(\d+)\.img_mlp\.fc1\.(.*)$":
r"double_blocks.\1.img_mlp.fc_in.\2",
r"^double_blocks\.(\d+)\.img_mlp\.fc2\.(.*)$":
r"double_blocks.\1.img_mlp.fc_out.\2",
r"^double_blocks\.(\d+)\.txt_mlp\.fc1\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_in.\2",
r"^double_blocks\.(\d+)\.txt_mlp\.fc2\.(.*)$":
r"double_blocks.\1.txt_mlp.fc_out.\2",
# 8. Final layer mapping:
r"^final_layer\.adaLN_modulation\.1\.(.*)$":
r"final_layer.adaLN_modulation.linear.\1",
})
# Reverse mapping for saving checkpoints: custom -> hf
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
# Parameters from HY-WorldPlay config.json (loaded from checkpoint)
patch_size: list | tuple | int = field(default_factory=lambda: [1, 1, 1])
# Base latent channels - will be expanded in __post_init__ if concat_condition=True
in_channels: int = 32
concat_condition: bool = True
out_channels: int = 32
hidden_size: int = 2048
heads_num: int = 16
mlp_width_ratio: float = 4.0
mlp_act_type: str = "gelu_tanh"
mm_double_blocks_depth: int = 54
mm_single_blocks_depth: int = 0
rope_dim_list: list | tuple = field(default_factory=lambda: [16, 56, 56])
qkv_bias: bool = True
qk_norm: bool | str = True
qk_norm_type: str = "rms"
guidance_embed: bool = False
use_meanflow: bool = False
text_projection: str = "single_refiner"
use_attention_mask: bool = True
text_states_dim: int = 3584
text_states_dim_2: int | None = None
text_pool_type: str | None = None
rope_theta: float = 256.0
attn_mode: str = "flash"
attn_param: str | None = None
glyph_byT5_v2: bool = True
vision_projection: str = "linear"
vision_states_dim: int = 1152
is_reshape_temporal_channels: bool = False
use_cond_type_embedding: bool = True
ideal_resolution: str = "480p"
ideal_task: str = "i2v"
task_type: str = "i2v"
exclude_lora_layers: list[str] = field(
default_factory=lambda: ["img_in", "txt_in", "time_in", "vector_in"])
def __post_init__(self):
super().__post_init__()
# Convert HY-WorldPlay naming to FastVideo naming conventions
self.num_attention_heads: int = self.heads_num
self.attention_head_dim: int = self.hidden_size // self.heads_num
self.num_layers: int = self.mm_double_blocks_depth
self.num_single_layers: int = self.mm_single_blocks_depth
self.num_refiner_layers: int = 2 # Default for HYWorld
self.mlp_ratio: float = float(self.mlp_width_ratio)
self.text_embed_dim: int = self.text_states_dim
self.text_embed_2_dim: int = self.text_states_dim_2 if self.text_states_dim_2 else 1472
self.image_embed_dim: int = self.vision_states_dim
self.rope_axes_dim: tuple[int, ...] = tuple(self.rope_dim_list)
self.num_channels_latents: int = self.out_channels
self.target_size: int = 640
# Handle concat_condition: when True, actual in_channels = base * 2 + 1
# (base latent + condition latent + mask channel)
# config.json has base in_channels (32), but img_in needs full (65)
if self.concat_condition and self.in_channels == 32:
if self.is_reshape_temporal_channels:
self.in_channels = self.in_channels + self.in_channels // 2 + 1
else:
self.in_channels = self.in_channels * 2 + 1 # 32 * 2 + 1 = 65
# Handle patch_size (can be list/tuple or int)
if isinstance(self.patch_size, list | tuple):
self.patch_size_t: int = self.patch_size[0]
# assume square patch size for height and width
patch_size_hw: int = self.patch_size[1]
object.__setattr__(self, 'patch_size', patch_size_hw)
else:
self.patch_size_t = 1
# Convert qk_norm to string format
if isinstance(self.qk_norm, bool):
if self.qk_norm:
self.qk_norm = "rms_norm" if self.qk_norm_type == "rms" else self.qk_norm_type
else:
self.qk_norm = "none"
@dataclass
class HYWorldConfig(DiTConfig):
arch_config: DiTArchConfig = field(default_factory=HYWorldArchConfig)
prefix: str = "HYWorld"
+84
View File
@@ -0,0 +1,84 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 Transformer configuration for native FastVideo integration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
def is_ltx2_blocks(name: str, _module) -> bool:
"""FSDP shard condition for LTX-2 transformer blocks."""
return "transformer_blocks" in name
@dataclass
class LTX2VideoArchConfig(DiTArchConfig):
"""Architecture configuration for LTX-2 video transformer."""
_fsdp_shard_conditions: list = field(
default_factory=lambda: [is_ltx2_blocks])
_compile_conditions: list = field(default_factory=lambda: [is_ltx2_blocks])
# Parameter name mapping for weight conversion (hf/comfy -> FastVideo)
param_names_mapping: dict = field(
default_factory=lambda: {
r"^model\.diffusion_model\.(.*)$": r"model.\1",
r"^diffusion_model\.(.*)$": r"model.\1",
r"^model\.(.*)$": r"model.\1",
r"^(.*)$": r"model.\1",
})
reverse_param_names_mapping: dict = field(default_factory=lambda: {})
lora_param_names_mapping: dict = field(default_factory=lambda: {})
# Core transformer settings (defaults from LTX-2 metadata)
num_attention_heads: int = 32
attention_head_dim: int = 128
num_layers: int = 48
cross_attention_dim: int = 4096
caption_channels: int = 3840
norm_eps: float = 1e-6
attention_type: str = "default"
rope_type: str = "split"
double_precision_rope: bool = True
positional_embedding_theta: float = 10000.0
positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [20, 2048, 2048])
timestep_scale_multiplier: int = 1000
use_middle_indices_grid: bool = True
# Patchification (video-only path)
patch_size: tuple[int, int, int] = (1, 1, 1)
num_channels_latents: int = 128
in_channels: int | None = None
out_channels: int | None = None
# Audio defaults (reserved for joint AV ports)
audio_num_attention_heads: int = 32
audio_attention_head_dim: int = 64
audio_in_channels: int = 128
audio_out_channels: int = 128
audio_cross_attention_dim: int = 2048
audio_positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [20])
av_ca_timestep_scale_multiplier: int = 1
def __post_init__(self):
super().__post_init__()
patch_volume = self.patch_size[0] * self.patch_size[
1] * self.patch_size[2]
if self.in_channels is None:
self.in_channels = self.num_channels_latents * patch_volume
if self.out_channels is None:
self.out_channels = self.in_channels
@dataclass
class LTX2VideoConfig(DiTConfig):
"""Main configuration for LTX-2 transformer."""
arch_config: DiTArchConfig = field(default_factory=LTX2VideoArchConfig)
prefix: str = "ltx2"
+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"
@@ -7,11 +7,14 @@ from fastvideo.configs.models.encoders.clip import (
from fastvideo.configs.models.encoders.llama import LlamaConfig
from fastvideo.configs.models.encoders.t5 import T5Config, T5LargeConfig
from fastvideo.configs.models.encoders.qwen2_5 import Qwen2_5_VLConfig
from fastvideo.configs.models.encoders.siglip import SiglipVisionConfig
from fastvideo.configs.models.encoders.reason1 import Reason1ArchConfig, Reason1Config
from fastvideo.configs.models.encoders.gemma import LTX2GemmaConfig
__all__ = [
"EncoderConfig", "TextEncoderConfig", "ImageEncoderConfig",
"BaseEncoderOutput", "CLIPTextConfig", "CLIPVisionConfig",
"WAN2_1ControlCLIPVisionConfig", "LlamaConfig", "T5Config", "T5LargeConfig",
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config"
"Qwen2_5_VLConfig", "Reason1ArchConfig", "Reason1Config", "LTX2GemmaConfig",
"SiglipVisionConfig"
]
@@ -0,0 +1,48 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (
TextEncoderArchConfig,
TextEncoderConfig,
)
@dataclass
class LTX2GemmaArchConfig(TextEncoderArchConfig):
architectures: list[str] = field(
default_factory=lambda: ["LTX2GemmaTextEncoderModel"])
hidden_size: int = 3840
num_hidden_layers: int = 48
num_attention_heads: int = 30
text_len: int = 1024
pad_token_id: int = 0
eos_token_id: int = 2
gemma_model_path: str = ""
gemma_dtype: str = "bfloat16"
padding_side: str = "left"
feature_extractor_in_features: int = 3840 * 49
feature_extractor_out_features: int = 3840
connector_num_attention_heads: int = 30
connector_attention_head_dim: int = 128
connector_num_layers: int = 2
connector_positional_embedding_theta: float = 10000.0
connector_positional_embedding_max_pos: list[int] = field(
default_factory=lambda: [4096])
connector_rope_type: str = "split"
connector_double_precision_rope: bool = False
connector_num_learnable_registers: int | None = 128
def __post_init__(self) -> None:
super().__post_init__()
self.tokenizer_kwargs["padding"] = "max_length"
@dataclass
class LTX2GemmaConfig(TextEncoderConfig):
arch_config: TextEncoderArchConfig = field(
default_factory=LTX2GemmaArchConfig)
prefix: str = "ltx2_gemma"
@@ -0,0 +1,53 @@
# SPDX-License-Identifier: Apache-2.0
"""SigLIP vision encoder configuration for FastVideo."""
from dataclasses import dataclass, field
from fastvideo.configs.models.encoders.base import (ImageEncoderArchConfig,
ImageEncoderConfig)
@dataclass
class SiglipVisionArchConfig(ImageEncoderArchConfig):
"""Architecture configuration for SigLIP vision encoder.
Fields match the config.json from HuggingFace SigLIP checkpoints.
"""
# From config.json
architectures: list[str] = field(
default_factory=lambda: ["SiglipVisionModel"])
attention_dropout: float = 0.0
dtype: str | None = None
hidden_act: str = "gelu_pytorch_tanh"
hidden_size: int = 1152
image_size: int = 384
intermediate_size: int = 4304
layer_norm_eps: float = 1e-6
model_type: str = "siglip_vision_model"
num_attention_heads: int = 16
num_channels: int = 3
num_hidden_layers: int = 27
patch_size: int = 14
# FastVideo specific - QKV fusion mapping
stacked_params_mapping: list = field(default_factory=lambda: [
("qkv_proj", "q_proj", "q"),
("qkv_proj", "k_proj", "k"),
("qkv_proj", "v_proj", "v"),
])
@dataclass
class SiglipVisionConfig(ImageEncoderConfig):
"""Configuration for SigLIP vision encoder."""
arch_config: ImageEncoderArchConfig = field(
default_factory=SiglipVisionArchConfig)
# FastVideo specific
num_hidden_layers_override: int | None = None
require_post_norm: bool | None = None
enable_scale: bool = True
is_causal: bool = False
prefix: str = "siglip"
@@ -2,6 +2,7 @@ from fastvideo.configs.models.vaes.cosmosvae import CosmosVAEConfig
from fastvideo.configs.models.vaes.cosmos2_5vae import Cosmos25VAEConfig
from fastvideo.configs.models.vaes.hunyuanvae import HunyuanVAEConfig
from fastvideo.configs.models.vaes.hunyuan15vae import Hunyuan15VAEConfig
from fastvideo.configs.models.vaes.ltx2vae import LTX2VAEConfig
from fastvideo.configs.models.vaes.stepvideovae import StepVideoVAEConfig
from fastvideo.configs.models.vaes.wanvae import WanVAEConfig
@@ -12,4 +13,5 @@ __all__ = [
"CosmosVAEConfig",
"Cosmos25VAEConfig",
"Hunyuan15VAEConfig",
"LTX2VAEConfig",
]
+45
View File
@@ -0,0 +1,45 @@
# SPDX-License-Identifier: Apache-2.0
"""
LTX-2 VAE configuration.
"""
from dataclasses import dataclass, field
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
@dataclass
class LTX2VAEArchConfig(VAEArchConfig):
# Mirrors LTX-2 safetensors metadata config under "vae"
_class_name: str = "CausalVideoAutoencoder"
dims: int = 3
in_channels: int = 3
out_channels: int = 3
latent_channels: int = 128
encoder_blocks: list = field(default_factory=list)
decoder_blocks: list = field(default_factory=list)
patch_size: int = 4
norm_layer: str = "pixel_norm"
latent_log_var: str = "uniform"
encoder_spatial_padding_mode: str = "zeros"
decoder_spatial_padding_mode: str = "reflect"
causal_decoder: bool = False
timestep_conditioning: bool = True
use_quant_conv: bool = False
scaling_factor: float = 1.0
normalize_latent_channels: bool = False
# Match FastVideo naming for compression ratios (LTX-2 default)
temporal_compression_ratio: int = 8
spatial_compression_ratio: int = 32
@dataclass
class LTX2VAEConfig(VAEConfig):
arch_config: VAEArchConfig = field(default_factory=LTX2VAEArchConfig)
# LTX-2 tiling defaults (match ltx_core.video_vae.TilingConfig.default()).
ltx2_spatial_tile_size_in_pixels: int = 512
ltx2_spatial_tile_overlap_in_pixels: int = 64
ltx2_temporal_tile_size_in_frames: int = 64
ltx2_temporal_tile_overlap_in_frames: int = 24
+4 -1
View File
@@ -4,6 +4,8 @@ from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.registry import (
get_pipeline_config_cls_from_name)
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
@@ -16,5 +18,6 @@ __all__ = [
"Hunyuan15T2V480PConfig", "Hunyuan15T2V720PConfig", "SlidingTileAttnConfig",
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
"CosmosConfig", "Cosmos25Config", "get_pipeline_config_cls_from_name"
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
"get_pipeline_config_cls_from_name"
]
+29
View File
@@ -0,0 +1,29 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.models import DiTConfig, EncoderConfig
from fastvideo.configs.models.dits import HYWorldConfig as HYWorldDiTConfig
from fastvideo.configs.models.encoders import SiglipVisionConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig
@dataclass
class HYWorldConfig(Hunyuan15T2V480PConfig):
"""Base configuration for HYWorld pipeline architecture."""
# HYWorldConfig-specific parameters with defaults
dit_config: DiTConfig = field(default_factory=HYWorldDiTConfig)
# SigLIP image encoder for I2V
image_encoder_config: EncoderConfig = field(
default_factory=SiglipVisionConfig)
image_encoder_precision: str = "fp16"
# vae_precision: str = "fp32"
# Text encoding
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("fp16", "fp32"))
def __post_init__(self):
self.vae_config.load_encoder = True
self.vae_config.load_decoder = True
+50
View File
@@ -0,0 +1,50 @@
# SPDX-License-Identifier: Apache-2.0
from collections.abc import Callable
from dataclasses import dataclass, field
import torch
from fastvideo.configs.models import (DiTConfig, EncoderConfig, ModelConfig,
LTX2AudioDecoderConfig, LTX2VocoderConfig,
VAEConfig)
from fastvideo.configs.models.dits import LTX2VideoConfig
from fastvideo.configs.models.encoders import BaseEncoderOutput, LTX2GemmaConfig
from fastvideo.configs.models.vaes import LTX2VAEConfig
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
def ltx2_postprocess_text(outputs: BaseEncoderOutput) -> torch.Tensor:
return outputs.last_hidden_state
@dataclass
class LTX2T2VConfig(PipelineConfig):
"""Configuration for LTX-2 T2V pipeline."""
dit_config: DiTConfig = field(default_factory=LTX2VideoConfig)
vae_config: VAEConfig = field(default_factory=LTX2VAEConfig)
vae_tiling: bool = True
vae_sp: bool = False
text_encoder_configs: tuple[EncoderConfig, ...] = field(
default_factory=lambda: (LTX2GemmaConfig(), ))
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
default_factory=lambda: (preprocess_text, ))
postprocess_text_funcs: tuple[Callable[[BaseEncoderOutput], torch.Tensor],
...] = field(default_factory=lambda:
(ltx2_postprocess_text, ))
dit_precision: str = "bf16"
vae_precision: str = "bf16"
text_encoder_precisions: tuple[str, ...] = field(
default_factory=lambda: ("bf16", ))
audio_decoder_config: ModelConfig = field(
default_factory=LTX2AudioDecoderConfig)
vocoder_config: ModelConfig = field(default_factory=LTX2VocoderConfig)
audio_decoder_precision: str = "bf16"
vocoder_precision: str = "bf16"
def __post_init__(self) -> None:
self.vae_config.load_encoder = False
self.vae_config.load_decoder = True
+14
View File
@@ -9,6 +9,8 @@ from fastvideo.configs.pipelines.cosmos import CosmosConfig
from fastvideo.configs.pipelines.cosmos2_5 import Cosmos25Config
from fastvideo.configs.pipelines.hunyuan import FastHunyuanConfig, HunyuanConfig
from fastvideo.configs.pipelines.hunyuan15 import Hunyuan15T2V480PConfig, Hunyuan15T2V720PConfig
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
from fastvideo.configs.pipelines.turbodiffusion import (
@@ -37,8 +39,10 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
Hunyuan15T2V480PConfig,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
Hunyuan15T2V720PConfig,
"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,
@@ -64,6 +68,9 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
"FastVideo/LongCat-Video-T2V-Diffusers": LongCatT2V480PConfig,
"FastVideo/LongCat-Video-I2V-Diffusers": LongCatT2V480PConfig,
"FastVideo/LongCat-Video-VC-Diffusers": LongCatT2V480PConfig,
# LTX-2 models
"Lightricks/LTX-2": LTX2T2VConfig,
"converted/ltx2_diffusers": LTX2T2VConfig,
# TurboDiffusion models
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers": TurboDiffusionT2V_1_3B_Config,
"loayrashid/TurboWan2.1-T2V-14B-Diffusers": TurboDiffusionT2V_14B_Config,
@@ -83,6 +90,8 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"hyworld":
lambda id: "hyworld" in id.lower(),
"matrixgame":
lambda id: "matrix-game" in id.lower() or "matrixgame" in id.lower(),
"wanpipeline":
@@ -102,6 +111,8 @@ PIPELINE_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "cosmos25" in id.lower(),
"turbodiffusion":
lambda id: "turbodiffusion" in id.lower() or "turbowan" in id.lower(),
"ltx2":
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -116,6 +127,8 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"matrixgame": MatrixGameI2V480PConfig,
"hunyuan15":
Hunyuan15T2V480PConfig, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
"hyworld":
HYWorldConfig, # HYWorld-specific config as fallback for any HYWorld variant
"wanpipeline":
WanT2V480PConfig, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V480PConfig,
@@ -123,6 +136,7 @@ PIPELINE_FALLBACK_CONFIG: dict[str, type[PipelineConfig]] = {
"wancausaldmdpipeline": SelfForcingWanT2V480PConfig,
"stepvideo": StepVideoT2VConfig,
"turbodiffusion": TurboDiffusionT2V_1_3B_Config,
"ltx2": LTX2T2VConfig,
# Other fallbacks by architecture
}
+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)
+5
View File
@@ -22,6 +22,11 @@ class Hunyuan15_480P_SamplingParam(SamplingParam):
negative_prompt: str = ""
def __post_init__(self):
super().__post_init__()
self.sigmas = list(
np.linspace(1.0, 0.0, self.num_inference_steps + 1)[:-1])
@dataclass
class Hunyuan15_720P_SamplingParam(Hunyuan15_480P_SamplingParam):
+26
View File
@@ -0,0 +1,26 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass, field
from fastvideo.configs.sample.base import SamplingParam
import numpy as np
@dataclass
class HYWorld_SamplingParam(SamplingParam):
num_inference_steps: int = 50
num_frames: int = 125
height: int = 480
width: int = 832
fps: int = 24
# Camera trajectory: pose string (e.g., 'w-31' means generating [1 + 31] latents) or JSON file path
pose: str = 'w-31'
guidance_scale: float = 6.0
prompt_attention_mask: list = field(default_factory=list)
negative_attention_mask: list = field(default_factory=list)
sigmas: list[float] | None = field(
default_factory=lambda: list(np.linspace(1.0, 0.0, 50 + 1)[:-1]))
negative_prompt: str = ""
+20
View File
@@ -0,0 +1,20 @@
# SPDX-License-Identifier: Apache-2.0
from dataclasses import dataclass
from fastvideo.configs.sample.base import SamplingParam
@dataclass
class LTX2SamplingParam(SamplingParam):
"""Default sampling parameters for LTX-2 distilled T2V.
"""
seed: int = 10
num_frames: int = 121
height: int = 1024
width: int = 1536
fps: int = 24
num_inference_steps: int = 8
guidance_scale: float = 1.0
# No default negative_prompt for distilled models
negative_prompt: str = ""
+31 -31
View File
@@ -6,10 +6,12 @@ from typing import Any
from fastvideo.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.configs.sample.hunyuan15 import Hunyuan15_480P_SamplingParam, Hunyuan15_720P_SamplingParam
from fastvideo.configs.sample.hyworld import HYWorld_SamplingParam
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.configs.sample.cosmos import Cosmos_Predict2_2B_Video2World_SamplingParam
from fastvideo.configs.sample.cosmos2_5 import Cosmos_Predict2_5_2B_Diffusers_SamplingParam
from fastvideo.configs.sample.ltx2 import LTX2SamplingParam
# isort: off
from fastvideo.configs.sample.wan import (
@@ -40,48 +42,39 @@ from fastvideo.utils import (maybe_download_model_index,
logger = init_logger(__name__)
# Registry maps specific model weights to their config classes
SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"FastVideo/FastHunyuan-diffusers":
FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo":
HunyuanSamplingParam,
"FastVideo/FastHunyuan-diffusers": FastHunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo": HunyuanSamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-480p_t2v":
Hunyuan15_480P_SamplingParam,
"hunyuanvideo-community/HunyuanVideo-1.5-Diffusers-720p_t2v":
Hunyuan15_720P_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers":
StepVideoT2VSamplingParam,
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers": HYWorld_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
# Wan2.1
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers":
WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers":
WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers":
WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers":
WanI2V_14B_720P_SamplingParam,
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V_1_3B_SamplingParam,
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"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,
# Wan2.2
"Wan-AI/Wan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-FullAttn-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers":
Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers":
Wan2_2_I2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-T2V-A14B-Diffusers": Wan2_2_T2V_A14B_SamplingParam,
"Wan-AI/Wan2.2-I2V-A14B-Diffusers": Wan2_2_I2V_A14B_SamplingParam,
# FastWan2.1
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers":
FastWanT2V480P_SamplingParam,
"FastVideo/FastWan2.1-T2V-1.3B-Diffusers": FastWanT2V480P_SamplingParam,
# FastWan2.2
"FastVideo/FastWan2.2-TI2V-5B-Diffusers":
Wan2_2_TI2V_5B_SamplingParam,
"FastVideo/FastWan2.2-TI2V-5B-Diffusers": Wan2_2_TI2V_5B_SamplingParam,
# Causal Self-Forcing Wan2.1
"wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers":
@@ -102,12 +95,9 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
# MatrixGame2.0 models
"FastVideo/Matrix-Game-2.0-Base-Diffusers":
MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers":
MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers":
MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-Base-Diffusers": MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-GTA-Diffusers": MatrixGame2_SamplingParam,
"FastVideo/Matrix-Game-2.0-TempleRun-Diffusers": MatrixGame2_SamplingParam,
# TurboDiffusion models
"loayrashid/TurboWan2.1-T2V-1.3B-Diffusers":
@@ -117,6 +107,10 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"loayrashid/TurboWan2.2-I2V-A14B-Diffusers":
TurboDiffusionI2V_A14B_SamplingParam,
# LTX-2 models
"Lightricks/LTX-2": LTX2SamplingParam,
"FastVideo/LTX2-Distilled-Diffusers": LTX2SamplingParam,
# Add other specific weight variants
}
@@ -126,6 +120,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "hunyuan" in id.lower(),
"hunyuan15":
lambda id: "hunyuan15" in id.lower(),
"hyworld":
lambda id: "hyworld" in id.lower(),
"wanpipeline":
lambda id: "wanpipeline" in id.lower(),
"wanimagetovideo":
@@ -144,6 +140,8 @@ SAMPLING_PARAM_DETECTOR: dict[str, Callable[[str], bool]] = {
lambda id: "cosmos2_5" in id.lower(),
"cosmos":
lambda id: "cosmos" in id.lower() and "2_5" not in id.lower(),
"ltx2":
lambda id: "ltx2" in id.lower() or "ltx-2" in id.lower(),
# Add other pipeline architecture detectors
}
@@ -153,6 +151,8 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
HunyuanSamplingParam, # Base Hunyuan config as fallback for any Hunyuan variant
"hunyuan15":
Hunyuan15_480P_SamplingParam, # Base Hunyuan15 config as fallback for any Hunyuan15 variant
"hyworld":
HYWorld_SamplingParam, # HYWorld-specific config as fallback for any HYWorld variant
"wanpipeline":
WanT2V_1_3B_SamplingParam, # Base Wan config as fallback for any Wan variant
"wanimagetovideo": WanI2V_14B_480P_SamplingParam,
@@ -164,13 +164,13 @@ SAMPLING_FALLBACK_PARAM: dict[str, Any] = {
TurboDiffusionT2V_1_3B_SamplingParam, # Default to T2V for fallback
"cosmos25": Cosmos_Predict2_5_2B_Diffusers_SamplingParam,
"cosmos": Cosmos_Predict2_2B_Video2World_SamplingParam,
"ltx2": LTX2SamplingParam,
# Other fallbacks by architecture
}
def get_sampling_param_cls_for_name(pipeline_name_or_path: str) -> Any | None:
"""Get the appropriate sampling param for specific pretrained weights."""
# First try exact match for specific weights
if pipeline_name_or_path in SAMPLING_PARAM_REGISTRY:
return SAMPLING_PARAM_REGISTRY[pipeline_name_or_path]
+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
@@ -57,6 +57,9 @@ class DistributedAutograd:
ctx.dim = dim
ctx.input_shape = input_.shape
# NCCL all_gather_into_tensor requires contiguous tensors.
if not input_.is_contiguous():
input_ = input_.contiguous()
input_size = input_.size()
output_size = (input_size[0] * world_size, ) + input_size[1:]
output_tensor = torch.empty(output_size,
+100
View File
@@ -18,6 +18,8 @@ import numpy as np
import torch
import torchvision
from einops import rearrange
import shutil
import tempfile
from fastvideo.configs.sample import SamplingParam
from fastvideo.fastvideo_args import FastVideoArgs
@@ -389,6 +391,11 @@ class VideoGenerator:
if batch.save_video:
imageio.mimsave(output_path, frames, fps=batch.fps, format="mp4")
logger.info("Saved video to %s", output_path)
audio = output_batch.extra.get("audio")
audio_sample_rate = output_batch.extra.get("audio_sample_rate")
if (audio is not None and audio_sample_rate is not None and
not self._mux_audio(output_path, audio, audio_sample_rate)):
logger.warning("Audio mux failed; saved video without audio.")
if batch.return_frames:
return frames
@@ -396,6 +403,7 @@ class VideoGenerator:
return {
"samples": samples,
"frames": frames,
"audio": output_batch.extra.get("audio"),
"prompts": prompt,
"size": (target_height, target_width, batch.num_frames),
"generation_time": gen_time,
@@ -405,6 +413,98 @@ class VideoGenerator:
"trajectory_decoded": output_batch.trajectory_decoded,
}
@staticmethod
def _mux_audio(
video_path: str,
audio: torch.Tensor | np.ndarray,
sample_rate: int,
) -> bool:
"""Mux audio into video using PyAV."""
try:
import av
except ImportError:
logger.warning("PyAV not installed; cannot mux audio. "
"Install with: pip install av")
return False
if torch.is_tensor(audio):
audio_np = audio.detach().cpu().float().numpy()
else:
audio_np = np.asarray(audio, dtype=np.float32)
if audio_np.ndim == 1:
audio_np = audio_np[:, None]
elif audio_np.ndim == 2:
if audio_np.shape[0] <= 8 and audio_np.shape[1] > audio_np.shape[0]:
audio_np = audio_np.T
else:
logger.warning("Unexpected audio shape %s; skipping mux.",
audio_np.shape)
return False
audio_np = np.clip(audio_np, -1.0, 1.0)
audio_int16 = (audio_np * 32767.0).astype(np.int16)
num_channels = audio_int16.shape[1]
layout = "stereo" if num_channels == 2 else "mono"
try:
import wave
with tempfile.TemporaryDirectory() as tmpdir:
out_path = os.path.join(tmpdir, "muxed.mp4")
wav_path = os.path.join(tmpdir, "audio.wav")
# Write audio to WAV file
with wave.open(wav_path, "wb") as wav_file:
wav_file.setnchannels(num_channels)
wav_file.setsampwidth(2)
wav_file.setframerate(sample_rate)
wav_file.writeframes(audio_int16.tobytes())
# Open input video and audio
input_video = av.open(video_path)
input_audio = av.open(wav_path)
# Create output with both streams
output = av.open(out_path, mode="w")
# Add video stream (copy codec from input)
in_video_stream = input_video.streams.video[0]
out_video_stream = output.add_stream(
codec_name=in_video_stream.codec_context.name,
rate=in_video_stream.average_rate,
)
out_video_stream.width = in_video_stream.width
out_video_stream.height = in_video_stream.height
out_video_stream.pix_fmt = in_video_stream.pix_fmt
# Add audio stream (AAC)
out_audio_stream = output.add_stream("aac", rate=sample_rate)
out_audio_stream.layout = layout
# Remux video (decode and re-encode to be safe)
for frame in input_video.decode(video=0):
for packet in out_video_stream.encode(frame):
output.mux(packet)
for packet in out_video_stream.encode():
output.mux(packet)
# Encode audio
for frame in input_audio.decode(audio=0):
frame.pts = None # Let encoder assign PTS
for packet in out_audio_stream.encode(frame):
output.mux(packet)
for packet in out_audio_stream.encode():
output.mux(packet)
input_video.close()
input_audio.close()
output.close()
shutil.move(out_path, video_path)
return True
except Exception as e:
logger.warning("Audio mux failed: %s", e)
return False
def set_lora_adapter(self,
lora_nickname: str,
lora_path: str | None = None) -> None:
+82
View File
@@ -166,6 +166,14 @@ class FastVideoArgs:
# Prompt text file for batch processing
prompt_txt: str | None = None
# LTX-2 VAE tiling overrides
ltx2_vae_tiling: bool | None = None
ltx2_vae_spatial_tile_size_in_pixels: int | None = None
ltx2_vae_spatial_tile_overlap_in_pixels: int | None = None
ltx2_vae_temporal_tile_size_in_frames: int | None = None
ltx2_vae_temporal_tile_overlap_in_frames: int | None = None
ltx2_initial_latent_path: str | None = None
# model paths for correct deallocation
model_paths: dict[str, str] = field(default_factory=dict)
model_loaded: dict[str, bool] = field(default_factory=lambda: {
@@ -203,8 +211,44 @@ class FastVideoArgs:
logger.error("Failed to load V-MoBA config from %s: %s",
self.moba_config_path, e)
raise
self._apply_ltx2_vae_overrides()
self.check_fastvideo_args()
def _apply_ltx2_vae_overrides(self) -> None:
if self.pipeline_config is None:
return
vae_config = self.pipeline_config.vae_config
has_any = any(value is not None for value in (
self.ltx2_vae_spatial_tile_size_in_pixels,
self.ltx2_vae_spatial_tile_overlap_in_pixels,
self.ltx2_vae_temporal_tile_size_in_frames,
self.ltx2_vae_temporal_tile_overlap_in_frames,
))
if self.ltx2_vae_tiling is not None and hasattr(self.pipeline_config,
"vae_tiling"):
self.pipeline_config.vae_tiling = self.ltx2_vae_tiling
elif has_any and hasattr(self.pipeline_config, "vae_tiling"):
self.pipeline_config.vae_tiling = True
if hasattr(vae_config, "ltx2_spatial_tile_size_in_pixels"
) and self.ltx2_vae_spatial_tile_size_in_pixels is not None:
vae_config.ltx2_spatial_tile_size_in_pixels = (
self.ltx2_vae_spatial_tile_size_in_pixels)
if hasattr(
vae_config, "ltx2_spatial_tile_overlap_in_pixels"
) and self.ltx2_vae_spatial_tile_overlap_in_pixels is not None:
vae_config.ltx2_spatial_tile_overlap_in_pixels = (
self.ltx2_vae_spatial_tile_overlap_in_pixels)
if hasattr(vae_config, "ltx2_temporal_tile_size_in_frames"
) and self.ltx2_vae_temporal_tile_size_in_frames is not None:
vae_config.ltx2_temporal_tile_size_in_frames = (
self.ltx2_vae_temporal_tile_size_in_frames)
if hasattr(
vae_config, "ltx2_temporal_tile_overlap_in_frames"
) and self.ltx2_vae_temporal_tile_overlap_in_frames is not None:
vae_config.ltx2_temporal_tile_overlap_in_frames = (
self.ltx2_vae_temporal_tile_overlap_in_frames)
@staticmethod
def add_cli_args(parser: FlexibleArgumentParser) -> FlexibleArgumentParser:
# Model and path configuration
@@ -325,6 +369,44 @@ class FastVideoArgs:
"Path to a text file containing prompts (one per line) for batch processing",
)
# LTX-2 VAE tiling overrides
parser.add_argument(
"--ltx2-vae-tiling",
action=StoreBoolean,
default=FastVideoArgs.ltx2_vae_tiling,
help="Enable LTX-2 VAE tiling overrides.",
)
parser.add_argument(
"--ltx2-vae-spatial-tile-size-in-pixels",
type=int,
default=FastVideoArgs.ltx2_vae_spatial_tile_size_in_pixels,
help="LTX-2 VAE spatial tile size in pixels.",
)
parser.add_argument(
"--ltx2-vae-spatial-tile-overlap-in-pixels",
type=int,
default=FastVideoArgs.ltx2_vae_spatial_tile_overlap_in_pixels,
help="LTX-2 VAE spatial tile overlap in pixels.",
)
parser.add_argument(
"--ltx2-vae-temporal-tile-size-in-frames",
type=int,
default=FastVideoArgs.ltx2_vae_temporal_tile_size_in_frames,
help="LTX-2 VAE temporal tile size in frames.",
)
parser.add_argument(
"--ltx2-vae-temporal-tile-overlap-in-frames",
type=int,
default=FastVideoArgs.ltx2_vae_temporal_tile_overlap_in_frames,
help="LTX-2 VAE temporal tile overlap in frames.",
)
parser.add_argument(
"--ltx2-initial-latent-path",
type=str,
default=FastVideoArgs.ltx2_initial_latent_path,
help="Path to load/save a precomputed LTX-2 initial latent.",
)
# LoRA parameters (inference-time adapter loading)
parser.add_argument(
"--lora-path",
+23 -5
View File
@@ -168,9 +168,15 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor | int, shift: torch.Tensor,
scale: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
def forward(
self,
residual: torch.Tensor,
x: torch.Tensor,
gate: torch.Tensor | int,
shift: torch.Tensor,
scale: torch.Tensor,
convert_modulation_dtype: bool = False
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Apply gated residual connection, followed by layernorm and
scale/shift in a single fused operation.
@@ -205,6 +211,11 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
# Apply normalization
normalized = self.norm(residual_output)
if convert_modulation_dtype:
scale = scale.to(normalized.dtype)
shift = shift.to(normalized.dtype)
# Apply scale and shift
if isinstance(scale, torch.Tensor) and scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
@@ -254,14 +265,21 @@ class LayerNormScaleShift(nn.Module):
else:
raise NotImplementedError(f"Norm type {norm_type} not implemented")
def forward(self, x: torch.Tensor, shift: torch.Tensor,
scale: torch.Tensor) -> torch.Tensor:
def forward(self,
x: torch.Tensor,
shift: torch.Tensor,
scale: torch.Tensor,
convert_modulation_dtype: bool = False) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
# x.shape: [batch_size, seq_len, inner_dim]
normalized = self.norm(x)
if self.compute_dtype == torch.float32:
normalized = normalized.float()
if convert_modulation_dtype:
scale = scale.to(normalized.dtype)
shift = shift.to(normalized.dtype)
if scale.dim() == 4:
# scale.shape: [batch_size, num_frames, 1, inner_dim]
num_frames = scale.shape[1]
+9
View File
@@ -0,0 +1,9 @@
# SPDX-License-Identifier: Apache-2.0
from fastvideo.models.audio.ltx2_audio_vae import (
LTX2AudioDecoder,
LTX2AudioEncoder,
LTX2Vocoder,
)
__all__ = ["LTX2AudioEncoder", "LTX2AudioDecoder", "LTX2Vocoder"]
File diff suppressed because it is too large Load Diff
+2
View File
@@ -626,6 +626,8 @@ class HunyuanVideo15Transformer3DModel(CachableDiT):
)
# Final layer processing
if get_sp_world_size() > 1:
hidden_states = hidden_states.contiguous()
hidden_states = sequence_model_parallel_all_gather_with_unpad(hidden_states, original_seq_len, dim=1)
hidden_states = self.final_layer(hidden_states, temb)
# Unpatchify to get original shape
+24
View File
@@ -0,0 +1,24 @@
# SPDX-License-Identifier: Apache-2.0
"""
HYWorld (HY-WorldPlay) model components for FastVideo.
This module provides:
- HYWorldTransformer3DModel: The main transformer model with ProPE and action conditioning
- HYWorldVideoGenerator: Extended VideoGenerator for HYWorld inference
- Utilities for pose processing and camera trajectory generation
"""
from .hyworld import HYWorldTransformer3DModel, HYWorldDoubleStreamBlock
# Inference utilities (used by examples)
from .resolution_utils import (
get_resolution_from_image,
)
__all__ = [
# Model (used by model registry)
"HYWorldTransformer3DModel",
"HYWorldDoubleStreamBlock",
# Inference utilities (used by examples)
"get_resolution_from_image",
]
@@ -0,0 +1,261 @@
# HY-WorldPlay/hyvideo/prope/camera_rope.py
# MIT License
#
# Copyright (c) Authors of
# "PRoPE: Projective Positional Encoding for Multiview Transformers"
#
# Permission is hereby granted, free of charge, to any person obtaining a copy
# of this software and associated documentation files (the "Software"), to deal
# in the Software without restriction, including without limitation the rights
# to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
# copies of the Software, and to permit persons to whom the Software is
# furnished to do so, subject to the following conditions:
#
# The above copyright notice and this permission notice shall be included in all
# copies or substantial portions of the Software.
#
# THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
# IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
# FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
# AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
# LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
# OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
# SOFTWARE.
# How to use PRoPE attention for self-attention:
#
# 1. Easiest way (fast):
# attn = PropeDotProductAttention(...)
# o = attn(q, k, v, viewmats, Ks)
#
# 2. More flexible way (fast):
# attn = PropeDotProductAttention(...)
# attn._precompute_and_cache_apply_fns(viewmats, Ks)
# q = attn._apply_to_q(q)
# k = attn._apply_to_kv(k)
# v = attn._apply_to_kv(v)
# o = F.scaled_dot_product_attention(q, k, v, **kwargs)
# o = attn._apply_to_o(o)
#
# 3. The most flexible way (but slower because repeated computation of RoPE coefficients):
# o = prope_dot_product_attention(q, k, v, ...)
#
# How to use PRoPE attention for cross-attention:
#
# attn_src = PropeDotProductAttention(...)
# attn_tgt = PropeDotProductAttention(...)
# attn_src._precompute_and_cache_apply_fns(viewmats_src, Ks_src)
# attn_tgt._precompute_and_cache_apply_fns(viewmats_tgt, Ks_tgt)
# q_src = attn_src._apply_to_q(q_src)
# k_tgt = attn_tgt._apply_to_kv(k_tgt)
# v_tgt = attn_tgt._apply_to_kv(v_tgt)
# o_src = F.scaled_dot_product_attention(q_src, k_tgt, v_tgt, **kwargs)
# o_src = attn_src._apply_to_o(o_src)
from functools import partial
from typing import Callable, Optional, Tuple, List
import torch
import torch.nn.functional as F
def prope_qkv(
q: torch.Tensor, # (batch, num_heads, seqlen, head_dim)
k: torch.Tensor, # (batch, num_heads, seqlen, head_dim)
v: torch.Tensor, # (batch, num_heads, seqlen, head_dim)
*,
viewmats: torch.Tensor, # (batch, cameras, 4, 4)
Ks: Optional[torch.Tensor], # (batch, cameras, 3, 3)
patches_x: int = None, # How many patches wide is each image?
patches_y: int = None, # How many patches tall is each image?
image_width: int = None, # Width of the image. Used to normalize intrinsics.
image_height: int = None, # Height of the image. Used to normalize intrinsics.
coeffs_x: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
coeffs_y: Optional[Tuple[torch.Tensor, torch.Tensor]] = None,
mask: Optional[torch.Tensor] = None,
kv_cache=None,
is_cache: bool = False,
**kwargs,
) -> torch.Tensor:
"""Similar to torch.nn.functional.scaled_dot_product_attention, but applies PRoPE-style
positional encoding.
Currently, we assume that the sequence length is equal to:
cameras * patches_x * patches_y
And token ordering allows the `(seqlen,)` axis to be reshaped into
`(cameras, patches_x, patches_y)`.
"""
# We're going to assume self-attention: all inputs are the same shape.
(batch, num_heads, seqlen, head_dim) = q.shape
cameras = viewmats.shape[1]
assert q.shape == k.shape == v.shape
assert viewmats.shape == (batch, cameras, 4, 4)
assert Ks is None or Ks.shape == (batch, cameras, 3, 3)
# assert seqlen == cameras * patches_x * patches_y
apply_fn_q, apply_fn_kv, apply_fn_o = _prepare_apply_fns_all_dim(
head_dim=head_dim,
viewmats=viewmats,
Ks=Ks,
patches_x=patches_x,
patches_y=patches_y,
image_width=image_width,
image_height=image_height,
coeffs_x=coeffs_x,
coeffs_y=coeffs_y,
)
query = apply_fn_q(q)
key = apply_fn_kv(k)
value = apply_fn_kv(v)
return query, key, value, apply_fn_o
def _prepare_apply_fns_all_dim(
head_dim: int, # Q/K/V will have this last dimension
viewmats: torch.Tensor, # (batch, cameras, 4, 4)
Ks: Optional[torch.Tensor], # (batch, cameras, 3, 3)
patches_x: int, # How many patches wide is each image?
patches_y: int, # How many patches tall is each image?
image_width: int, # Width of the image. Used to normalize intrinsics.
image_height: int, # Height of the image. Used to normalize intrinsics.
coeffs_x: Optional[torch.Tensor] = None,
coeffs_y: Optional[torch.Tensor] = None,
) -> Tuple[
Callable[[torch.Tensor], torch.Tensor],
Callable[[torch.Tensor], torch.Tensor],
Callable[[torch.Tensor], torch.Tensor],
]:
"""Prepare transforms for PRoPE-style positional encoding."""
device = viewmats.device
(batch, cameras, _, _) = viewmats.shape
# Normalize camera intrinsics.
if Ks is not None:
Ks_norm = torch.zeros_like(Ks)
Ks_norm[..., 0, 0] = Ks[..., 0, 0]
Ks_norm[..., 1, 1] = Ks[..., 1, 1]
Ks_norm[..., 0, 2] = 0
Ks_norm[..., 1, 2] = 0
Ks_norm[..., 2, 2] = 1.0
Ks_norm = Ks_norm.to(dtype=Ks.dtype)
del Ks
# Compute the camera projection matrices we use in PRoPE.
# - K is an `image<-camera` transform.
# - viewmats is a `camera<-world` transform.
# - P = lift(K) @ viewmats is an `image<-world` transform.
P = torch.einsum("...ij,...jk->...ik", _lift_K(Ks_norm), viewmats)
P_T = P.transpose(-1, -2).to(dtype=viewmats.dtype)
P_inv = torch.einsum(
"...ij,...jk->...ik",
_invert_SE3(viewmats),
_lift_K(_invert_K(Ks_norm)),
).to(dtype=viewmats.dtype)
else:
# GTA formula. P is `camera<-world` transform.
P = viewmats
P_T = P.transpose(-1, -2)
P_inv = _invert_SE3(viewmats)
assert P.shape == P_inv.shape == (batch, cameras, 4, 4)
# Block-diagonal transforms to the inputs and outputs of the attention operator.
assert head_dim % 4 == 0
transforms_q = [
(partial(_apply_tiled_projmat, matrix=P_T), head_dim),
]
transforms_kv = [
(partial(_apply_tiled_projmat, matrix=P_inv), head_dim),
]
transforms_o = [
(partial(_apply_tiled_projmat, matrix=P), head_dim),
]
apply_fn_q = partial(_apply_block_diagonal, func_size_pairs=transforms_q)
apply_fn_kv = partial(_apply_block_diagonal, func_size_pairs=transforms_kv)
apply_fn_o = partial(_apply_block_diagonal, func_size_pairs=transforms_o)
return apply_fn_q, apply_fn_kv, apply_fn_o
def _apply_tiled_projmat(
feats: torch.Tensor, # (batch, num_heads, seqlen, feat_dim)
matrix: torch.Tensor, # (batch, cameras, D, D)
) -> torch.Tensor:
"""Apply projection matrix to features."""
# - seqlen => (cameras, patches_x * patches_y)
# - feat_dim => (feat_dim // 4, 4)
(batch, num_heads, seqlen, feat_dim) = feats.shape
cameras = matrix.shape[1]
assert seqlen >= cameras and seqlen % cameras == 0
D = matrix.shape[-1]
assert matrix.shape == (batch, cameras, D, D)
assert feat_dim % D == 0
return torch.einsum(
"bcij,bncpkj->bncpki",
matrix,
feats.reshape((batch, num_heads, cameras, -1, feat_dim // D, D)),
).reshape(feats.shape)
def _apply_block_diagonal(
feats: torch.Tensor, # (..., dim)
func_size_pairs: List[Tuple[Callable[[torch.Tensor], torch.Tensor], int]],
) -> torch.Tensor:
"""Apply a block-diagonal function to an input array.
Each function is specified as a tuple with form:
((Tensor) -> Tensor, int)
Where the integer is the size of the input to the function.
"""
funcs, block_sizes = zip(*func_size_pairs)
assert feats.shape[-1] == sum(block_sizes)
x_blocks = torch.split(feats, block_sizes, dim=-1)
out = torch.cat(
[f(x_block) for f, x_block in zip(funcs, x_blocks)],
dim=-1,
)
assert out.shape == feats.shape, "Input/output shapes should match."
return out
def _invert_SE3(transforms: torch.Tensor) -> torch.Tensor:
"""Invert a 4x4 SE(3) matrix."""
assert transforms.shape[-2:] == (4, 4)
Rinv = transforms[..., :3, :3].transpose(-1, -2)
out = torch.zeros_like(transforms)
out[..., :3, :3] = Rinv
out[..., :3, 3] = -torch.einsum("...ij,...j->...i", Rinv, transforms[..., :3, 3])
out[..., 3, 3] = 1.0
out = out.to(dtype=transforms.dtype)
return out
def _lift_K(Ks: torch.Tensor) -> torch.Tensor:
"""Lift 3x3 matrices to homogeneous 4x4 matrices."""
assert Ks.shape[-2:] == (3, 3)
out = torch.zeros(Ks.shape[:-2] + (4, 4), device=Ks.device)
out[..., :3, :3] = Ks
out[..., 3, 3] = 1.0
out = out.to(dtype=Ks.dtype)
return out
def _invert_K(Ks: torch.Tensor) -> torch.Tensor:
"""Invert 3x3 intrinsics matrices. Assumes no skew."""
assert Ks.shape[-2:] == (3, 3)
out = torch.zeros_like(Ks)
out[..., 0, 0] = 1.0 / Ks[..., 0, 0]
out[..., 1, 1] = 1.0 / Ks[..., 1, 1]
out[..., 0, 2] = -Ks[..., 0, 2] / Ks[..., 0, 0]
out[..., 1, 2] = -Ks[..., 1, 2] / Ks[..., 1, 1]
out[..., 2, 2] = 1.0
out = out.to(dtype=Ks.dtype)
return out
@@ -0,0 +1,76 @@
# HY-WorldPlay/hyvideo/utils/data_utils.py
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
import numpy as np
from PIL import Image
def resize_and_center_crop(image, target_width, target_height):
if target_height == image.shape[0] and target_width == image.shape[1]:
return image
pil_image = Image.fromarray(image)
original_width, original_height = pil_image.size
scale_factor = max(target_width / original_width, target_height / original_height)
resized_width = int(round(original_width * scale_factor))
resized_height = int(round(original_height * scale_factor))
resized_image = pil_image.resize((resized_width, resized_height), Image.LANCZOS)
left = (resized_width - target_width) / 2
top = (resized_height - target_height) / 2
right = (resized_width + target_width) / 2
bottom = (resized_height + target_height) / 2
cropped_image = resized_image.crop((left, top, right, bottom))
return np.array(cropped_image)
def get_closest_ratio(height: float, width: float, ratios: list, buckets: list):
"""
Get the closest ratio in the buckets.
Args:
height (float): video height
width (float): video width
ratios (list): video aspect ratio
buckets (list): buckets generated by `generate_crop_size_list`
Returns:
the closest size in the buckets and the corresponding ratio
"""
aspect_ratio = float(height) / float(width)
ratios_array = np.array(ratios)
closest_ratio_id = np.abs(ratios_array - aspect_ratio).argmin()
closest_size = buckets[closest_ratio_id]
closest_ratio = ratios_array[closest_ratio_id]
return closest_size, closest_ratio
def generate_crop_size_list(base_size=256, patch_size=16, max_ratio=4.0):
num_patches = round((base_size / patch_size) ** 2)
assert max_ratio >= 1.0
crop_size_list = []
wp, hp = num_patches, 1
while wp > 0:
if max(wp, hp) / min(wp, hp) <= max_ratio:
crop_size_list.append((wp * patch_size, hp * patch_size))
if (hp + 1) * wp <= num_patches:
hp += 1
else:
wp -= 1
return crop_size_list
+569
View File
@@ -0,0 +1,569 @@
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
from typing import Any, Optional
import torch
import torch.nn as nn
from einops import rearrange, repeat
from fastvideo.distributed.communication_op import (
sequence_model_parallel_all_gather_with_unpad,
sequence_model_parallel_shard)
from fastvideo.configs.models.dits import HYWorldConfig
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
from fastvideo.layers.visual_embedding import TimestepEmbedder, unpatchify
from fastvideo.models.dits.hunyuanvideo15 import (
MMDoubleStreamBlock,
HunyuanVideo15Transformer3DModel,
)
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.logger import init_logger
from fastvideo.forward_context import set_forward_context
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.distributed.utils import create_attention_mask_for_padding
from .camera_rope import prope_qkv
logger = init_logger(__name__)
class HYWorldDoubleStreamBlock(MMDoubleStreamBlock):
"""
Extended MMDoubleStreamBlock with ProPE (Projective Positional Encoding) support
for camera-aware attention in HY-World/WorldPlay models.
"""
def __init__(
self,
hidden_size: int,
num_attention_heads: int,
mlp_ratio: float,
dtype: torch.dtype | None = None,
supported_attention_backends: tuple[AttentionBackendEnum, ...] | None = None,
prefix: str = "",
):
super().__init__(
hidden_size=hidden_size,
num_attention_heads=num_attention_heads,
mlp_ratio=mlp_ratio,
dtype=dtype,
supported_attention_backends=supported_attention_backends,
prefix=prefix,
)
self.hidden_size = hidden_size
# Add ProPE projection layer for camera-aware attention
self.img_attn_prope_proj = ReplicatedLinear(
hidden_size,
hidden_size,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.img_attn_prope_proj"
)
# Zero-initialize ProPE projection (starts as identity)
nn.init.zeros_(self.img_attn_prope_proj.weight)
if self.img_attn_prope_proj.bias is not None:
nn.init.zeros_(self.img_attn_prope_proj.bias)
def forward(
self,
img: torch.Tensor,
txt: torch.Tensor,
encoder_attention_mask: torch.Tensor,
vec: torch.Tensor,
vec_txt: torch.Tensor,
freqs_cis: tuple,
seq_attention_mask: torch.Tensor,
viewmats: torch.Tensor,
Ks: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor]:
"""
Forward pass with ProPE camera conditioning.
Args:
img: Image/video tokens
txt: Text tokens
encoder_attention_mask: Text attention mask
vec: Modulation vector
freqs_cis: Rotary embedding frequencies
seq_attention_mask: Sequence attention mask
viewmats: Camera view matrices for ProPE [B, T, 4, 4]
Ks: Camera intrinsics for ProPE [B, T, 3, 3]
Returns:
Tuple of (img, txt) output tokens
"""
# Process modulation vectors (inherited from parent)
img_mod_outputs = self.img_mod(vec)
(
img_attn_shift,
img_attn_scale,
img_attn_gate,
img_mlp_shift,
img_mlp_scale,
img_mlp_gate,
) = torch.chunk(img_mod_outputs, 6, dim=-1)
txt_mod_outputs = self.txt_mod(vec_txt)
(
txt_attn_shift,
txt_attn_scale,
txt_attn_gate,
txt_mlp_shift,
txt_mlp_scale,
txt_mlp_gate,
) = torch.chunk(txt_mod_outputs, 6, dim=-1)
# Prepare image for attention using fused operation
img_attn_input = self.img_attn_norm(img, img_attn_shift, img_attn_scale, convert_modulation_dtype=True)
# Get QKV for image
img_qkv, _ = self.img_attn_qkv(img_attn_input)
batch_size, image_seq_len = img_qkv.shape[0], img_qkv.shape[1]
# Split QKV
img_qkv = img_qkv.view(batch_size, image_seq_len, 3,
self.num_attention_heads, -1)
img_q, img_k, img_v = img_qkv[:, :, 0], img_qkv[:, :, 1], img_qkv[:, :,
2]
# Apply QK-Norm if needed
img_q = self.img_attn_q_norm(img_q).to(img_v)
img_k = self.img_attn_k_norm(img_k).to(img_v)
# Prepare text for attention using fused operation
txt_attn_input = self.txt_attn_norm(txt, txt_attn_shift, txt_attn_scale, convert_modulation_dtype=True)
# Get QKV for text
txt_qkv, _ = self.txt_attn_qkv(txt_attn_input)
batch_size, text_seq_len = txt_qkv.shape[0], txt_qkv.shape[1]
# Split QKV
txt_qkv = txt_qkv.view(batch_size, text_seq_len, 3,
self.num_attention_heads, -1)
txt_q, txt_k, txt_v = txt_qkv[:, :, 0], txt_qkv[:, :, 1], txt_qkv[:, :,
2]
# Apply QK-Norm if needed
txt_q = self.txt_attn_q_norm(txt_q).to(txt_q.dtype)
txt_k = self.txt_attn_k_norm(txt_k).to(txt_k.dtype)
# begin hyworld: add camera pose through prope
img_q_prope, img_k_prope, img_v_prope, apply_fn_o = prope_qkv(
img_q.permute(0, 2, 1, 3),
img_k.permute(0, 2, 1, 3),
img_v.permute(0, 2, 1, 3),
viewmats=viewmats,
Ks=Ks,
) # [batch, num_heads, seqlen, head_dim]
img_q_prope = img_q_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
img_k_prope = img_k_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
img_v_prope = img_v_prope.permute(0, 2, 1, 3) # [batch, seqlen, num_heads, head_dim]
# end hyworld
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
# Run distributed attention
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
img_attn, txt_attn = self.attn(img_q, img_k, img_v, txt_q, txt_k, txt_v, freqs_cis=freqs_cis, attention_mask=seq_attention_mask)
# begin hyworld
# attention with prope
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
attn_metadata_prope = FlashAttnMetadataBuilder().build(
current_timestep=0,
attn_mask=encoder_attention_mask,
)
# NOTE: Do NOT pass freqs_cis to prope attention - HY-WorldPlay does not apply RoPE to prope
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata_prope):
img_attn_prope, _ = self.attn(
img_q_prope, img_k_prope, img_v_prope, txt_q, txt_k, txt_v,
freqs_cis=None, attention_mask=seq_attention_mask # No RoPE for prope attention
)
img_attn_prope = img_attn_prope.reshape(batch_size, image_seq_len, -1)
img_attn_prope = rearrange(
img_attn_prope, "B L (H D) -> B H L D", H=self.num_attention_heads
)
img_attn_prope = apply_fn_o(img_attn_prope) # [batch, num_heads, seqlen, head_dim]
img_attn_prope = rearrange(img_attn_prope, "B H L D -> B L (H D)")
# add prope to img_attn
img_attn_out, _ = self.img_attn_proj(img_attn.view(batch_size, image_seq_len, -1))
img_attn_prope_out, _ = self.img_attn_prope_proj(img_attn_prope)
img_attn_out = img_attn_out + img_attn_prope_out
# end hyworld
# Use fused operation for residual connection, normalization, and modulation
img_mlp_input, img_residual = self.img_attn_residual_mlp_norm(
img, img_attn_out, img_attn_gate, img_mlp_shift, img_mlp_scale, convert_modulation_dtype=True)
# Process image MLP
img_mlp_out = self.img_mlp(img_mlp_input)
img = self.img_mlp_residual(img_residual, img_mlp_out, img_mlp_gate)
# Process text attention output
txt_attn_out, _ = self.txt_attn_proj(
txt_attn.reshape(batch_size, text_seq_len, -1))
# Use fused operation for residual connection, normalization, and modulation
txt_mlp_input, txt_residual = self.txt_attn_residual_mlp_norm(
txt, txt_attn_out, txt_attn_gate, txt_mlp_shift, txt_mlp_scale, convert_modulation_dtype=True)
# Process text MLP
txt_mlp_out = self.txt_mlp(txt_mlp_input)
txt = self.txt_mlp_residual(txt_residual, txt_mlp_out, txt_mlp_gate)
return img, txt
class HYWorldFinalLayer(nn.Module):
"""
Final layer for HYWorld that uses modulate() to handle per-token conditioning.
This matches HY-WorldPlay's FinalLayer behavior.
"""
def __init__(self,
hidden_size,
patch_size,
out_channels,
dtype=None,
prefix: str = "") -> None:
super().__init__()
from fastvideo.layers.linear import ReplicatedLinear
from fastvideo.layers.visual_embedding import ModulateProjection
from fastvideo.layers.layernorm import LayerNormScaleShift
self.norm_final = LayerNormScaleShift(
hidden_size,
norm_type="layer",
eps=1e-6,
elementwise_affine=False,
dtype=dtype,
prefix=f"{prefix}.norm_final")
output_dim = patch_size[0] * patch_size[1] * patch_size[2] * out_channels
self.linear = ReplicatedLinear(hidden_size,
output_dim,
bias=True,
params_dtype=dtype,
prefix=f"{prefix}.linear")
# Modulation projection to get shift/scale
self.adaLN_modulation = ModulateProjection(
hidden_size,
factor=2,
act_layer="silu",
dtype=dtype,
prefix=f"{prefix}.adaLN_modulation")
def forward(self, x, c):
shift, scale = self.adaLN_modulation(c).chunk(2, dim=-1)
x = self.norm_final(x, shift, scale, convert_modulation_dtype=True)
x, _ = self.linear(x)
return x
class HYWorldTransformer3DModel(HunyuanVideo15Transformer3DModel):
r"""
HY-World Transformer extending HunyuanVideo15 with:
- ProPE (Projective Positional Encoding) for camera-aware attention
- Action conditioning for interactive video generation
"""
# Class attributes for weight loading - use HYWorld-specific mapping
_fsdp_shard_conditions = HYWorldConfig().arch_config._fsdp_shard_conditions
_compile_conditions = HYWorldConfig().arch_config._compile_conditions
param_names_mapping = HYWorldConfig().arch_config.param_names_mapping
reverse_param_names_mapping = HYWorldConfig().arch_config.reverse_param_names_mapping
def __init__(
self,
config: HYWorldConfig,
hf_config: dict[str, Any],
) -> None:
super().__init__(config=config, hf_config=hf_config)
# Replace double_blocks with HY-World version that supports ProPE
self.double_blocks = nn.ModuleList([
HYWorldDoubleStreamBlock(
hidden_size=self.hidden_size,
num_attention_heads=self.num_attention_heads,
mlp_ratio=config.arch_config.mlp_ratio,
dtype=None,
supported_attention_backends=self._supported_attention_backends,
prefix=f"{config.prefix}.double_blocks.{i}"
)
for i in range(config.arch_config.num_layers)
])
# Add action conditioning module
self.action_in = TimestepEmbedder(
self.hidden_size,
act_layer="silu",
dtype=None,
prefix=f"{config.prefix}.action_in"
)
# Zero-initialize action embedding (starts with no effect)
nn.init.zeros_(self.action_in.mlp.fc_out.weight)
if self.action_in.mlp.fc_out.bias is not None:
nn.init.zeros_(self.action_in.mlp.fc_out.bias)
# Override final_layer with HYWorld version that uses per-token modulate()
self.final_layer = HYWorldFinalLayer(
hidden_size=self.hidden_size,
patch_size=self.patch_size,
out_channels=self.out_channels,
dtype=None,
prefix=f"{config.prefix}.final_layer"
)
self.gradient_checkpointing = False
def forward(
self,
hidden_states: torch.Tensor,
encoder_hidden_states: list[torch.Tensor],
timestep: torch.LongTensor,
encoder_hidden_states_image: list[torch.Tensor],
encoder_attention_mask: list[torch.Tensor],
action: torch.Tensor,
viewmats: torch.Tensor,
Ks: torch.Tensor,
timestep_txt: torch.LongTensor,
guidance: Optional[torch.Tensor] = None,
timestep_r: Optional[torch.LongTensor] = None,
attention_kwargs: Optional[dict[str, Any]] = None,
):
"""
Forward pass with action and camera conditioning.
Args:
action: Action tensor for action conditioning [B, T] or [B*T]
viewmats: Camera view matrices [B, T, 4, 4]
Ks: Camera intrinsics [B, T, 3, 3]
... (other args same as parent)
"""
encoder_hidden_states_image = encoder_hidden_states_image[0]
encoder_hidden_states, encoder_hidden_states_2 = encoder_hidden_states
encoder_attention_mask, encoder_attention_mask_2 = encoder_attention_mask
batch_size, num_channels, num_frames, height, width = hidden_states.shape
p_t, p_h, p_w = self.config.patch_size_t, self.config.patch_size, self.config.patch_size
post_patch_num_frames = num_frames // p_t
post_patch_height = height // p_h
post_patch_width = width // p_w
# 1. RoPE
# Get rotary embeddings
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames, post_patch_height, post_patch_width),
self.hidden_size,
self.num_attention_heads,
self.config.rope_axes_dim,
self.config.rope_theta
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
# NOTE: freqs_cis does NOT need sharding because FastVideo's DistributedAttention
# uses all-to-all to gather the full sequence before applying RoPE
freqs_cis = (freqs_cos, freqs_sin) if freqs_cos is not None else None
# 2. Conditional embeddings
temb = self.time_in(timestep, timestep_r=timestep_r)
temb_txt = self.time_in(timestep_txt, timestep_r=timestep_r)
# Add action conditioning if provided
# temb shape: [B*T, C] where T = num_frames
temb = temb + self.action_in(action.reshape(-1))
# Broadcast timestep embedding for transformer blocks (one per spatial token)
# [B*T, C] -> [B, T*H*W, C] -> [B*T*H*W, C]
temb = repeat(temb, "(B T) C -> B (T H W) C", B=batch_size, H=post_patch_height, W=post_patch_width)
hidden_states = self.img_in(hidden_states)
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
current_seq_len = hidden_states.shape[1]
sp_world_size = get_sp_world_size()
padded_seq_len = current_seq_len * sp_world_size
if padded_seq_len > original_seq_len:
seq_attention_mask = create_attention_mask_for_padding(
seq_len=original_seq_len,
padded_seq_len=padded_seq_len,
batch_size=batch_size,
device=hidden_states.device,
)
else:
seq_attention_mask = None
viewmats_seq = repeat(
viewmats, "B T M N->B (T H W) M N",
H=post_patch_height,
W=post_patch_width
)
Ks_seq = repeat(
Ks, "B T M N->B (T H W) M N",
H=post_patch_height,
W=post_patch_width
)
# Shard viewmats, Ks, and temb for sequence parallelism (shard along sequence dim=1)
# Note that temb in HY1.5 does not need sharding because it is per-sample modulation
# In HYWorld, temb is per-token modulation.
if sp_world_size > 1:
viewmats_seq, _ = sequence_model_parallel_shard(viewmats_seq, dim=1)
Ks_seq, _ = sequence_model_parallel_shard(Ks_seq, dim=1)
temb, _ = sequence_model_parallel_shard(temb, dim=1)
# Rearrange temb after sharding to match expected shape
temb = rearrange(temb, "B S C -> (B S) C")
# qwen text embedding
encoder_hidden_states = self.txt_in(encoder_hidden_states, timestep_txt, encoder_attention_mask)
encoder_hidden_states_cond_emb = self.cond_type_embed(
torch.zeros_like(encoder_hidden_states[:, :, 0], dtype=torch.long)
)
encoder_hidden_states = encoder_hidden_states + encoder_hidden_states_cond_emb
# byt5 text embedding
encoder_hidden_states_2 = self.txt_in_2(encoder_hidden_states_2)
encoder_hidden_states_2_cond_emb = self.cond_type_embed(
torch.ones_like(encoder_hidden_states_2[:, :, 0], dtype=torch.long)
)
encoder_hidden_states_2 = encoder_hidden_states_2 + encoder_hidden_states_2_cond_emb
# image embed
encoder_hidden_states_3 = self.image_embedder(encoder_hidden_states_image)
is_t2v = torch.all(encoder_hidden_states_image == 0)
if is_t2v:
encoder_hidden_states_3 = encoder_hidden_states_3 * 0.0
encoder_attention_mask_3 = torch.zeros(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
else:
encoder_attention_mask_3 = torch.ones(
(batch_size, encoder_hidden_states_3.shape[1]),
dtype=encoder_attention_mask.dtype,
device=encoder_attention_mask.device,
)
encoder_hidden_states_3_cond_emb = self.cond_type_embed(
2
* torch.ones_like(
encoder_hidden_states_3[:, :, 0],
dtype=torch.long,
)
)
encoder_hidden_states_3 = encoder_hidden_states_3 + encoder_hidden_states_3_cond_emb
# reorder and combine text tokens: combine valid tokens first, then padding
encoder_attention_mask = encoder_attention_mask.bool()
encoder_attention_mask_2 = encoder_attention_mask_2.bool()
encoder_attention_mask_3 = encoder_attention_mask_3.bool()
new_encoder_hidden_states = []
new_encoder_attention_mask = []
for text, text_mask, text_2, text_mask_2, image, image_mask in zip(
encoder_hidden_states,
encoder_attention_mask,
encoder_hidden_states_2,
encoder_attention_mask_2,
encoder_hidden_states_3,
encoder_attention_mask_3,
):
# Concatenate: [valid_image, valid_byt5, valid_mllm, invalid_image, invalid_byt5, invalid_mllm]
new_encoder_hidden_states.append(
torch.cat(
[
image[image_mask], # valid image
text_2[text_mask_2], # valid byt5
text[text_mask], # valid mllm
image[~image_mask], # invalid image (zeroed)
torch.zeros_like(text_2[~text_mask_2]), # invalid byt5 (zeroed)
torch.zeros_like(text[~text_mask]), # invalid mllm (zeroed)
],
dim=0,
)
)
# Apply same reordering to attention masks
new_encoder_attention_mask.append(
torch.cat(
[
image_mask[image_mask],
text_mask_2[text_mask_2],
text_mask[text_mask],
image_mask[~image_mask],
text_mask_2[~text_mask_2],
text_mask[~text_mask],
],
dim=0,
)
)
encoder_hidden_states = torch.stack(new_encoder_hidden_states)
encoder_attention_mask = torch.stack(new_encoder_attention_mask)
# 4. Transformer blocks
if torch.is_grad_enabled() and self.gradient_checkpointing:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = self._gradient_checkpointing_func(
block,
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
temb_txt,
freqs_cis,
seq_attention_mask,
viewmats_seq, # hyworld
Ks_seq, # hyworld
)
else:
for block in self.double_blocks:
hidden_states, encoder_hidden_states = block(
hidden_states,
encoder_hidden_states,
encoder_attention_mask,
temb,
temb_txt,
freqs_cis,
seq_attention_mask,
viewmats=viewmats_seq, # hyworld
Ks=Ks_seq, # hyworld
)
# Final layer processing (per-token conditioning via HYWorldFinalLayer)
# Apply final_layer on sharded data first, then gather
hidden_states = self.final_layer(hidden_states, temb)
# Gather the output from all ranks
if get_sp_world_size() > 1:
hidden_states = hidden_states.contiguous()
hidden_states = sequence_model_parallel_all_gather_with_unpad(hidden_states, original_seq_len, dim=1)
# Unpatchify to get original shape
hidden_states = unpatchify(hidden_states, post_patch_num_frames, post_patch_height, post_patch_width, self.patch_size, self.out_channels)
return hidden_states
+679
View File
@@ -0,0 +1,679 @@
# Some functions from HY-WorldPlay/hyvideo/generate.py
"""
Pose processing utilities for HYWorld video generation.
This module provides functions to convert camera poses to model input tensors,
including viewmats, intrinsics, and action labels.
Adapted from HY-WorldPlay: https://github.com/Tencent-Hunyuan/HY-WorldPlay
"""
import json
import numpy as np
import torch
from scipy.spatial.transform import Rotation as R
from typing import Union, Optional
from fastvideo.models.dits.hyworld.trajectory import generate_camera_trajectory_local
# Mapping from one-hot action encoding to single label
mapping = {
(0, 0, 0, 0): 0,
(1, 0, 0, 0): 1,
(0, 1, 0, 0): 2,
(0, 0, 1, 0): 3,
(0, 0, 0, 1): 4,
(1, 0, 1, 0): 5,
(1, 0, 0, 1): 6,
(0, 1, 1, 0): 7,
(0, 1, 0, 1): 8,
}
# Default camera intrinsic matrix (for 1920x1080 resolution)
DEFAULT_INTRINSIC = [
[969.6969696969696, 0.0, 960.0],
[0.0, 969.6969696969696, 540.0],
[0.0, 0.0, 1.0],
]
# Default movement speeds
DEFAULT_FORWARD_SPEED = 0.08 # units per frame
DEFAULT_YAW_SPEED = np.deg2rad(3) # radians per frame
DEFAULT_PITCH_SPEED = np.deg2rad(3) # radians per frame
def one_hot_to_one_dimension(one_hot: torch.Tensor) -> torch.Tensor:
"""Convert one-hot action encoding to single dimension labels."""
return torch.tensor([mapping[tuple(row.tolist())] for row in one_hot])
def parse_pose_string(
pose_string: str,
forward_speed: float = DEFAULT_FORWARD_SPEED,
yaw_speed: float = DEFAULT_YAW_SPEED,
pitch_speed: float = DEFAULT_PITCH_SPEED,
) -> list[dict]:
"""
Parse pose string to motions list.
Format: "w-3, right-0.5, d-4"
- w: forward movement
- s: backward movement
- a: left movement
- d: right movement
- up: pitch up rotation
- down: pitch down rotation
- left: yaw left rotation
- right: yaw right rotation
- number after dash: duration in frames/latents
Args:
pose_string: Comma-separated pose commands
forward_speed: Movement amount per frame
yaw_speed: Yaw rotation amount per frame (radians)
pitch_speed: Pitch rotation amount per frame (radians)
Returns:
List of motion dictionaries for generate_camera_trajectory_local
"""
motions = []
commands = [cmd.strip() for cmd in pose_string.split(",")]
for cmd in commands:
if not cmd:
continue
parts = cmd.split("-")
if len(parts) != 2:
raise ValueError(
f"Invalid pose command: {cmd}. Expected format: 'action-duration'"
)
action = parts[0].strip()
try:
duration = float(parts[1].strip())
except ValueError:
raise ValueError(f"Invalid duration in command: {cmd}")
num_frames = int(duration)
# Parse action and create motion dicts
if action == "w":
# Forward
for _ in range(num_frames):
motions.append({"forward": forward_speed})
elif action == "s":
# Backward
for _ in range(num_frames):
motions.append({"forward": -forward_speed})
elif action == "a":
# Left
for _ in range(num_frames):
motions.append({"right": -forward_speed})
elif action == "d":
# Right
for _ in range(num_frames):
motions.append({"right": forward_speed})
elif action == "up":
# Pitch up
for _ in range(num_frames):
motions.append({"pitch": pitch_speed})
elif action == "down":
# Pitch down
for _ in range(num_frames):
motions.append({"pitch": -pitch_speed})
elif action == "left":
# Yaw left
for _ in range(num_frames):
motions.append({"yaw": -yaw_speed})
elif action == "right":
# Yaw right
for _ in range(num_frames):
motions.append({"yaw": yaw_speed})
else:
raise ValueError(
f"Unknown action: {action}. "
f"Supported actions: w, s, a, d, up, down, left, right"
)
return motions
def pose_string_to_json(
pose_string: str,
intrinsic: Optional[list[list[float]]] = None,
) -> dict:
"""
Convert pose string to pose JSON format.
Args:
pose_string: Comma-separated pose commands
intrinsic: Camera intrinsic matrix (default: DEFAULT_INTRINSIC from trajectory)
Returns:
Dict with frame indices as keys, containing extrinsic and K (intrinsic) matrices
"""
if intrinsic is None:
intrinsic = DEFAULT_INTRINSIC
motions = parse_pose_string(pose_string)
poses = generate_camera_trajectory_local(motions)
pose_json = {}
for i, p in enumerate(poses):
pose_json[str(i)] = {"extrinsic": p.tolist(), "K": intrinsic}
return pose_json
def pose_to_input(
pose_data: Union[str, dict],
latent_num: int,
tps: bool = False,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
"""
Convert pose data to model input tensors.
Args:
pose_data: One of:
- str ending with '.json': path to JSON file
- str: pose string (e.g., "w-3, right-0.5, d-4")
- dict: pose JSON data
latent_num: Number of latents (frames in latent space)
tps: Third-person mode flag
Returns:
Tuple of (viewmats, intrinsics, action_labels):
- viewmats: World-to-camera matrices [T, 4, 4]
- intrinsics: Normalized camera intrinsics [T, 3, 3]
- action_labels: Action labels for each frame [T]
"""
# Handle different input types
if isinstance(pose_data, str):
if pose_data.endswith(".json"):
# Load from JSON file
with open(pose_data, "r") as f:
pose_json = json.load(f)
else:
# Parse pose string
pose_json = pose_string_to_json(pose_data)
elif isinstance(pose_data, dict):
pose_json = pose_data
else:
raise ValueError(
f"Invalid pose_data type: {type(pose_data)}. Expected str or dict."
)
pose_keys = list(pose_json.keys())
latent_num_from_pose = len(pose_keys)
assert latent_num_from_pose == latent_num, (
f"pose corresponds to {latent_num_from_pose * 4 - 3} frames, num_frames "
f"must be set to {latent_num_from_pose * 4 - 3} to ensure alignment."
)
intrinsic_list = []
w2c_list = []
for i in range(latent_num):
t_key = pose_keys[i]
c2w = np.array(pose_json[t_key]["extrinsic"])
w2c = np.linalg.inv(c2w)
w2c_list.append(w2c)
# Normalize intrinsics
intrinsic = np.array(pose_json[t_key]["K"])
intrinsic[0, 0] /= intrinsic[0, 2] * 2
intrinsic[1, 1] /= intrinsic[1, 2] * 2
intrinsic[0, 2] = 0.5
intrinsic[1, 2] = 0.5
intrinsic_list.append(intrinsic)
w2c_list = np.array(w2c_list)
intrinsic_list = torch.tensor(np.array(intrinsic_list))
# Compute relative camera-to-world transforms
c2ws = np.linalg.inv(w2c_list)
C_inv = np.linalg.inv(c2ws[:-1])
relative_c2w = np.zeros_like(c2ws)
relative_c2w[0, ...] = c2ws[0, ...]
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
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 and rotation actions
if move_norms > move_norm_valid: # threshold for movement
if (not tps) or (
tps and abs(rot_angles_deg[1]) < 5e-2 and abs(rot_angles_deg[0]) < 5e-2
):
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
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
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
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 one-hot to single-dimension labels
trans_one_label = one_hot_to_one_dimension(trans_one_hot)
rotate_one_label = one_hot_to_one_dimension(rotate_one_hot)
action_one_label = trans_one_label * 9 + rotate_one_label
return (
torch.as_tensor(w2c_list),
torch.as_tensor(intrinsic_list),
action_one_label,
)
def camera_center_normalization(w2c: np.ndarray) -> np.ndarray:
"""Normalize camera centers relative to the first camera."""
c2w = np.linalg.inv(w2c)
C0_inv = np.linalg.inv(c2w[0])
c2w_aligned = np.array([C0_inv @ C for C in c2w])
return np.linalg.inv(c2w_aligned)
def parse_pose_string_to_actions(pose_string: str, fps: int = 24) -> list[dict]:
"""
Parse pose string to frame-level action timeline.
Format: pose string uses latent counts, where:
- 1 latent = 4 frames
- Special rule: first frame of entire video is extra (frame 0)
- Example: "w-4,d-4" means:
- w-4: forward for frames 0-16 (17 frames total: 1 extra + 4*4)
- d-4: right for frames 17-32 (16 frames total: 4*4)
Args:
pose_string: Comma-separated pose commands (e.g., "w-4,d-4")
fps: Frames per second for video (default: 24)
Returns:
List of dicts with action values for each frame
"""
commands = [cmd.strip() for cmd in pose_string.split(",")]
frame_actions = []
is_first_command = True
for cmd in commands:
if not cmd:
continue
parts = cmd.split("-")
if len(parts) != 2:
raise ValueError(
f"Invalid pose command: {cmd}. Expected format: 'action-duration'"
)
action = parts[0].strip()
try:
num_latents = int(parts[1].strip())
except ValueError:
raise ValueError(f"Invalid duration in command: {cmd}")
# Convert latents to frames
# First command gets 1 extra frame (the special frame 0)
if is_first_command:
num_frames = 1 + num_latents * 4
is_first_command = False
else:
num_frames = num_latents * 4
# Map action to action values
action_values = {"forward": 0, "left": 0, "yaw": 0, "pitch": 0}
if action == "w":
action_values["forward"] = 1
elif action == "s":
action_values["forward"] = -1
elif action == "a":
action_values["left"] = 1
elif action == "d":
action_values["left"] = -1
elif action == "up":
action_values["pitch"] = 1
elif action == "down":
action_values["pitch"] = -1
elif action == "left":
action_values["yaw"] = -1
elif action == "right":
action_values["yaw"] = 1
else:
raise ValueError(f"Unknown action: {action}")
# Add frame-level actions
for _ in range(num_frames):
frame_actions.append(action_values.copy())
return frame_actions
def compute_latent_num(num_frames: int) -> int:
"""
Compute the number of latents from number of frames.
Formula: num_frames = (latent_num - 1) * 4 + 1
So: latent_num = (num_frames - 1) // 4 + 1
Args:
num_frames: Number of video frames
Returns:
Number of latents
"""
return (num_frames - 1) // 4 + 1
def compute_num_frames(latent_num: int) -> int:
"""
Compute the number of frames from number of latents.
Formula: num_frames = (latent_num - 1) * 4 + 1
Args:
latent_num: Number of latents
Returns:
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)
@@ -0,0 +1,64 @@
import numpy as np
from PIL import Image
import requests
from io import BytesIO
from fastvideo.models.dits.hyworld.data_utils import generate_crop_size_list
# Target resolution configs (matching HY-WorldPlay)
TARGET_SIZE_CONFIG = {
"360p": {"bucket_hw_base_size": 480, "bucket_hw_bucket_stride": 16},
"480p": {"bucket_hw_base_size": 640, "bucket_hw_bucket_stride": 16},
"720p": {"bucket_hw_base_size": 960, "bucket_hw_bucket_stride": 16},
"1080p": {"bucket_hw_base_size": 1440, "bucket_hw_bucket_stride": 16},
}
def get_closest_resolution(image_height, image_width, target_resolution="480p"):
"""
Get closest supported resolution for given image dimensions.
Args:
image_height: Height of input image
image_width: Width of input image
target_resolution: Target resolution string (e.g., "480p", "720p")
Returns:
tuple[int, int]: (height, width) of closest supported resolution
"""
config = TARGET_SIZE_CONFIG[target_resolution]
bucket_hw_base_size = config["bucket_hw_base_size"]
bucket_hw_bucket_stride = config["bucket_hw_bucket_stride"]
crop_size_list = generate_crop_size_list(bucket_hw_base_size, bucket_hw_bucket_stride)
aspect_ratios = np.array([round(float(h) / float(w), 5) for h, w in crop_size_list])
# Find closest aspect ratio
image_ratio = float(image_height) / float(image_width)
closest_idx = np.abs(aspect_ratios - image_ratio).argmin()
closest_size = crop_size_list[closest_idx]
return closest_size[0], closest_size[1] # (height, width)
def get_resolution_from_image(image_path, target_resolution="480p"):
"""
Automatically determine resolution from input image.
Args:
image_path: Path or URL to input image
target_resolution: Target resolution tier ("480p", "720p", etc.)
Returns:
tuple[int, int]: (height, width) matching HY-WorldPlay's bucket selection
"""
# Handle URL inputs
if isinstance(image_path, str) and image_path.startswith(('http://', 'https://')):
response = requests.get(image_path)
response.raise_for_status()
img = Image.open(BytesIO(response.content))
else:
img = Image.open(image_path)
img_width, img_height = img.size
return get_closest_resolution(img_height, img_width, target_resolution)
@@ -0,0 +1,316 @@
# HY-WorldPlay/hyvideo/utils/retrieval_context.py
# Licensed under the TENCENT HUNYUAN COMMUNITY LICENSE AGREEMENT (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# https://github.com/Tencent-Hunyuan/HunyuanVideo-1.5/blob/main/LICENSE
#
# Unless and only to the extent required by applicable law, the Tencent Hunyuan works and any
# output and results therefrom are provided "AS IS" without any express or implied warranties of
# any kind including any warranties of title, merchantability, noninfringement, course of dealing,
# usage of trade, or fitness for a particular purpose. You are solely responsible for determining the
# appropriateness of using, reproducing, modifying, performing, displaying or distributing any of
# the Tencent Hunyuan works or outputs and assume any and all risks associated with your or a
# third party's use or distribution of any of the Tencent Hunyuan works or outputs and your exercise
# of rights and permissions under this agreement.
# See the License for the specific language governing permissions and limitations under the License.
import torch
import numpy as np
from typing import List, Tuple, Dict
import math
def generate_points_in_sphere(n_points: int, radius: float) -> torch.Tensor:
"""
Uniformly sample points within a sphere of a specified radius.
:param n_points: The number of points to generate.
:param radius: The radius of the sphere.
:return: A tensor of shape (n_points, 3), representing the (x, y, z) coordinates of the points.
"""
samples_r = torch.rand(n_points)
samples_phi = torch.rand(n_points)
samples_u = torch.rand(n_points)
r = radius * torch.pow(samples_r, 1 / 3)
phi = 2 * math.pi * samples_phi
theta = torch.acos(1 - 2 * samples_u)
# transfer the coordinates from spherical to cartesian
x = r * torch.sin(theta) * torch.cos(phi)
y = r * torch.sin(theta) * torch.sin(phi)
z = r * torch.cos(theta)
points = torch.stack((x, y, z), dim=1)
return points
def rotation_matrix_to_angles(R: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
"""
Estimate the Pitch and Yaw angles from a 3x3 rotation matrix R in the camera coordinate system.
Assumed Camera Coordinate System: X=Right, Y=Up, Z=Backward
(or NeRF style: X=Right, Y=Down, Z=Forward).
Here we adopt the common Computer Vision convention: Z-axis is Forward.
Note: The angle calculations here are directly based on the conventions of your `is_inside_fov_3d_hv` function:
- Yaw/Azimuth angle is in the XZ plane (atan2(x, z)).
- Pitch/Elevation angle is relative to the horizontal plane (atan2(y, sqrt(x^2 + z^2))).
For the third column R[:, 2] of the W2C matrix R (the direction of the World Z-axis in the Camera frame),
this typically corresponds to the direction the camera is looking
(the representation of the world Z-axis in the camera frame).
To simplify and match your `is_inside_fov` logic, we directly use the camera's Z-axis vector:
Camera Z-axis direction in World Frame (Forward Vector): fwd = R_w2c_inv @ [0, 0, 1]
More simply, the Z-axis vector of the C2W matrix is the camera's forward vector in the world frame.
C2W = W2C_inv
"""
R_c2w = R.T
fwd = R_c2w[:, 2]
x = fwd[0]
y = fwd[1]
z = fwd[2]
# compute yaw and pitch
yaw_rad = torch.atan2(x, z)
yaw_deg = yaw_rad * (180.0 / math.pi)
pitch_rad = torch.atan2(y, torch.sqrt(x**2 + z**2))
pitch_deg = pitch_rad * (180.0 / math.pi)
return pitch_deg, yaw_deg
def is_inside_fov_3d_hv(
points: torch.Tensor,
center: torch.Tensor,
center_pitch: torch.Tensor,
center_yaw: torch.Tensor,
fov_half_h: torch.Tensor,
fov_half_v: torch.Tensor,
) -> torch.Tensor:
"""
Check whether points are inside a 3D view frustum defined by a center coordinate, pitch angle, and yaw angle.
:param points: Tensor of shape (N, 3) or (N, B, 3) representing the coordinates of the sampled points.
:param center: Tensor of shape (3) or (B, 3) representing the camera center coordinates.
:param center_pitch: Tensor of shape (1) or (B) representing the pitch angle of center view direction.
:param center_yaw: Tensor of shape (1) or (B) representing the yaw angle of the center view direction.
:param fov_half_h: The horizontal half field-of-view angle (in degrees).
:param fov_half_v: The vertical half field-of-view angle (in degrees).
:return: Boolean tensor of shape (N) or (N, B), indicating whether each point is inside the FOV.
"""
if points.ndim == 2: # N, 3
vectors = points - center[None, :]
C = 1
elif points.ndim == 3: # N, B, 3
vectors = points - center[None, ...]
center_pitch = center_pitch[None, :] if center_pitch.ndim == 1 else center_pitch
center_yaw = center_yaw[None, :] if center_yaw.ndim == 1 else center_yaw
else:
raise ValueError("points' shape should be (N, 3) or (N, B, 3)")
x = vectors[..., 0]
y = vectors[..., 1]
z = vectors[..., 2]
# Calculate the horizontal angle (yaw/azimuth), assuming the Z-axis is forward.
azimuth = torch.atan2(x, z) * (180 / math.pi)
# Calculate the vertical angle (pitch/elevation).
elevation = torch.atan2(y, torch.sqrt(x**2 + z**2)) * (180 / math.pi)
# Calculate the angular difference from the center view direction (handling angle wrapping).
diff_azimuth = azimuth - center_yaw
diff_azimuth = torch.remainder(diff_azimuth + 180, 360) - 180
diff_elevation = elevation - center_pitch
diff_elevation = torch.remainder(diff_elevation + 180, 360) - 180
# Check if within FOV
in_fov_h = diff_azimuth.abs() < fov_half_h
in_fov_v = diff_elevation.abs() < fov_half_v
return in_fov_h & in_fov_v
def calculate_fov_overlap_similarity(
w2c_matrix_curr: torch.Tensor,
w2c_matrix_hist: torch.Tensor,
fov_h_deg: float = 105.0,
fov_v_deg: float = 75.0,
device=None,
points_local=None,
) -> float:
"""
Calculate the Field-of-View (FOV) overlap similarity between two W2C poses using Monte Carlo sampling.
Similarity = (Number of points in Curr_FOV ∩ Hist_FOV) / (Number of points in Curr_FOV).
:param w2c_matrix_curr: The (4, 4) W2C matrix for the current frame.
:param w2c_matrix_hist: The (4, 4) W2C matrix for the historical frame.
:param num_samples, radius, fov_h_deg, fov_v_deg: Sampling and FOV parameters.
:return: The overlap ratio (a float between 0.0 and 1.0).
"""
w2c_matrix_curr = torch.tensor(w2c_matrix_curr, device=device)
w2c_matrix_hist = torch.tensor(w2c_matrix_hist, device=device)
c2w_matrix_curr = torch.linalg.inv(w2c_matrix_curr)
c2w_matrix_hist = torch.linalg.inv(w2c_matrix_hist)
C_inv = w2c_matrix_curr
w2c_matrix_curr = torch.linalg.inv(C_inv @ c2w_matrix_curr)
w2c_matrix_hist = torch.linalg.inv(C_inv @ c2w_matrix_hist)
R_curr, t_curr = w2c_matrix_curr[:3, :3], w2c_matrix_curr[:3, 3]
R_hist, t_hist = w2c_matrix_hist[:3, :3], w2c_matrix_hist[:3, 3]
P_w_curr = -R_curr.T @ t_curr
P_w_hist = -R_hist.T @ t_hist
# pitch, yaw
pitch_curr, yaw_curr = rotation_matrix_to_angles(R_curr)
pitch_hist, yaw_hist = rotation_matrix_to_angles(R_hist)
fov_half_h = torch.tensor(fov_h_deg / 2.0, device=device)
fov_half_v = torch.tensor(fov_v_deg / 2.0, device=device)
# move to P_w_curr (N, 3)
points_world = points_local + P_w_curr[None, :]
in_fov_curr = is_inside_fov_3d_hv(
points_world,
P_w_curr[None, :],
pitch_curr[None],
yaw_curr[None],
fov_half_h,
fov_half_v,
)
# compute based on angle
in_fov_hist = is_inside_fov_3d_hv(
points_world,
P_w_hist[None, :],
pitch_hist[None],
yaw_hist[None],
fov_half_h,
fov_half_v,
)
# compute based on distance
dist = torch.norm(points_world - P_w_hist.reshape(1, -1), dim=1) < 8.0
in_fov_hist = in_fov_hist.bool() & dist.reshape(1, -1).bool()
overlap_count = (in_fov_curr.bool() & in_fov_hist.bool()).sum().float()
fov_curr_count = in_fov_curr.sum().float()
if fov_curr_count == 0:
return 0.0
overlap_ratio = overlap_count / fov_curr_count
return overlap_ratio.item()
def select_aligned_memory_frames(
w2c_list: List[np.ndarray],
current_frame_idx: int,
memory_frames: int,
temporal_context_size: int,
pred_latent_size: int,
pos_weight: float = 1.0,
ang_weight: float = 1.0,
device=None,
points_local=None,
) -> List[int]:
"""
Selects memory and context frames for a given frame based on a four-frame segment distance calculation.
:param w2c_list: List of all N 4x4 World-to-Camera (W2C) extrinsic matrices (np.ndarray).
:param current_frame_idx: The index of the current frame to be processed.
:param memory_frames: The total number of memory frames to select.
:param context_size: The total number of context frames to select.
:param pos_weight: The weight applied to the spatial (position) distance component.
:param ang_weight: The weight applied to the angular distance component.
:return: List[int]: A list containing the indices of the selected memory frames and context frames.
"""
if current_frame_idx <= memory_frames:
return list(range(0, current_frame_idx))
num_total_frames = len(w2c_list)
if current_frame_idx >= num_total_frames or current_frame_idx < 3:
raise ValueError(
f"The current frame index must be within the valid range of w2c_list and must be at least 3."
f"{current_frame_idx}, {len(w2c_list)}"
)
start_context_idx = max(0, current_frame_idx - temporal_context_size)
context_frames_indices = list(range(start_context_idx, current_frame_idx))
candidate_distances = []
query_clip_indices = list(
range(
current_frame_idx,
(
current_frame_idx + pred_latent_size
if current_frame_idx + pred_latent_size <= num_total_frames
else num_total_frames
),
)
)
historical_clip_indices = list(
range(4, current_frame_idx - temporal_context_size, 4)
)
memory_frames_indices = [0, 1, 2, 3] # add the first chunk as context
memory_frames = memory_frames - temporal_context_size
for hist_idx in historical_clip_indices:
total_dist = 0
hist_w2c_1 = w2c_list[hist_idx]
hist_w2c_2 = w2c_list[hist_idx + 2]
for query_idx in query_clip_indices:
dist_1_for_query_idx = 1.0 - calculate_fov_overlap_similarity(
w2c_list[query_idx],
hist_w2c_1,
fov_h_deg=60.0,
fov_v_deg=35.0,
device=device,
points_local=points_local,
)
dist_2_for_query_idx = 1.0 - calculate_fov_overlap_similarity(
w2c_list[query_idx],
hist_w2c_2,
fov_h_deg=60.0,
fov_v_deg=35.0,
device=device,
points_local=points_local,
)
dist_for_query_idx = (dist_1_for_query_idx + dist_2_for_query_idx) / 2.0
total_dist += dist_for_query_idx
final_clip_distance = total_dist / len(query_clip_indices)
candidate_distances.append((hist_idx, final_clip_distance))
candidate_distances.sort(key=lambda x: x[1])
for start_idx, _ in candidate_distances:
# check the memory frame number
if len(memory_frames_indices) >= memory_frames:
break
if start_idx not in memory_frames_indices:
memory_frames_indices.extend(range(start_idx, start_idx + 4))
# exclude the repeated frames
selected_frames_set = set(context_frames_indices)
selected_frames_set.update(memory_frames_indices)
final_selected_frames = sorted(list(selected_frames_set))
return final_selected_frames
+112
View File
@@ -0,0 +1,112 @@
# HY-WorldPlay/hyvideo/generate_custom_trajectory.py
import numpy as np
import json
def rot_x(theta):
c, s = np.cos(theta), np.sin(theta)
return np.array([[1, 0, 0], [0, c, -s], [0, s, c]])
def rot_y(theta):
c, s = np.cos(theta), np.sin(theta)
return np.array([[c, 0, s], [0, 1, 0], [-s, 0, c]])
def rot_z(theta):
c, s = np.cos(theta), np.sin(theta)
return np.array([[c, -s, 0], [s, c, 0], [0, 0, 1]])
def generate_camera_trajectory_local(motions):
"""
motions: list of dict
{"forward": 1.0}, {"yaw": np.pi/2}, {"pitch": np.pi/6}, {"right": 1.0}
- forward: Translation (Forward or Backward)
- yaw: Rotate (Left or Right)
- pitch: Rotate (Up or Down)
- right: Translation (Right or Left)
- third_yaw: Third Perspective Rotate (Left or Right)
"""
poses = []
T = np.eye(4)
poses.append(T.copy())
for move in motions:
# Rotate (Left or Right)
if "yaw" in move:
R = rot_y(move["yaw"])
T[:3, :3] = T[:3, :3] @ R
# Rotate (Up or Down)
if "pitch" in move:
R = rot_x(move["pitch"])
T[:3, :3] = T[:3, :3] @ R
# Translation (Z-direction of the camera's local coordinate system)
forward = move.get("forward", 0.0)
if forward != 0:
local_t = np.array([0, 0, forward])
world_t = T[:3, :3] @ local_t
T[:3, 3] += world_t
# Translation (Z-direction of the camera's local coordinate system)
right = move.get("right", 0.0)
if right != 0:
local_t = np.array([right, 0, 0])
world_t = T[:3, :3] @ local_t
T[:3, 3] += world_t
# Third Perspective Rotate (Left or Right)
third_yaw = move.get("third_yaw", 0.0)
if third_yaw != 0:
theta = -third_yaw
C = np.array([[1, 0.0, 0, 0], [0, 1, 0, 0], [0, 0, 1, -1.0], [0, 0, 0, 1]])
c_origin = C.copy()
# Rotation around the Y-axis
R_y = np.array(
[
[np.cos(theta), 0, np.sin(theta)],
[0, 1, 0],
[-np.sin(theta), 0, np.cos(theta)],
]
)
# Translation
C[:3, :3] = C[:3, :3] @ R_y
C[:3, 3] = R_y @ C[:3, 3]
c_inv = np.linalg.inv(c_origin)
c_relative = c_inv @ C
T = T @ c_relative
poses.append(T.copy())
return poses
if __name__ == "__main__":
# Examples: Forward 0.08 * 16 -> Right Rotate 3 degree * 16
motions = []
for i in range(15):
motions.append({"forward": 0.08})
for i in range(16):
motions.append({"yaw": np.deg2rad(3)})
intrinsic = [
[969.6969696969696, 0.0, 960.0],
[0.0, 969.6969696969696, 540.0],
[0.0, 0.0, 1.0],
]
poses = generate_camera_trajectory_local(motions)
custom_c2w = {}
for i, p in enumerate(poses):
custom_c2w[str(i)] = {"extrinsic": p.tolist(), "K": intrinsic}
json.dump(
custom_c2w,
open("./assets/pose/pose.json", "w"),
indent=4,
ensure_ascii=False,
)
File diff suppressed because it is too large Load Diff
@@ -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)

Some files were not shown because too many files have changed in this diff Show More