Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
bdaf2e2e45 | ||
|
|
f3a7a37312 | ||
|
|
b1abb5c251 | ||
|
|
a80824255e | ||
|
|
caa1c402ba | ||
|
|
38a6bd93d3 | ||
|
|
b867ef7e7c | ||
|
|
3ae58c277a | ||
|
|
0c6862ca55 | ||
|
|
e8c854bcf1 | ||
|
|
06860e96fe |
|
After Width: | Height: | Size: 211 KiB |
|
After Width: | Height: | Size: 117 KiB |
|
Before Width: | Height: | Size: 18 KiB After Width: | Height: | Size: 461 KiB |
|
After Width: | Height: | Size: 210 KiB |
|
Before Width: | Height: | Size: 27 KiB After Width: | Height: | Size: 437 KiB |
@@ -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
|
||||
|
||||
|
||||
@@ -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.
|
||||
@@ -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`
|
||||
@@ -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.
|
||||

|
||||
|
||||
- Use "Additional Filters" to select CUDA 12.8.
|
||||

|
||||
|
||||
- Click "Deploy" and Pick a single A40 or RTX 4090 GPU.
|
||||

|
||||
|
||||
- Select the "FastVideo" or "fastvideo-dev" Pod Template.
|
||||

|
||||
|
||||
- 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.
|
||||

|
||||
|
||||
## 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
|
||||
```
|
||||
|
||||

|
||||
|
||||
After deploying, the pod will take a few minutes to pull the image and start the SSH service.
|
||||
|
||||

|
||||
|
||||
## 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/
|
||||
```
|
||||
|
||||
@@ -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`.
|
||||
|
||||
@@ -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
|
||||
```
|
||||
|
||||

|
||||
|
||||
### 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.
|
||||
|
||||

|
||||
|
||||
## 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 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)
|
||||
|
||||
@@ -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)
|
||||
```
|
||||
@@ -103,7 +103,7 @@ python setup.py install # or pip install -e .
|
||||
|
||||
**`SAGE_ATTN_THREE`**
|
||||
|
||||
[SageAttention 3](https://huggingface.co/jt-zhang/SageAttention3) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
|
||||
[SageAttention 3](https://github.com/thu-ml/SageAttention/tree/main/sageattention3_blackwell) is an advanced attention mechanism that leverages FP4 quantization and Blackwell GPU Tensor Cores for significant performance improvements.
|
||||
|
||||
#### Hardware Requirements
|
||||
|
||||
@@ -113,11 +113,7 @@ python setup.py install # or pip install -e .
|
||||
|
||||
Note that Sage Attention 3 requires `python>=3.13`, `torch>=2.8.0`, `CUDA >=12.8`. If you are using `uv` and using `torch==2.8.0` make sure that `sentencepiece==0.2.1` in the pyproject.toml file.
|
||||
|
||||
To use Sage Attention 3 in FastVideo, first get access to the SageAttention3 code, then move `sageattn/` and `setup.py` to the directory `fastvideo/attention/backends`, then install from using:
|
||||
|
||||
```bash
|
||||
python setup.py install
|
||||
```
|
||||
To use Sage Attention 3 in FastVideo, follow the `README.md` in the linked repository to install the package from source.
|
||||
|
||||
## Teacache
|
||||
|
||||
|
||||
@@ -1,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
|
||||
@@ -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
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
@@ -42,6 +42,7 @@ PIPE_NAME_TO_CONFIG: dict[str, type[PipelineConfig]] = {
|
||||
"FastVideo/HY-WorldPlay-Bidirectional-Diffusers": HYWorldConfig,
|
||||
"Wan-AI/Wan2.1-T2V-1.3B-Diffusers": WanT2V480PConfig,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
|
||||
"weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers": WanI2V480PConfig,
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers": WANV2VConfig,
|
||||
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V480PConfig,
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V720PConfig,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -58,6 +58,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
|
||||
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
|
||||
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"weizhou03/Wan2.1-Game-Fun-1.3B-InP-Diffusers":
|
||||
Wan2_1_Fun_1_3B_InP_SamplingParam,
|
||||
"IRMChen/Wan2.1-Fun-1.3B-Control-Diffusers":
|
||||
Wan2_1_Fun_1_3B_Control_SamplingParam,
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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()),
|
||||
])
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -15,7 +15,7 @@ import torch
|
||||
from scipy.spatial.transform import Rotation as R
|
||||
from typing import Union, Optional
|
||||
|
||||
from .trajectory import generate_camera_trajectory_local
|
||||
from fastvideo.models.dits.hyworld.trajectory import generate_camera_trajectory_local
|
||||
|
||||
|
||||
# Mapping from one-hot action encoding to single label
|
||||
@@ -411,3 +411,269 @@ def compute_num_frames(latent_num: int) -> int:
|
||||
Number of video frames
|
||||
"""
|
||||
return (latent_num - 1) * 4 + 1
|
||||
|
||||
def reformat_keyboard_and_mouse_tensors(keyboard_tensor, mouse_tensor):
|
||||
"""
|
||||
Reformat the keyboard and mouse tensors to the format compatible with HyWorld.
|
||||
"""
|
||||
num_frames = keyboard_tensor.shape[0]
|
||||
assert (num_frames - 1) % 4 == 0, "num_frames must be a multiple of 4"
|
||||
assert mouse_tensor.shape[0] == num_frames, "mouse_tensor must have the same number of frames as keyboard_tensor"
|
||||
keyboard_tensor = keyboard_tensor[1:, :]
|
||||
mouse_tensor = mouse_tensor[1:, :]
|
||||
groups = keyboard_tensor.view(-1, 4, keyboard_tensor.shape[1])
|
||||
assert (groups == groups[:, 0:1]).all(dim=1).all(), "keyboard_tensor must have the same value for each group"
|
||||
groups = mouse_tensor.view(-1, 4, mouse_tensor.shape[1])
|
||||
assert (groups == groups[:, 0:1]).all(dim=1).all(), "mouse_tensor must have the same value for each group"
|
||||
|
||||
return keyboard_tensor[::4], mouse_tensor[::4]
|
||||
|
||||
def process_custom_actions(keyboard_tensor, mouse_tensor, forward_speed=DEFAULT_FORWARD_SPEED):
|
||||
"""
|
||||
Process custom keyboard and mouse tensors into model inputs (viewmats, intrinsics, action_labels).
|
||||
Assumes inputs correspond to each LATENT frame.
|
||||
"""
|
||||
if keyboard_tensor.ndim == 3:
|
||||
keyboard_tensor = keyboard_tensor.squeeze(0)
|
||||
if mouse_tensor.ndim == 3:
|
||||
mouse_tensor = mouse_tensor.squeeze(0)
|
||||
keyboard_tensor, mouse_tensor = reformat_keyboard_and_mouse_tensors(keyboard_tensor, mouse_tensor)
|
||||
|
||||
motions = []
|
||||
|
||||
# 1. Translate tensors to motions for trajectory generation
|
||||
for t in range(keyboard_tensor.shape[0]):
|
||||
frame_motion = {}
|
||||
|
||||
# --- Translation ---
|
||||
# MatrixGame convention: 0:W, 1:S, 2:A, 3:D
|
||||
fwd = 0.0
|
||||
if keyboard_tensor[t, 0] > 0.5: fwd += forward_speed # W
|
||||
if keyboard_tensor[t, 1] > 0.5: fwd -= forward_speed # S
|
||||
if fwd != 0: frame_motion["forward"] = fwd
|
||||
|
||||
rgt = 0.0
|
||||
if keyboard_tensor[t, 2] > 0.5: rgt -= forward_speed # A (Left is negative Right)
|
||||
if keyboard_tensor[t, 3] > 0.5: rgt += forward_speed # D (Right)
|
||||
if rgt != 0: frame_motion["right"] = rgt
|
||||
|
||||
# --- Rotation ---
|
||||
# MatrixGame convention: mouse is [Pitch, Yaw] (or Y, X)
|
||||
# Apply scaling (e.g. to match HyWorld distribution)
|
||||
pitch = mouse_tensor[t, 0].item()
|
||||
yaw = mouse_tensor[t, 1].item()
|
||||
|
||||
if abs(pitch) > 1e-4: frame_motion["pitch"] = pitch
|
||||
if abs(yaw) > 1e-4: frame_motion["yaw"] = yaw
|
||||
|
||||
motions.append(frame_motion)
|
||||
|
||||
# 2. Generate Camera Trajectory
|
||||
# generate_camera_trajectory_local returns T+1 poses (starting at Identity)
|
||||
# We take the first T poses to match the latent count.
|
||||
# Pose 0 is Identity. Pose 1 is Identity + Motion[0].
|
||||
poses = generate_camera_trajectory_local(motions)
|
||||
# poses = np.array(poses[:T])
|
||||
|
||||
# 3. Compute Viewmats (w2c) and Intrinsics
|
||||
w2c_list = []
|
||||
intrinsic_list = []
|
||||
|
||||
# Setup default intrinsic (normalized)
|
||||
K = np.array(DEFAULT_INTRINSIC)
|
||||
K[0, 0] /= K[0, 2] * 2
|
||||
K[1, 1] /= K[1, 2] * 2
|
||||
K[0, 2] = 0.5
|
||||
K[1, 2] = 0.5
|
||||
|
||||
for i in range(len(poses)):
|
||||
c2w = np.array(poses[i])
|
||||
w2c = np.linalg.inv(c2w)
|
||||
w2c_list.append(w2c)
|
||||
intrinsic_list.append(K)
|
||||
|
||||
viewmats = torch.as_tensor(np.array(w2c_list))
|
||||
intrinsics = torch.as_tensor(np.array(intrinsic_list))
|
||||
|
||||
# 4. Generate Action Labels by analyzing the generated trajectory
|
||||
# This ensures consistency with complex simultaneous movements, exactly as pose_to_input does.
|
||||
|
||||
# Calculate relative camera-to-world transforms
|
||||
# c2ws = inverse(viewmats)
|
||||
c2ws = np.linalg.inv(np.array(w2c_list))
|
||||
|
||||
# Calculate relative movement between frames
|
||||
# relative_c2w[i] = inv(c2ws[i-1]) @ c2ws[i]
|
||||
C_inv = np.linalg.inv(c2ws[:-1])
|
||||
relative_c2w = np.zeros_like(c2ws)
|
||||
relative_c2w[0, ...] = c2ws[0, ...] # First is anchor
|
||||
relative_c2w[1:, ...] = C_inv @ c2ws[1:, ...]
|
||||
|
||||
# Initialize one-hot action encodings
|
||||
trans_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
|
||||
rotate_one_hot = np.zeros((relative_c2w.shape[0], 4), dtype=np.int32)
|
||||
|
||||
move_norm_valid = 0.0001
|
||||
|
||||
# Skip index 0 (anchor/identity)
|
||||
for i in range(1, relative_c2w.shape[0]):
|
||||
move_dirs = relative_c2w[i, :3, 3] # direction vector
|
||||
move_norms = np.linalg.norm(move_dirs)
|
||||
|
||||
if move_norms > move_norm_valid: # threshold for movement
|
||||
move_norm_dirs = move_dirs / move_norms
|
||||
angles_rad = np.arccos(move_norm_dirs.clip(-1.0, 1.0))
|
||||
trans_angles_deg = angles_rad * (180.0 / np.pi) # convert to degrees
|
||||
else:
|
||||
trans_angles_deg = np.zeros(3)
|
||||
|
||||
R_rel = relative_c2w[i, :3, :3]
|
||||
r = R.from_matrix(R_rel)
|
||||
rot_angles_deg = r.as_euler("xyz", degrees=True)
|
||||
|
||||
# Determine movement actions based on trajectory
|
||||
# Note: HyWorld logic checks if rotation is small before assigning translation labels
|
||||
# to avoid ambiguity in TPS mode, but here we generally want to capture the dominant movement.
|
||||
tps = False # Default assumption, can be made an arg if needed
|
||||
|
||||
if move_norms > move_norm_valid:
|
||||
if (not tps) or (
|
||||
tps and abs(rot_angles_deg[1]) < 5e-2 and abs(rot_angles_deg[0]) < 5e-2
|
||||
):
|
||||
# Z-axis (Forward/Back)
|
||||
if trans_angles_deg[2] < 60:
|
||||
trans_one_hot[i, 0] = 1 # forward
|
||||
elif trans_angles_deg[2] > 120:
|
||||
trans_one_hot[i, 1] = 1 # backward
|
||||
|
||||
# X-axis (Right/Left)
|
||||
if trans_angles_deg[0] < 60:
|
||||
trans_one_hot[i, 2] = 1 # right
|
||||
elif trans_angles_deg[0] > 120:
|
||||
trans_one_hot[i, 3] = 1 # left
|
||||
|
||||
# Determine rotation actions
|
||||
# Y-axis (Yaw)
|
||||
if rot_angles_deg[1] > 5e-2:
|
||||
rotate_one_hot[i, 0] = 1 # right
|
||||
elif rot_angles_deg[1] < -5e-2:
|
||||
rotate_one_hot[i, 1] = 1 # left
|
||||
|
||||
# X-axis (Pitch)
|
||||
if rot_angles_deg[0] > 5e-2:
|
||||
rotate_one_hot[i, 2] = 1 # up
|
||||
elif rot_angles_deg[0] < -5e-2:
|
||||
rotate_one_hot[i, 3] = 1 # down
|
||||
|
||||
trans_one_hot = torch.tensor(trans_one_hot)
|
||||
rotate_one_hot = torch.tensor(rotate_one_hot)
|
||||
|
||||
# Convert to single labels
|
||||
trans_label = one_hot_to_one_dimension(trans_one_hot)
|
||||
rotate_label = one_hot_to_one_dimension(rotate_one_hot)
|
||||
action_labels = trans_label * 9 + rotate_label
|
||||
|
||||
return viewmats, intrinsics, action_labels
|
||||
|
||||
if __name__ == "__main__":
|
||||
print("Running comparison test between process_custom_actions and pose_to_input...")
|
||||
|
||||
def test_process_custom_actions(pose_string: str, keyboard: torch.Tensor, mouse: torch.Tensor, latent_num: int):
|
||||
# Run process_custom_actions
|
||||
# Note: We need to pass float tensors
|
||||
print("Running process_custom_actions...")
|
||||
viewmats_1, intrinsics_1, labels_1 = process_custom_actions(
|
||||
keyboard, mouse
|
||||
)
|
||||
|
||||
print(f"Running pose_to_input with string: '{pose_string}'...")
|
||||
viewmats_2, intrinsics_2, labels_2 = pose_to_input(
|
||||
pose_string, latent_num=latent_num
|
||||
)
|
||||
|
||||
# print(f"Viewmats: {viewmats_1} vs \n {viewmats_2}")
|
||||
# print(f"Intrinsics: {intrinsics_1} vs \n {intrinsics_2}")
|
||||
# print(f"Labels: {labels_1} vs \n {labels_2}")
|
||||
# 3. Compare Results
|
||||
print("\nComparison Results:")
|
||||
|
||||
# Check Shapes
|
||||
print(f"Shapes (Viewmats): {viewmats_1.shape} vs {viewmats_2.shape}")
|
||||
assert viewmats_1.shape == viewmats_2.shape, "Shape mismatch for viewmats"
|
||||
|
||||
# Check Values
|
||||
# Viewmats
|
||||
diff_viewmats = (viewmats_1 - viewmats_2).abs().max().item()
|
||||
print(f"Max difference in Viewmats: {diff_viewmats}")
|
||||
if diff_viewmats < 1e-5:
|
||||
print("✅ Viewmats match.")
|
||||
else:
|
||||
print("❌ Viewmats mismatch.")
|
||||
|
||||
# Check intrinsics
|
||||
diff_intrinsics = (intrinsics_1 - intrinsics_2).abs().max().item()
|
||||
print(f"Max difference in Intrinsics: {diff_intrinsics}")
|
||||
if diff_intrinsics < 1e-5:
|
||||
print("✅ Intrinsics match.")
|
||||
else:
|
||||
print("❌ Intrinsics mismatch.")
|
||||
|
||||
# Check labels
|
||||
diff_labels = (labels_1 - labels_2).abs().max().item()
|
||||
print(f"Max difference in Labels: {diff_labels}")
|
||||
if diff_labels < 1e-5:
|
||||
print("✅ Labels match.")
|
||||
else:
|
||||
print("❌ Labels mismatch.")
|
||||
|
||||
print("All checks passed.")
|
||||
|
||||
# Define shared parameters
|
||||
|
||||
latent_num = 13
|
||||
pose_string = "w-2, a-3, s-1, d-6"
|
||||
|
||||
num_frames = 4 * (latent_num - 1) + 1
|
||||
keyboard = torch.zeros((num_frames, 6))
|
||||
mouse = torch.zeros((num_frames, 2))
|
||||
|
||||
# Frame 0 is ignored/start
|
||||
# Frames 1-8: Press W (index 0)
|
||||
keyboard[1:9, 0] = 1.0
|
||||
# Frames 9-20: Press A (index 2)
|
||||
keyboard[9:21, 2] = 1.0
|
||||
# Frames 21-24: Press S (index 1)
|
||||
keyboard[21:25, 1] = 1.0
|
||||
# Frames 25-48: Press D (index 3)
|
||||
keyboard[25:49, 3] = 1.0
|
||||
|
||||
test_process_custom_actions(pose_string, keyboard, mouse, latent_num)
|
||||
|
||||
# Test keyboard AND mouse
|
||||
latent_num = 25
|
||||
pose_string = "w-2, up-2, a-3, down-4, s-1, left-2, d-6, right-4"
|
||||
|
||||
num_frames = 4 * (latent_num - 1) + 1
|
||||
keyboard = torch.zeros((num_frames, 6))
|
||||
mouse = torch.zeros((num_frames, 2))
|
||||
|
||||
# Frame 0 is ignored/start
|
||||
# Frames 1-8: Press W (index 0)
|
||||
keyboard[1:9, 0] = 1.0
|
||||
# Frames 17-28: Press A (index 2)
|
||||
keyboard[17:29, 2] = 1.0
|
||||
# Frames 45-48: Press S (index 1)
|
||||
keyboard[45:49, 1] = 1.0
|
||||
# Frames 57-80: Press D (index 3)
|
||||
keyboard[57:81, 3] = 1.0
|
||||
|
||||
# Frames 9-16: Press Up (index 4)
|
||||
mouse[9:17, 0] = DEFAULT_PITCH_SPEED
|
||||
# Frames 25-32: Press Down (index 5)
|
||||
mouse[29:45, 0] = -DEFAULT_PITCH_SPEED
|
||||
# Frames 41-48: Press Left (index 6)
|
||||
mouse[49:57, 1] = -DEFAULT_YAW_SPEED
|
||||
# Frames 57-64: Press Right (index 7)
|
||||
mouse[81:, 1] = DEFAULT_YAW_SPEED
|
||||
|
||||
test_process_custom_actions(pose_string, keyboard, mouse, latent_num)
|
||||
|
||||
@@ -257,6 +257,17 @@ class ActionModule(nn.Module):
|
||||
'''
|
||||
assert use_rope_keyboard
|
||||
|
||||
target_device = x.device
|
||||
target_dtype = x.dtype
|
||||
if mouse_condition is not None:
|
||||
mouse_condition = mouse_condition.to(device=target_device,
|
||||
dtype=target_dtype)
|
||||
if keyboard_condition is not None:
|
||||
keyboard_condition = keyboard_condition.to(
|
||||
device=target_device, dtype=target_dtype)
|
||||
else:
|
||||
return x
|
||||
|
||||
B, N_frames, C = keyboard_condition.shape
|
||||
assert tt*th*tw == x.shape[1]
|
||||
assert ((N_frames - 1) + self.vae_time_compression_ratio) % self.vae_time_compression_ratio == 0
|
||||
@@ -272,7 +283,9 @@ class ActionModule(nn.Module):
|
||||
# Defined freqs_cis early so it's available for both mouse and keyboard
|
||||
freqs_cis = (self._freqs_cos, self._freqs_sin)
|
||||
|
||||
assert (N_feats == tt and ((is_causal and kv_cache_mouse is None) or not is_causal)) or ((N_frames - 1) // self.vae_time_compression_ratio + 1 == start_frame + num_frame_per_block and is_causal)
|
||||
if is_causal:
|
||||
assert (N_feats == tt and kv_cache_mouse is None) or ((N_frames - 1) // self.vae_time_compression_ratio + 1 == start_frame + num_frame_per_block)
|
||||
# For non-causal (training), we trust that the caller provides correctly shaped inputs
|
||||
|
||||
if self.enable_mouse and mouse_condition is not None:
|
||||
hidden_states = rearrange(x, "B (T S) C -> (B S) T C", T=tt, S=th*tw) # 65*272*480 -> 17*(272//16)*(480//16) -> 8670
|
||||
@@ -289,13 +302,15 @@ class ActionModule(nn.Module):
|
||||
mouse_condition = mouse_condition[:, self.vae_time_compression_ratio*(N_feats - num_frame_per_block - self.windows_size) + pad_t:, :]
|
||||
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(num_frame_per_block)]
|
||||
else:
|
||||
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(N_feats)]
|
||||
local_num_frames = tt
|
||||
group_mouse = [mouse_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(local_num_frames)]
|
||||
|
||||
group_mouse = torch.stack(group_mouse, dim = 1)
|
||||
actual_num_frames = group_mouse.shape[1] # Use actual stacked frame count
|
||||
|
||||
S = th * tw
|
||||
group_mouse = group_mouse.unsqueeze(-1).expand(B, num_frame_per_block, pad_t, C, S)
|
||||
group_mouse = group_mouse.permute(0, 4, 1, 2, 3).reshape(B * S, num_frame_per_block, pad_t * C)
|
||||
group_mouse = group_mouse.unsqueeze(-1).expand(B, actual_num_frames, pad_t, C, S)
|
||||
group_mouse = group_mouse.permute(0, 4, 1, 2, 3).reshape(B * S, actual_num_frames, pad_t * C)
|
||||
|
||||
group_mouse = torch.cat([hidden_states, group_mouse], dim = -1)
|
||||
group_mouse = self.mouse_mlp(group_mouse)
|
||||
@@ -314,7 +329,7 @@ class ActionModule(nn.Module):
|
||||
## TODO: adding cache here
|
||||
if is_causal:
|
||||
if kv_cache_mouse is None:
|
||||
assert q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0 # == 880, f"{q.shape[0]},{k.shape[0]}"
|
||||
assert q.shape[0] == k.shape[0] and q.shape[0] % S == 0 # == 880, f"{q.shape[0]},{k.shape[0]}"
|
||||
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
|
||||
padded_q = torch.cat(
|
||||
[q,
|
||||
@@ -395,7 +410,8 @@ class ActionModule(nn.Module):
|
||||
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(num_frame_per_block)]
|
||||
else:
|
||||
keyboard_condition = self.keyboard_embed(keyboard_condition)
|
||||
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(N_feats)]
|
||||
local_num_frames = tt
|
||||
group_keyboard = [keyboard_condition[:, self.vae_time_compression_ratio*(i - self.windows_size) + pad_t:i * self.vae_time_compression_ratio + pad_t,:] for i in range(local_num_frames)]
|
||||
group_keyboard = torch.stack(group_keyboard, dim = 1) # B F RW C
|
||||
group_keyboard = group_keyboard.reshape(shape=(group_keyboard.shape[0],group_keyboard.shape[1],-1))
|
||||
# apply cross attn
|
||||
@@ -415,7 +431,7 @@ class ActionModule(nn.Module):
|
||||
q = self.key_attn_q_norm(q).to(v)
|
||||
k = self.key_attn_k_norm(k).to(v)
|
||||
S = th * tw
|
||||
assert S == 880
|
||||
# assert S == 880
|
||||
# position embed
|
||||
if use_rope_keyboard:
|
||||
B, TS, H, D = q.shape
|
||||
@@ -424,13 +440,13 @@ class ActionModule(nn.Module):
|
||||
q, k = _apply_rotary_emb_qk(q, k, freqs_cis[0], freqs_cis[1], start_offset=start_frame)
|
||||
|
||||
k1, k2, k3, k4 = k.shape
|
||||
k = k.expand(S, k2, k3, k4)
|
||||
v = v.expand(S, k2, k3, k4)
|
||||
k = k.repeat_interleave(S, dim=0)
|
||||
v = v.repeat_interleave(S, dim=0)
|
||||
|
||||
|
||||
if is_causal:
|
||||
if kv_cache_keyboard is None:
|
||||
assert q.shape[0] == k.shape[0] and q.shape[0] % 880 == 0
|
||||
assert q.shape[0] == k.shape[0] and q.shape[0] % S == 0
|
||||
|
||||
padded_length = math.ceil(q.shape[1] / 32) * 32 - q.shape[1]
|
||||
padded_q = torch.cat(
|
||||
@@ -480,7 +496,7 @@ class ActionModule(nn.Module):
|
||||
else:
|
||||
local_end_index = kv_cache_keyboard["local_end_index"].item() + current_end - kv_cache_keyboard["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
assert k.shape[0] == 880 # BS == 1 or the cache should not be saved/ load method should be modified
|
||||
assert k.shape[0] == S # BS == 1 or the cache should not be saved/ load method should be modified
|
||||
kv_cache_keyboard["k"][:, local_start_index:local_end_index] = k[:1]
|
||||
kv_cache_keyboard["v"][:, local_start_index:local_end_index] = v[:1]
|
||||
|
||||
@@ -1096,4 +1096,4 @@ class CausalMatrixGameWanModel(BaseDiT):
|
||||
if block.action_model.proj_keyboard.bias is not None:
|
||||
nn.init.zeros_(block.action_model.proj_keyboard.bias)
|
||||
except AttributeError:
|
||||
pass
|
||||
pass
|
||||
@@ -61,14 +61,17 @@ class MatrixGameTimeImageEmbedding(nn.Module):
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
|
||||
# MatrixGame does not use text embeddings, so we ignore encoder_hidden_states
|
||||
# and return None for the text embedding part
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
assert self.image_embedder is not None
|
||||
encoder_hidden_states_image = self.image_embedder(
|
||||
encoder_hidden_states_image)
|
||||
|
||||
return temb, timestep_proj, None, encoder_hidden_states_image
|
||||
encoder_hidden_states = torch.zeros((timestep.shape[0], 0, temb.shape[-1]),
|
||||
device=temb.device,
|
||||
dtype=temb.dtype)
|
||||
|
||||
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
|
||||
|
||||
|
||||
class MatrixGameCrossAttention(WanSelfAttention):
|
||||
@@ -313,6 +316,7 @@ class MatrixGameWanModel(BaseDiT):
|
||||
inner_dim = config.num_attention_heads * config.attention_head_dim
|
||||
self.hidden_size = config.hidden_size
|
||||
self.num_attention_heads = config.num_attention_heads
|
||||
self.num_channels_latents = config.out_channels
|
||||
self.patch_size = config.patch_size
|
||||
|
||||
# 1. Patch & position embedding
|
||||
@@ -395,7 +399,8 @@ class MatrixGameWanModel(BaseDiT):
|
||||
self.num_attention_heads,
|
||||
rope_dim_list,
|
||||
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
|
||||
rope_theta=10000)
|
||||
rope_theta=10000,
|
||||
do_sp_sharding=True)
|
||||
freqs_cos = freqs_cos.to(hidden_states.device)
|
||||
freqs_sin = freqs_sin.to(hidden_states.device)
|
||||
freqs_cis = (freqs_cos.float(),
|
||||
@@ -157,8 +157,8 @@ def load_initial_image(image_path: str = None) -> Image.Image:
|
||||
return Image.new("RGB", (640, 352), (128, 128, 128))
|
||||
|
||||
def create_action_presets(num_frames: int, keyboard_dim: int = 4, seed: int = None):
|
||||
if keyboard_dim not in (2, 4, 7):
|
||||
raise ValueError(f"keyboard_dim must be 2, 4, or 7, got {keyboard_dim}")
|
||||
if keyboard_dim not in (2, 4, 6, 7):
|
||||
raise ValueError(f"keyboard_dim must be 2, 4, 6, or 7, got {keyboard_dim}")
|
||||
if num_frames % 4 != 1:
|
||||
raise ValueError("Matrix-Game conditioning expects num_frames to be 4k+1.")
|
||||
|
||||
@@ -181,6 +181,11 @@ def create_action_presets(num_frames: int, keyboard_dim: int = 4, seed: int = No
|
||||
actions_double_action = []
|
||||
actions_single_camera = ["camera_l", "camera_r"]
|
||||
keyboard_idx = {"forward": 0, "back": 1}
|
||||
elif keyboard_dim == 6:
|
||||
actions_single_action = ["forward", "back", "left", "right"]
|
||||
actions_double_action = ["forward_left", "forward_right"]
|
||||
actions_single_camera = ["camera_l", "camera_r"]
|
||||
keyboard_idx = {"forward": 0, "back": 1, "left": 2, "right": 3, "t1": 4, "t2": 5}
|
||||
else: # keyboard_dim == 7
|
||||
# Temple Run model: still, w, s, left, right, a, d (no mouse)
|
||||
actions_single_action = ["forward", "back", "left", "right"]
|
||||
@@ -0,0 +1,5 @@
|
||||
from .model import WanGameActionTransformer3DModel
|
||||
|
||||
__all__ = [
|
||||
"WanGameActionTransformer3DModel",
|
||||
]
|
||||
@@ -0,0 +1,231 @@
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.layers.visual_embedding import TimestepEmbedder, ModulateProjection, timestep_embedding
|
||||
from fastvideo.platforms import AttentionBackendEnum
|
||||
from fastvideo.attention import DistributedAttention
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.models.dits.wanvideo import WanImageEmbedding
|
||||
|
||||
from fastvideo.models.dits.hyworld.camera_rope import prope_qkv
|
||||
from fastvideo.layers.rotary_embedding import _apply_rotary_emb
|
||||
from fastvideo.layers.mlp import MLP
|
||||
|
||||
class WanGameActionTimeImageEmbedding(nn.Module):
|
||||
def __init__(
|
||||
self,
|
||||
dim: int,
|
||||
time_freq_dim: int,
|
||||
image_embed_dim: int | None = None,
|
||||
):
|
||||
super().__init__()
|
||||
|
||||
self.time_freq_dim = time_freq_dim
|
||||
self.time_embedder = TimestepEmbedder(
|
||||
dim, frequency_embedding_size=time_freq_dim, act_layer="silu")
|
||||
self.time_modulation = ModulateProjection(dim,
|
||||
factor=6,
|
||||
act_layer="silu")
|
||||
|
||||
self.image_embedder = None
|
||||
if image_embed_dim is not None:
|
||||
self.image_embedder = WanImageEmbedding(image_embed_dim, dim)
|
||||
|
||||
self.action_embedder = MLP(
|
||||
time_freq_dim,
|
||||
dim,
|
||||
dim,
|
||||
bias=True,
|
||||
act_type="silu"
|
||||
)
|
||||
# Initialize with zeros for residual-like behavior
|
||||
nn.init.zeros_(self.action_embedder.fc_out.weight)
|
||||
if self.action_embedder.fc_out.bias is not None:
|
||||
nn.init.zeros_(self.action_embedder.fc_out.bias)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
timestep: torch.Tensor,
|
||||
action: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor, # Kept for interface compatibility
|
||||
encoder_hidden_states_image: torch.Tensor | None = None,
|
||||
timestep_seq_len: int | None = None,
|
||||
):
|
||||
temb = self.time_embedder(timestep, timestep_seq_len)
|
||||
|
||||
action_emb = timestep_embedding(action.flatten(), self.time_freq_dim)
|
||||
action_embedder_dtype = next(iter(self.action_embedder.parameters())).dtype
|
||||
if (
|
||||
action_emb.dtype != action_embedder_dtype
|
||||
and action_embedder_dtype != torch.int8
|
||||
):
|
||||
action_emb = action_emb.to(action_embedder_dtype)
|
||||
action_emb = self.action_embedder(action_emb).type_as(temb)
|
||||
temb = temb + action_emb
|
||||
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
|
||||
# MatrixGame does not use text embeddings, so we ignore encoder_hidden_states
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
assert self.image_embedder is not None
|
||||
encoder_hidden_states_image = self.image_embedder(
|
||||
encoder_hidden_states_image)
|
||||
|
||||
encoder_hidden_states = torch.zeros((timestep.shape[0], 0, temb.shape[-1]),
|
||||
device=temb.device,
|
||||
dtype=temb.dtype)
|
||||
|
||||
return temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image
|
||||
|
||||
class WanGameActionSelfAttention(nn.Module):
|
||||
"""
|
||||
Self-attention module with support for:
|
||||
- Standard RoPE-based attention
|
||||
- Camera PRoPE-based attention (when viewmats and Ks are provided)
|
||||
- KV caching for autoregressive generation
|
||||
"""
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
num_heads: int,
|
||||
local_attn_size: int = -1,
|
||||
sink_size: int = 0,
|
||||
qk_norm=True,
|
||||
eps=1e-6) -> None:
|
||||
assert dim % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = dim
|
||||
self.num_heads = num_heads
|
||||
self.head_dim = dim // num_heads
|
||||
self.local_attn_size = local_attn_size
|
||||
self.sink_size = sink_size
|
||||
self.qk_norm = qk_norm
|
||||
self.eps = eps
|
||||
self.max_attention_size = 32760 if local_attn_size == -1 else local_attn_size * 1560
|
||||
|
||||
# Scaled dot product attention (using DistributedAttention for SP support)
|
||||
self.attn = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=self.head_dim,
|
||||
softmax_scale=None,
|
||||
causal=False,
|
||||
supported_attention_backends=(AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA))
|
||||
|
||||
def forward(self,
|
||||
q: torch.Tensor,
|
||||
k: torch.Tensor,
|
||||
v: torch.Tensor,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor],
|
||||
kv_cache: dict | None = None,
|
||||
current_start: int = 0,
|
||||
cache_start: int | None = None,
|
||||
viewmats: torch.Tensor | None = None,
|
||||
Ks: torch.Tensor | None = None,
|
||||
is_cache: bool = False,
|
||||
attention_mask: torch.Tensor | None = None):
|
||||
"""
|
||||
Forward pass with camera PRoPE attention combining standard RoPE and projective positional encoding.
|
||||
|
||||
Args:
|
||||
q, k, v: Query, key, value tensors [B, L, num_heads, head_dim]
|
||||
freqs_cis: RoPE frequency cos/sin tensors
|
||||
kv_cache: KV cache dict (may have None values for training)
|
||||
current_start: Current position for KV cache
|
||||
cache_start: Cache start position
|
||||
viewmats: Camera view matrices for PRoPE [B, cameras, 4, 4]
|
||||
Ks: Camera intrinsics for PRoPE [B, cameras, 3, 3]
|
||||
is_cache: Whether to store to KV cache (for inference)
|
||||
attention_mask: Attention mask [B, L] (1 = attend, 0 = mask)
|
||||
"""
|
||||
if cache_start is None:
|
||||
cache_start = current_start
|
||||
|
||||
# Apply RoPE manually
|
||||
cos, sin = freqs_cis
|
||||
query_rope = _apply_rotary_emb(q, cos, sin, is_neox_style=False).type_as(v)
|
||||
key_rope = _apply_rotary_emb(k, cos, sin, is_neox_style=False).type_as(v)
|
||||
value_rope = v
|
||||
|
||||
# Get PRoPE transformed q, k, v
|
||||
query_prope, key_prope, value_prope, apply_fn_o = prope_qkv(
|
||||
q.transpose(1, 2), # [B, num_heads, L, head_dim]
|
||||
k.transpose(1, 2),
|
||||
v.transpose(1, 2),
|
||||
viewmats=viewmats,
|
||||
Ks=Ks,
|
||||
patches_x=40, # hardcoded for now
|
||||
patches_y=22,
|
||||
)
|
||||
# PRoPE returns [B, num_heads, L, head_dim], convert to [B, L, num_heads, head_dim]
|
||||
query_prope = query_prope.transpose(1, 2)
|
||||
key_prope = key_prope.transpose(1, 2)
|
||||
value_prope = value_prope.transpose(1, 2)
|
||||
|
||||
# KV cache handling
|
||||
if kv_cache is not None:
|
||||
cache_key = kv_cache.get("k", None)
|
||||
cache_value = kv_cache.get("v", None)
|
||||
|
||||
if cache_value is not None and not is_cache:
|
||||
cache_key_rope, cache_key_prope = cache_key.chunk(2, dim=-1)
|
||||
cache_value_rope, cache_value_prope = cache_value.chunk(2, dim=-1)
|
||||
|
||||
key_rope = torch.cat([cache_key_rope, key_rope], dim=1)
|
||||
value_rope = torch.cat([cache_value_rope, value_rope], dim=1)
|
||||
key_prope = torch.cat([cache_key_prope, key_prope], dim=1)
|
||||
value_prope = torch.cat([cache_value_prope, value_prope], dim=1)
|
||||
|
||||
if is_cache:
|
||||
# Store to cache (update input dict directly)
|
||||
kv_cache["k"] = torch.cat([key_rope, key_prope], dim=-1)
|
||||
kv_cache["v"] = torch.cat([value_rope, value_prope], dim=-1)
|
||||
|
||||
# Concatenate rope and prope paths (matching original)
|
||||
query_all = torch.cat([query_rope, query_prope], dim=0)
|
||||
key_all = torch.cat([key_rope, key_prope], dim=0)
|
||||
value_all = torch.cat([value_rope, value_prope], dim=0)
|
||||
|
||||
# Check if Q and KV have different sequence lengths (KV cache mode)
|
||||
# In this case, use LocalAttention (supports different Q/KV lengths)
|
||||
if query_all.shape[1] != key_all.shape[1]:
|
||||
raise ValueError("Q and KV have different sequence lengths")
|
||||
# KV cache mode: Q has new tokens only, KV has cached + new tokens
|
||||
# Use LocalAttention which supports different Q/KV lengths
|
||||
# LocalAttention will use the appropriate backend (SageAttn, FlashAttn, etc.)
|
||||
if not hasattr(self, '_kv_cache_attn'):
|
||||
from fastvideo.attention import LocalAttention
|
||||
self._kv_cache_attn = LocalAttention(
|
||||
num_heads=self.num_heads,
|
||||
head_size=self.head_dim,
|
||||
causal=False,
|
||||
supported_attention_backends=(AttentionBackendEnum.SAGE_ATTN,
|
||||
AttentionBackendEnum.FLASH_ATTN,
|
||||
AttentionBackendEnum.TORCH_SDPA)
|
||||
)
|
||||
hidden_states_all = self._kv_cache_attn(query_all, key_all, value_all)
|
||||
else:
|
||||
# Same sequence length: use DistributedAttention (supports SP)
|
||||
# Create default attention mask if not provided
|
||||
if attention_mask is None:
|
||||
batch_size, seq_len = q.shape[0], q.shape[1]
|
||||
attention_mask = torch.ones(batch_size, seq_len, device=q.device, dtype=q.dtype)
|
||||
|
||||
if q.dtype == torch.float32:
|
||||
from fastvideo.attention.backends.sdpa import SDPAMetadataBuilder
|
||||
attn_metadata_builder = SDPAMetadataBuilder
|
||||
else:
|
||||
from fastvideo.attention.backends.flash_attn import FlashAttnMetadataBuilder
|
||||
attn_metadata_builder = FlashAttnMetadataBuilder
|
||||
attn_metadata = attn_metadata_builder().build(
|
||||
current_timestep=0,
|
||||
attn_mask=attention_mask,
|
||||
)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=attn_metadata):
|
||||
hidden_states_all, _ = self.attn(query_all, key_all, value_all, attention_mask=attention_mask)
|
||||
|
||||
hidden_states_rope, hidden_states_prope = hidden_states_all.chunk(2, dim=0)
|
||||
hidden_states_prope = apply_fn_o(hidden_states_prope.transpose(1, 2)).transpose(1, 2)
|
||||
|
||||
return hidden_states_rope, hidden_states_prope
|
||||
@@ -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
|
||||
@@ -76,7 +76,7 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
dim,
|
||||
dim,
|
||||
bias=True,
|
||||
act_type="gelu_pytorch_tanh")
|
||||
act_type="gelu_pytorch_tanh") if text_embed_dim > 0 else None
|
||||
|
||||
self.image_embedder = None
|
||||
if image_embed_dim is not None:
|
||||
@@ -92,7 +92,10 @@ class WanTimeTextImageEmbedding(nn.Module):
|
||||
temb = self.time_embedder(timestep, timestep_seq_len)
|
||||
timestep_proj = self.time_modulation(temb)
|
||||
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
if self.text_embedder is not None:
|
||||
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
|
||||
else:
|
||||
encoder_hidden_states = torch.zeros((timestep.shape[0], 0, temb.shape[-1]), device=temb.device, dtype=temb.dtype)
|
||||
if encoder_hidden_states_image is not None:
|
||||
assert self.image_embedder is not None
|
||||
encoder_hidden_states_image = self.image_embedder(
|
||||
@@ -179,7 +182,10 @@ class WanT2VCrossAttention(WanSelfAttention):
|
||||
v = self.to_v(context)[0].view(b, -1, n, d)
|
||||
|
||||
# compute attention
|
||||
x = self.attn(q, k, v)
|
||||
if k.size(1) > 0:
|
||||
x = self.attn(q, k, v)
|
||||
else:
|
||||
x = torch.zeros_like(q)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
@@ -227,7 +233,10 @@ class WanI2VCrossAttention(WanSelfAttention):
|
||||
v_img = self.add_v_proj(context_img)[0].view(b, -1, n, d)
|
||||
img_x = self.attn(q, k_img, v_img)
|
||||
# compute attention
|
||||
x = self.attn(q, k, v)
|
||||
if k.size(1) > 0:
|
||||
x = self.attn(q, k, v)
|
||||
else:
|
||||
x = torch.zeros_like(q)
|
||||
|
||||
# output
|
||||
x = x.flatten(2)
|
||||
@@ -633,7 +642,7 @@ class WanTransformer3DModel(CachableDiT):
|
||||
enable_teacache = forward_batch is not None and forward_batch.enable_teacache
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
if not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
if encoder_hidden_states is not None and not isinstance(encoder_hidden_states, torch.Tensor):
|
||||
encoder_hidden_states = encoder_hidden_states[0]
|
||||
if isinstance(encoder_hidden_states_image,
|
||||
list) and len(encoder_hidden_states_image) > 0:
|
||||
@@ -705,8 +714,11 @@ class WanTransformer3DModel(CachableDiT):
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
encoder_hidden_states = torch.concat(
|
||||
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
if encoder_hidden_states is not None:
|
||||
encoder_hidden_states = torch.concat(
|
||||
[encoder_hidden_states_image, encoder_hidden_states], dim=1)
|
||||
else:
|
||||
encoder_hidden_states = encoder_hidden_states_image
|
||||
|
||||
if current_platform.is_mps() or current_platform.is_npu():
|
||||
encoder_hidden_states = encoder_hidden_states.to(orig_dtype)
|
||||
|
||||
@@ -806,6 +806,10 @@ class TransformerLoader(ComponentLoader):
|
||||
cls_name.startswith("Cosmos25")
|
||||
or cls_name == "Cosmos25Transformer3DModel"
|
||||
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "Cosmos25"
|
||||
) and not (
|
||||
cls_name.startswith("WanGame")
|
||||
or cls_name == "WanGameActionTransformer3DModel"
|
||||
or getattr(fastvideo_args.pipeline_config, "prefix", "") == "WanGame"
|
||||
)
|
||||
model = maybe_load_fsdp_model(
|
||||
model_cls=model_cls,
|
||||
|
||||
@@ -138,7 +138,7 @@ def maybe_load_fsdp_model(
|
||||
|
||||
weight_iterator = safetensors_weights_iterator(weight_dir_list)
|
||||
param_names_mapping_fn = get_param_names_mapping(model.param_names_mapping)
|
||||
load_model_from_full_model_state_dict(
|
||||
incompatible_keys, unexpected_keys = load_model_from_full_model_state_dict(
|
||||
model,
|
||||
weight_iterator,
|
||||
device,
|
||||
@@ -147,6 +147,9 @@ def maybe_load_fsdp_model(
|
||||
cpu_offload=cpu_offload,
|
||||
param_names_mapping=param_names_mapping_fn,
|
||||
)
|
||||
if incompatible_keys or unexpected_keys:
|
||||
logger.warning("Incompatible keys: %s", incompatible_keys)
|
||||
logger.warning("Unexpected keys: %s", unexpected_keys)
|
||||
for n, p in chain(model.named_parameters(), model.named_buffers()):
|
||||
if p.is_meta:
|
||||
raise RuntimeError(
|
||||
@@ -340,7 +343,7 @@ def load_model_from_full_model_state_dict(
|
||||
unused_keys)
|
||||
|
||||
# List of allowed parameter name patterns
|
||||
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l"] # Can be extended as needed
|
||||
ALLOWED_NEW_PARAM_PATTERNS = ["gate_compress", "proj_l", "to_out_prope", "action_embedder"] # Can be extended as needed
|
||||
for new_param_name in unused_keys:
|
||||
if not any(pattern in new_param_name
|
||||
for pattern in ALLOWED_NEW_PARAM_PATTERNS):
|
||||
|
||||
@@ -42,8 +42,9 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
# "HunyuanVideoTransformer3DModel": ("dits", "hunyuanvideo", "HunyuanVideoDiT"),
|
||||
"WanTransformer3DModel": ("dits", "wanvideo", "WanTransformer3DModel"),
|
||||
"CausalWanTransformer3DModel": ("dits", "causal_wanvideo", "CausalWanTransformer3DModel"),
|
||||
"MatrixGameWanModel": ("dits", "matrix_game", "MatrixGameWanModel"),
|
||||
"CausalMatrixGameWanModel": ("dits", "matrix_game", "CausalMatrixGameWanModel"),
|
||||
"WanGameActionTransformer3DModel": ("dits", "wangame", "WanGameActionTransformer3DModel"),
|
||||
"MatrixGameWanModel": ("dits", "matrixgame", "MatrixGameWanModel"),
|
||||
"CausalMatrixGameWanModel": ("dits", "matrixgame", "CausalMatrixGameWanModel"),
|
||||
}
|
||||
|
||||
_TEXT_ENCODER_MODELS = {
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Wan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.lora_pipeline import LoRAPipeline
|
||||
|
||||
# isort: off
|
||||
from fastvideo.pipelines.stages import (
|
||||
ImageEncodingStage, ConditioningStage, DecodingStage, DenoisingStage,
|
||||
ImageVAEEncodingStage, InputValidationStage, LatentPreparationStage,
|
||||
TimestepPreparationStage)
|
||||
# isort: on
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanGameActionImageToVideoPipeline(LoRAPipeline, ComposedPipelineBase):
|
||||
|
||||
_required_config_modules = [
|
||||
"vae", "transformer", "scheduler", \
|
||||
"image_encoder", "image_processor"
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
"""Set up pipeline stages with proper dependency injection."""
|
||||
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="image_encoding_stage",
|
||||
stage=ImageEncodingStage(
|
||||
image_encoder=self.get_module("image_encoder"),
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
self.add_stage(stage_name="conditioning_stage",
|
||||
stage=ConditioningStage())
|
||||
|
||||
self.add_stage(stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="latent_preparation_stage",
|
||||
stage=LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"),
|
||||
transformer=self.get_module("transformer")))
|
||||
|
||||
self.add_stage(stage_name="image_latent_preparation_stage",
|
||||
stage=ImageVAEEncodingStage(vae=self.get_module("vae")))
|
||||
|
||||
self.add_stage(stage_name="denoising_stage",
|
||||
stage=DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler")))
|
||||
|
||||
self.add_stage(stage_name="decoding_stage",
|
||||
stage=DecodingStage(vae=self.get_module("vae")))
|
||||
|
||||
|
||||
EntryClass = WanGameActionImageToVideoPipeline
|
||||
@@ -21,6 +21,7 @@ _PIPELINE_NAME_TO_ARCHITECTURE_NAME: dict[str, str] = {
|
||||
"WanPipeline": "wan",
|
||||
"WanDMDPipeline": "wan",
|
||||
"WanImageToVideoPipeline": "wan",
|
||||
"WanGameActionImageToVideoPipeline": "wan",
|
||||
"WanVideoToVideoPipeline": "wan",
|
||||
"WanCausalDMDPipeline": "wan",
|
||||
"TurboDiffusionPipeline": "turbodiffusion",
|
||||
|
||||
@@ -0,0 +1,303 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_matrixgame
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import ImageEncodingStage
|
||||
|
||||
|
||||
class PreprocessPipeline_MatrixGame(BasePreprocessPipeline):
|
||||
"""I2V preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = ["vae", "image_encoder", "image_processor"]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=ImageEncodingStage(
|
||||
image_encoder=self.get_module("image_encoder"),
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for I2V pipeline."""
|
||||
return pyarrow_schema_matrixgame
|
||||
|
||||
def get_extra_features(self, valid_data: dict[str, Any],
|
||||
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
|
||||
|
||||
# TODO(will): move these to cpu at some point
|
||||
self.get_module("image_encoder").to(get_local_torch_device())
|
||||
self.get_module("vae").to(get_local_torch_device())
|
||||
|
||||
features = {}
|
||||
"""Get CLIP features from the first frame of each video."""
|
||||
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
|
||||
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
|
||||
_, _, num_frames, height, width = valid_data["pixel_values"].shape
|
||||
# latent_height = height // self.get_module(
|
||||
# "vae").spatial_compression_ratio
|
||||
# latent_width = width // self.get_module("vae").spatial_compression_ratio
|
||||
|
||||
processed_images = []
|
||||
# Frame has values between -1 and 1
|
||||
for frame in first_frame:
|
||||
frame = (frame + 1) * 127.5
|
||||
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
|
||||
processed_img = self.get_module("image_processor")(
|
||||
images=frame_pil, return_tensors="pt")
|
||||
processed_images.append(processed_img)
|
||||
|
||||
# Get CLIP features
|
||||
pixel_values = torch.cat(
|
||||
[img['pixel_values'] for img in processed_images],
|
||||
dim=0).to(get_local_torch_device())
|
||||
with torch.no_grad():
|
||||
image_inputs = {'pixel_values': pixel_values}
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
clip_features = self.get_module("image_encoder")(**image_inputs)
|
||||
clip_features = clip_features.last_hidden_state
|
||||
|
||||
features["clip_feature"] = clip_features
|
||||
"""Get VAE features from the first frame of each video"""
|
||||
video_conditions = []
|
||||
for frame in first_frame:
|
||||
processed_img = frame.to(device="cpu", dtype=torch.float32)
|
||||
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
|
||||
2).unsqueeze(2)
|
||||
# (B, H, W, C) -> (B, C, 1, H, W)
|
||||
video_condition = torch.cat([
|
||||
processed_img,
|
||||
processed_img.new_zeros(processed_img.shape[0],
|
||||
processed_img.shape[1], num_frames - 1,
|
||||
height, width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(
|
||||
device=get_local_torch_device(), dtype=torch.float32)
|
||||
video_conditions.append(video_condition)
|
||||
|
||||
video_conditions = torch.cat(video_conditions, dim=0)
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=torch.float32,
|
||||
enabled=True):
|
||||
encoder_outputs = self.get_module("vae").encode(video_conditions)
|
||||
|
||||
latent_condition = encoder_outputs.mean
|
||||
if (hasattr(self.get_module("vae"), "shift_factor")
|
||||
and self.get_module("vae").shift_factor is not None):
|
||||
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
|
||||
latent_condition -= self.get_module("vae").shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.get_module("vae").shift_factor
|
||||
|
||||
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.get_module(
|
||||
"vae").scaling_factor.to(latent_condition.device,
|
||||
latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.get_module(
|
||||
"vae").scaling_factor
|
||||
|
||||
# mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
|
||||
# latent_width)
|
||||
# mask_lat_size[:, :, list(range(1, num_frames))] = 0
|
||||
# first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
# first_frame_mask = torch.repeat_interleave(
|
||||
# first_frame_mask,
|
||||
# dim=2,
|
||||
# repeats=self.get_module("vae").temporal_compression_ratio)
|
||||
# mask_lat_size = torch.concat(
|
||||
# [first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
|
||||
# mask_lat_size = mask_lat_size.view(
|
||||
# batch_size, -1,
|
||||
# self.get_module("vae").temporal_compression_ratio, latent_height,
|
||||
# latent_width)
|
||||
# mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
# mask_lat_size = mask_lat_size.to(latent_condition.device)
|
||||
|
||||
# image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
|
||||
|
||||
features["first_frame_latent"] = latent_condition
|
||||
|
||||
if "action_path" in valid_data and valid_data["action_path"]:
|
||||
keyboard_cond_list = []
|
||||
mouse_cond_list = []
|
||||
num_bits = 6
|
||||
for action_path in valid_data["action_path"]:
|
||||
if action_path:
|
||||
action_data = np.load(action_path, allow_pickle=True)
|
||||
if isinstance(
|
||||
action_data,
|
||||
np.ndarray) and action_data.dtype == np.dtype('O'):
|
||||
action_dict = action_data.item()
|
||||
if "keyboard" in action_dict:
|
||||
keyboard_raw = action_dict["keyboard"]
|
||||
# Convert 1D bit-flag values to 2D multi-hot encoding
|
||||
if isinstance(keyboard_raw, np.ndarray):
|
||||
if keyboard_raw.ndim == 1:
|
||||
# [T] -> [T, num_bits]
|
||||
T = len(keyboard_raw)
|
||||
multi_hot = np.zeros((T, num_bits),
|
||||
dtype=np.float32)
|
||||
action_values = keyboard_raw.astype(int)
|
||||
for bit_idx in range(num_bits):
|
||||
target_idx = (
|
||||
2 -
|
||||
(bit_idx % 3)) + 3 * (bit_idx // 3)
|
||||
if target_idx < num_bits:
|
||||
multi_hot[:, target_idx] = (
|
||||
(action_values >> bit_idx)
|
||||
& 1).astype(np.float32)
|
||||
keyboard_cond_list.append(multi_hot)
|
||||
else:
|
||||
# If already 2D, pad to num_bits if necessary
|
||||
k_data = keyboard_raw.astype(np.float32)
|
||||
if k_data.ndim == 2 and k_data.shape[
|
||||
-1] < num_bits:
|
||||
padding = np.zeros(
|
||||
(k_data.shape[0],
|
||||
num_bits - k_data.shape[-1]),
|
||||
dtype=np.float32)
|
||||
k_data = np.concatenate(
|
||||
[k_data, padding], axis=-1)
|
||||
keyboard_cond_list.append(k_data)
|
||||
else:
|
||||
keyboard_cond_list.append(keyboard_raw)
|
||||
if "mouse" in action_dict:
|
||||
mouse_cond_list.append(action_dict["mouse"])
|
||||
else:
|
||||
if isinstance(action_data,
|
||||
np.ndarray) and action_data.ndim == 1:
|
||||
T = len(action_data)
|
||||
multi_hot = np.zeros((T, num_bits),
|
||||
dtype=np.float32)
|
||||
action_values = action_data.astype(int)
|
||||
for bit_idx in range(num_bits):
|
||||
target_idx = (
|
||||
2 - (bit_idx % 3)) + 3 * (bit_idx // 3)
|
||||
if target_idx < num_bits:
|
||||
multi_hot[:, target_idx] = (
|
||||
(action_values >> bit_idx) & 1).astype(
|
||||
np.float32)
|
||||
keyboard_cond_list.append(multi_hot)
|
||||
else:
|
||||
# If already 2D, pad to num_bits if necessary
|
||||
k_data = action_data.astype(np.float32)
|
||||
if k_data.ndim == 2 and k_data.shape[-1] < num_bits:
|
||||
padding = np.zeros(
|
||||
(k_data.shape[0],
|
||||
num_bits - k_data.shape[-1]),
|
||||
dtype=np.float32)
|
||||
k_data = np.concatenate([k_data, padding],
|
||||
axis=-1)
|
||||
keyboard_cond_list.append(k_data)
|
||||
if keyboard_cond_list:
|
||||
features["keyboard_cond"] = keyboard_cond_list
|
||||
if mouse_cond_list:
|
||||
features["mouse_cond"] = mouse_cond_list
|
||||
|
||||
return features
|
||||
|
||||
def create_record(
|
||||
self,
|
||||
video_name: str,
|
||||
vae_latent: np.ndarray,
|
||||
text_embedding: np.ndarray,
|
||||
valid_data: dict[str, Any],
|
||||
idx: int,
|
||||
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""Create a record for the Parquet dataset with CLIP features."""
|
||||
record = super().create_record(video_name=video_name,
|
||||
vae_latent=vae_latent,
|
||||
text_embedding=text_embedding,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=extra_features)
|
||||
|
||||
if extra_features and "clip_feature" in extra_features:
|
||||
clip_feature = extra_features["clip_feature"]
|
||||
record.update({
|
||||
"clip_feature_bytes": clip_feature.tobytes(),
|
||||
"clip_feature_shape": list(clip_feature.shape),
|
||||
"clip_feature_dtype": str(clip_feature.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"clip_feature_bytes": b"",
|
||||
"clip_feature_shape": [],
|
||||
"clip_feature_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "first_frame_latent" in extra_features:
|
||||
first_frame_latent = extra_features["first_frame_latent"]
|
||||
record.update({
|
||||
"first_frame_latent_bytes":
|
||||
first_frame_latent.tobytes(),
|
||||
"first_frame_latent_shape":
|
||||
list(first_frame_latent.shape),
|
||||
"first_frame_latent_dtype":
|
||||
str(first_frame_latent.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"first_frame_latent_bytes": b"",
|
||||
"first_frame_latent_shape": [],
|
||||
"first_frame_latent_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "pil_image" in extra_features:
|
||||
pil_image = extra_features["pil_image"]
|
||||
record.update({
|
||||
"pil_image_bytes": pil_image.tobytes(),
|
||||
"pil_image_shape": list(pil_image.shape),
|
||||
"pil_image_dtype": str(pil_image.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"pil_image_bytes": b"",
|
||||
"pil_image_shape": [],
|
||||
"pil_image_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "keyboard_cond" in extra_features:
|
||||
keyboard_cond = extra_features["keyboard_cond"]
|
||||
record.update({
|
||||
"keyboard_cond_bytes": keyboard_cond.tobytes(),
|
||||
"keyboard_cond_shape": list(keyboard_cond.shape),
|
||||
"keyboard_cond_dtype": str(keyboard_cond.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"keyboard_cond_bytes": b"",
|
||||
"keyboard_cond_shape": [],
|
||||
"keyboard_cond_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "mouse_cond" in extra_features:
|
||||
mouse_cond = extra_features["mouse_cond"]
|
||||
record.update({
|
||||
"mouse_cond_bytes": mouse_cond.tobytes(),
|
||||
"mouse_cond_shape": list(mouse_cond.shape),
|
||||
"mouse_cond_dtype": str(mouse_cond.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"mouse_cond_bytes": b"",
|
||||
"mouse_cond_shape": [],
|
||||
"mouse_cond_dtype": "",
|
||||
})
|
||||
|
||||
return record
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_MatrixGame
|
||||
@@ -320,12 +320,18 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
"pixel_values":
|
||||
torch.stack(
|
||||
[data["pixel_values"][i] for i in valid_indices]),
|
||||
"text": [data["text"][i] for i in valid_indices],
|
||||
"text": [data["text"][i] for i in valid_indices]
|
||||
if "text" in data else ["" for _ in valid_indices],
|
||||
"path": [data["path"][i] for i in valid_indices],
|
||||
"fps": [data["fps"][i] for i in valid_indices],
|
||||
"duration": [data["duration"][i] for i in valid_indices],
|
||||
}
|
||||
|
||||
if "action_path" in data:
|
||||
valid_data["action_path"] = [
|
||||
data["action_path"][i] for i in valid_indices
|
||||
]
|
||||
|
||||
# VAE
|
||||
with torch.autocast("cuda", dtype=torch.float32):
|
||||
latents = self.get_module("vae").encode(
|
||||
@@ -343,28 +349,35 @@ class BasePreprocessPipeline(ComposedPipelineBase):
|
||||
prompt_embeds=[],
|
||||
prompt_attention_mask=[],
|
||||
)
|
||||
assert hasattr(self, "prompt_encoding_stage")
|
||||
result_batch = self.prompt_encoding_stage(batch, fastvideo_args)
|
||||
prompt_embeds, prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
assert prompt_embeds.shape[0] == prompt_attention_mask.shape[0]
|
||||
|
||||
# Get sequence lengths from attention masks (number of 1s)
|
||||
seq_lens = prompt_attention_mask.sum(dim=1)
|
||||
if hasattr(self, "prompt_encoding_stage"):
|
||||
result_batch = self.prompt_encoding_stage(
|
||||
batch, fastvideo_args)
|
||||
prompt_embeds, prompt_attention_mask = result_batch.prompt_embeds[
|
||||
0], result_batch.prompt_attention_mask[0]
|
||||
assert prompt_embeds.shape[
|
||||
0] == prompt_attention_mask.shape[0]
|
||||
|
||||
non_padded_embeds = []
|
||||
non_padded_masks = []
|
||||
# Get sequence lengths from attention masks (number of 1s)
|
||||
seq_lens = prompt_attention_mask.sum(dim=1)
|
||||
|
||||
# Process each item in the batch
|
||||
for i in range(prompt_embeds.size(0)):
|
||||
seq_len = seq_lens[i].item()
|
||||
# Slice the embeddings and masks to keep only non-padding parts
|
||||
non_padded_embeds.append(prompt_embeds[i, :seq_len])
|
||||
non_padded_masks.append(prompt_attention_mask[i, :seq_len])
|
||||
non_padded_embeds = []
|
||||
non_padded_masks = []
|
||||
|
||||
# Update the tensors with non-padded versions
|
||||
prompt_embeds = non_padded_embeds
|
||||
prompt_attention_mask = non_padded_masks
|
||||
# Process each item in the batch
|
||||
for i in range(prompt_embeds.size(0)):
|
||||
seq_len = seq_lens[i].item()
|
||||
# Slice the embeddings and masks to keep only non-padding parts
|
||||
non_padded_embeds.append(prompt_embeds[i, :seq_len])
|
||||
non_padded_masks.append(
|
||||
prompt_attention_mask[i, :seq_len])
|
||||
|
||||
# Update the tensors with non-padded versions
|
||||
prompt_embeds = non_padded_embeds
|
||||
prompt_attention_mask = non_padded_masks
|
||||
else:
|
||||
bs = len(valid_indices)
|
||||
prompt_embeds = [torch.zeros(0) for _ in range(bs)]
|
||||
|
||||
# Prepare batch data for Parquet dataset
|
||||
batch_data = []
|
||||
|
||||
@@ -16,6 +16,10 @@ from fastvideo.pipelines.preprocess.preprocess_pipeline_t2v import (
|
||||
PreprocessPipeline_T2V)
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_text import (
|
||||
PreprocessPipeline_Text)
|
||||
from fastvideo.pipelines.preprocess.matrixgame.matrixgame_preprocess_pipeline import (
|
||||
PreprocessPipeline_MatrixGame)
|
||||
from fastvideo.pipelines.preprocess.wangame.wangame_preprocess_pipeline import (
|
||||
PreprocessPipeline_WanGame)
|
||||
from fastvideo.utils import maybe_download_model
|
||||
|
||||
logger = init_logger(__name__)
|
||||
@@ -60,9 +64,14 @@ def main(args) -> None:
|
||||
assert args.flow_shift is not None, "flow_shift is required for ode_trajectory"
|
||||
fastvideo_args.pipeline_config.flow_shift = args.flow_shift
|
||||
PreprocessPipeline = PreprocessPipeline_ODE_Trajectory
|
||||
elif args.preprocess_task == "matrixgame":
|
||||
PreprocessPipeline = PreprocessPipeline_MatrixGame
|
||||
elif args.preprocess_task == "wangame":
|
||||
PreprocessPipeline = PreprocessPipeline_WanGame
|
||||
else:
|
||||
raise ValueError(f"Invalid preprocess task: {args.preprocess_task}. "
|
||||
f"Valid options: t2v, i2v, ode_trajectory, text_only")
|
||||
raise ValueError(
|
||||
f"Invalid preprocess task: {args.preprocess_task}. "
|
||||
f"Valid options: t2v, i2v, ode_trajectory, text_only, matrixgame, wangame")
|
||||
|
||||
logger.info("Preprocess task: %s using %s", args.preprocess_task,
|
||||
PreprocessPipeline.__name__)
|
||||
@@ -106,11 +115,12 @@ if __name__ == "__main__":
|
||||
parser.add_argument("--group_frame", action="store_true") # TODO
|
||||
parser.add_argument("--group_resolution", action="store_true") # TODO
|
||||
parser.add_argument("--flow_shift", type=float, default=None)
|
||||
parser.add_argument("--preprocess_task",
|
||||
type=str,
|
||||
default="t2v",
|
||||
choices=["t2v", "i2v", "text_only", "ode_trajectory"],
|
||||
help="Type of preprocessing task to run")
|
||||
parser.add_argument(
|
||||
"--preprocess_task",
|
||||
type=str,
|
||||
default="t2v",
|
||||
choices=["t2v", "i2v", "text_only", "ode_trajectory", "matrixgame", "wangame"],
|
||||
help="Type of preprocessing task to run")
|
||||
parser.add_argument("--train_fps", type=int, default=30)
|
||||
parser.add_argument("--use_image_num", type=int, default=0)
|
||||
parser.add_argument("--text_max_length", type=int, default=256)
|
||||
|
||||
@@ -0,0 +1,303 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
from PIL import Image
|
||||
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_wangame
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.pipelines.preprocess.preprocess_pipeline_base import (
|
||||
BasePreprocessPipeline)
|
||||
from fastvideo.pipelines.stages import ImageEncodingStage
|
||||
|
||||
|
||||
class PreprocessPipeline_WanGame(BasePreprocessPipeline):
|
||||
"""I2V preprocessing pipeline implementation."""
|
||||
|
||||
_required_config_modules = ["vae", "image_encoder", "image_processor"]
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
self.add_stage(stage_name="image_encoding_stage",
|
||||
stage=ImageEncodingStage(
|
||||
image_encoder=self.get_module("image_encoder"),
|
||||
image_processor=self.get_module("image_processor"),
|
||||
))
|
||||
|
||||
def get_pyarrow_schema(self):
|
||||
"""Return the PyArrow schema for I2V pipeline."""
|
||||
return pyarrow_schema_wangame
|
||||
|
||||
def get_extra_features(self, valid_data: dict[str, Any],
|
||||
fastvideo_args: FastVideoArgs) -> dict[str, Any]:
|
||||
|
||||
# TODO(will): move these to cpu at some point
|
||||
self.get_module("image_encoder").to(get_local_torch_device())
|
||||
self.get_module("vae").to(get_local_torch_device())
|
||||
|
||||
features = {}
|
||||
"""Get CLIP features from the first frame of each video."""
|
||||
first_frame = valid_data["pixel_values"][:, :, 0, :, :].permute(
|
||||
0, 2, 3, 1) # (B, C, T, H, W) -> (B, H, W, C)
|
||||
_, _, num_frames, height, width = valid_data["pixel_values"].shape
|
||||
# latent_height = height // self.get_module(
|
||||
# "vae").spatial_compression_ratio
|
||||
# latent_width = width // self.get_module("vae").spatial_compression_ratio
|
||||
|
||||
processed_images = []
|
||||
# Frame has values between -1 and 1
|
||||
for frame in first_frame:
|
||||
frame = (frame + 1) * 127.5
|
||||
frame_pil = Image.fromarray(frame.cpu().numpy().astype(np.uint8))
|
||||
processed_img = self.get_module("image_processor")(
|
||||
images=frame_pil, return_tensors="pt")
|
||||
processed_images.append(processed_img)
|
||||
|
||||
# Get CLIP features
|
||||
pixel_values = torch.cat(
|
||||
[img['pixel_values'] for img in processed_images],
|
||||
dim=0).to(get_local_torch_device())
|
||||
with torch.no_grad():
|
||||
image_inputs = {'pixel_values': pixel_values}
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
clip_features = self.get_module("image_encoder")(**image_inputs)
|
||||
clip_features = clip_features.last_hidden_state
|
||||
|
||||
features["clip_feature"] = clip_features
|
||||
"""Get VAE features from the first frame of each video"""
|
||||
video_conditions = []
|
||||
for frame in first_frame:
|
||||
processed_img = frame.to(device="cpu", dtype=torch.float32)
|
||||
processed_img = processed_img.unsqueeze(0).permute(0, 3, 1,
|
||||
2).unsqueeze(2)
|
||||
# (B, H, W, C) -> (B, C, 1, H, W)
|
||||
video_condition = torch.cat([
|
||||
processed_img,
|
||||
processed_img.new_zeros(processed_img.shape[0],
|
||||
processed_img.shape[1], num_frames - 1,
|
||||
height, width)
|
||||
],
|
||||
dim=2)
|
||||
video_condition = video_condition.to(
|
||||
device=get_local_torch_device(), dtype=torch.float32)
|
||||
video_conditions.append(video_condition)
|
||||
|
||||
video_conditions = torch.cat(video_conditions, dim=0)
|
||||
|
||||
with torch.autocast(device_type="cuda",
|
||||
dtype=torch.float32,
|
||||
enabled=True):
|
||||
encoder_outputs = self.get_module("vae").encode(video_conditions)
|
||||
|
||||
latent_condition = encoder_outputs.mean
|
||||
if (hasattr(self.get_module("vae"), "shift_factor")
|
||||
and self.get_module("vae").shift_factor is not None):
|
||||
if isinstance(self.get_module("vae").shift_factor, torch.Tensor):
|
||||
latent_condition -= self.get_module("vae").shift_factor.to(
|
||||
latent_condition.device, latent_condition.dtype)
|
||||
else:
|
||||
latent_condition -= self.get_module("vae").shift_factor
|
||||
|
||||
if isinstance(self.get_module("vae").scaling_factor, torch.Tensor):
|
||||
latent_condition = latent_condition * self.get_module(
|
||||
"vae").scaling_factor.to(latent_condition.device,
|
||||
latent_condition.dtype)
|
||||
else:
|
||||
latent_condition = latent_condition * self.get_module(
|
||||
"vae").scaling_factor
|
||||
|
||||
# mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
|
||||
# latent_width)
|
||||
# mask_lat_size[:, :, list(range(1, num_frames))] = 0
|
||||
# first_frame_mask = mask_lat_size[:, :, 0:1]
|
||||
# first_frame_mask = torch.repeat_interleave(
|
||||
# first_frame_mask,
|
||||
# dim=2,
|
||||
# repeats=self.get_module("vae").temporal_compression_ratio)
|
||||
# mask_lat_size = torch.concat(
|
||||
# [first_frame_mask, mask_lat_size[:, :, 1:, :]], dim=2)
|
||||
# mask_lat_size = mask_lat_size.view(
|
||||
# batch_size, -1,
|
||||
# self.get_module("vae").temporal_compression_ratio, latent_height,
|
||||
# latent_width)
|
||||
# mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
# mask_lat_size = mask_lat_size.to(latent_condition.device)
|
||||
|
||||
# image_latent = torch.concat([mask_lat_size, latent_condition], dim=1)
|
||||
|
||||
features["first_frame_latent"] = latent_condition
|
||||
|
||||
if "action_path" in valid_data and valid_data["action_path"]:
|
||||
keyboard_cond_list = []
|
||||
mouse_cond_list = []
|
||||
num_bits = 6
|
||||
for action_path in valid_data["action_path"]:
|
||||
if action_path:
|
||||
action_data = np.load(action_path, allow_pickle=True)
|
||||
if isinstance(
|
||||
action_data,
|
||||
np.ndarray) and action_data.dtype == np.dtype('O'):
|
||||
action_dict = action_data.item()
|
||||
if "keyboard" in action_dict:
|
||||
keyboard_raw = action_dict["keyboard"]
|
||||
# Convert 1D bit-flag values to 2D multi-hot encoding
|
||||
if isinstance(keyboard_raw, np.ndarray):
|
||||
if keyboard_raw.ndim == 1:
|
||||
# [T] -> [T, num_bits]
|
||||
T = len(keyboard_raw)
|
||||
multi_hot = np.zeros((T, num_bits),
|
||||
dtype=np.float32)
|
||||
action_values = keyboard_raw.astype(int)
|
||||
for bit_idx in range(num_bits):
|
||||
target_idx = (
|
||||
2 -
|
||||
(bit_idx % 3)) + 3 * (bit_idx // 3)
|
||||
if target_idx < num_bits:
|
||||
multi_hot[:, target_idx] = (
|
||||
(action_values >> bit_idx)
|
||||
& 1).astype(np.float32)
|
||||
keyboard_cond_list.append(multi_hot)
|
||||
else:
|
||||
# If already 2D, pad to num_bits if necessary
|
||||
k_data = keyboard_raw.astype(np.float32)
|
||||
if k_data.ndim == 2 and k_data.shape[
|
||||
-1] < num_bits:
|
||||
padding = np.zeros(
|
||||
(k_data.shape[0],
|
||||
num_bits - k_data.shape[-1]),
|
||||
dtype=np.float32)
|
||||
k_data = np.concatenate(
|
||||
[k_data, padding], axis=-1)
|
||||
keyboard_cond_list.append(k_data)
|
||||
else:
|
||||
keyboard_cond_list.append(keyboard_raw)
|
||||
if "mouse" in action_dict:
|
||||
mouse_cond_list.append(action_dict["mouse"])
|
||||
else:
|
||||
if isinstance(action_data,
|
||||
np.ndarray) and action_data.ndim == 1:
|
||||
T = len(action_data)
|
||||
multi_hot = np.zeros((T, num_bits),
|
||||
dtype=np.float32)
|
||||
action_values = action_data.astype(int)
|
||||
for bit_idx in range(num_bits):
|
||||
target_idx = (
|
||||
2 - (bit_idx % 3)) + 3 * (bit_idx // 3)
|
||||
if target_idx < num_bits:
|
||||
multi_hot[:, target_idx] = (
|
||||
(action_values >> bit_idx) & 1).astype(
|
||||
np.float32)
|
||||
keyboard_cond_list.append(multi_hot)
|
||||
else:
|
||||
# If already 2D, pad to num_bits if necessary
|
||||
k_data = action_data.astype(np.float32)
|
||||
if k_data.ndim == 2 and k_data.shape[-1] < num_bits:
|
||||
padding = np.zeros(
|
||||
(k_data.shape[0],
|
||||
num_bits - k_data.shape[-1]),
|
||||
dtype=np.float32)
|
||||
k_data = np.concatenate([k_data, padding],
|
||||
axis=-1)
|
||||
keyboard_cond_list.append(k_data)
|
||||
if keyboard_cond_list:
|
||||
features["keyboard_cond"] = keyboard_cond_list
|
||||
if mouse_cond_list:
|
||||
features["mouse_cond"] = mouse_cond_list
|
||||
|
||||
return features
|
||||
|
||||
def create_record(
|
||||
self,
|
||||
video_name: str,
|
||||
vae_latent: np.ndarray,
|
||||
text_embedding: np.ndarray,
|
||||
valid_data: dict[str, Any],
|
||||
idx: int,
|
||||
extra_features: dict[str, Any] | None = None) -> dict[str, Any]:
|
||||
"""Create a record for the Parquet dataset with CLIP features."""
|
||||
record = super().create_record(video_name=video_name,
|
||||
vae_latent=vae_latent,
|
||||
text_embedding=text_embedding,
|
||||
valid_data=valid_data,
|
||||
idx=idx,
|
||||
extra_features=extra_features)
|
||||
|
||||
if extra_features and "clip_feature" in extra_features:
|
||||
clip_feature = extra_features["clip_feature"]
|
||||
record.update({
|
||||
"clip_feature_bytes": clip_feature.tobytes(),
|
||||
"clip_feature_shape": list(clip_feature.shape),
|
||||
"clip_feature_dtype": str(clip_feature.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"clip_feature_bytes": b"",
|
||||
"clip_feature_shape": [],
|
||||
"clip_feature_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "first_frame_latent" in extra_features:
|
||||
first_frame_latent = extra_features["first_frame_latent"]
|
||||
record.update({
|
||||
"first_frame_latent_bytes":
|
||||
first_frame_latent.tobytes(),
|
||||
"first_frame_latent_shape":
|
||||
list(first_frame_latent.shape),
|
||||
"first_frame_latent_dtype":
|
||||
str(first_frame_latent.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"first_frame_latent_bytes": b"",
|
||||
"first_frame_latent_shape": [],
|
||||
"first_frame_latent_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "pil_image" in extra_features:
|
||||
pil_image = extra_features["pil_image"]
|
||||
record.update({
|
||||
"pil_image_bytes": pil_image.tobytes(),
|
||||
"pil_image_shape": list(pil_image.shape),
|
||||
"pil_image_dtype": str(pil_image.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"pil_image_bytes": b"",
|
||||
"pil_image_shape": [],
|
||||
"pil_image_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "keyboard_cond" in extra_features:
|
||||
keyboard_cond = extra_features["keyboard_cond"]
|
||||
record.update({
|
||||
"keyboard_cond_bytes": keyboard_cond.tobytes(),
|
||||
"keyboard_cond_shape": list(keyboard_cond.shape),
|
||||
"keyboard_cond_dtype": str(keyboard_cond.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"keyboard_cond_bytes": b"",
|
||||
"keyboard_cond_shape": [],
|
||||
"keyboard_cond_dtype": "",
|
||||
})
|
||||
|
||||
if extra_features and "mouse_cond" in extra_features:
|
||||
mouse_cond = extra_features["mouse_cond"]
|
||||
record.update({
|
||||
"mouse_cond_bytes": mouse_cond.tobytes(),
|
||||
"mouse_cond_shape": list(mouse_cond.shape),
|
||||
"mouse_cond_dtype": str(mouse_cond.dtype),
|
||||
})
|
||||
else:
|
||||
record.update({
|
||||
"mouse_cond_bytes": b"",
|
||||
"mouse_cond_shape": [],
|
||||
"mouse_cond_dtype": "",
|
||||
})
|
||||
|
||||
return record
|
||||
|
||||
|
||||
EntryClass = PreprocessPipeline_WanGame
|
||||
@@ -168,6 +168,20 @@ class DenoisingStage(PipelineStage):
|
||||
},
|
||||
)
|
||||
|
||||
if batch.mouse_cond is not None and batch.keyboard_cond is not None:
|
||||
from fastvideo.models.dits.hyworld.pose import process_custom_actions
|
||||
viewmats, intrinsics, action_labels = process_custom_actions(batch.keyboard_cond, batch.mouse_cond)
|
||||
camera_action_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
"viewmats": viewmats.unsqueeze(0).to(get_local_torch_device(), dtype=target_dtype),
|
||||
"Ks": intrinsics.unsqueeze(0).to(get_local_torch_device(), dtype=target_dtype),
|
||||
"action": action_labels.unsqueeze(0).to(get_local_torch_device(), dtype=target_dtype),
|
||||
},
|
||||
)
|
||||
else:
|
||||
camera_action_kwargs = {}
|
||||
|
||||
action_kwargs = self.prepare_extra_func_kwargs(
|
||||
self.transformer.forward,
|
||||
{
|
||||
@@ -406,6 +420,7 @@ class DenoisingStage(PipelineStage):
|
||||
**image_kwargs,
|
||||
**pos_cond_kwargs,
|
||||
**action_kwargs,
|
||||
**camera_action_kwargs,
|
||||
)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
@@ -423,6 +438,7 @@ class DenoisingStage(PipelineStage):
|
||||
**image_kwargs,
|
||||
**neg_cond_kwargs,
|
||||
**action_kwargs,
|
||||
**camera_action_kwargs,
|
||||
)
|
||||
|
||||
noise_pred_text = noise_pred
|
||||
|
||||
@@ -210,9 +210,9 @@ class InputValidationStage(PipelineStage):
|
||||
f"keyboard_cond must have 3 dimensions (B, T, K), but got {batch.keyboard_cond.dim()}"
|
||||
)
|
||||
keyboard_dim = batch.keyboard_cond.shape[-1]
|
||||
if keyboard_dim not in {2, 4, 6, 7}:
|
||||
if keyboard_dim not in {2, 3, 4, 6, 7}:
|
||||
raise ValueError(
|
||||
f"keyboard_cond last dimension must be 2, 4, 6, or 7, but got {keyboard_dim}"
|
||||
f"keyboard_cond last dimension must be 2, 3, 4, 6, or 7, but got {keyboard_dim}"
|
||||
)
|
||||
logger.info(
|
||||
"Action control: keyboard_cond validated - shape %s (dim=%d)",
|
||||
|
||||
@@ -155,7 +155,7 @@ class CudaPlatformBase(Platform):
|
||||
)
|
||||
elif selected_backend == AttentionBackendEnum.SAGE_ATTN_THREE:
|
||||
try:
|
||||
from fastvideo.attention.backends.sageattn.api import sageattn_blackwell # noqa: F401
|
||||
from sageattn3 import sageattn3_blackwell # noqa: F401
|
||||
|
||||
from fastvideo.attention.backends.sage_attn3 import ( # noqa: F401
|
||||
SageAttention3Backend)
|
||||
|
||||
@@ -0,0 +1,138 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Mismatch compatibility test (direction B):
|
||||
- WORLD_SIZE=4
|
||||
- SP: sp_size=2
|
||||
- HSDP/FSDP2 mesh: replicate_dim=1, shard_dim=4 (mesh_shape=(1,4))
|
||||
|
||||
We verify gradients from:
|
||||
(A) full MSE: mean((pred-target)^2)
|
||||
(B) SP-sharded loss: sp_size * local_sse / total_numel
|
||||
match under this mismatch.
|
||||
|
||||
Run:
|
||||
source /home/hao_lab/miniconda3/etc/profile.d/conda.sh
|
||||
conda activate alexfv
|
||||
torchrun --standalone --nproc_per_node=4 -m pytest -q fastvideo/tests/distributed/test_hsdp_sp_mismatch_sp2_hsdp_shard4.py
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.distributed import cleanup_dist_env_and_memory
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
get_sp_group,
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
)
|
||||
from fastvideo.training.training_utils import shard_latents_across_sp
|
||||
|
||||
|
||||
def _world_size() -> int:
|
||||
return int(os.environ.get("WORLD_SIZE", "1"))
|
||||
|
||||
|
||||
def _rank() -> int:
|
||||
return int(os.environ.get("RANK", "0"))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def dist_setup():
|
||||
if _world_size() != 4:
|
||||
pytest.skip("Designed for torchrun WORLD_SIZE=4")
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("Requires CUDA")
|
||||
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
|
||||
torch.cuda.set_device(local_rank)
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=2)
|
||||
assert get_sp_group().world_size == 2
|
||||
yield
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
dist.barrier()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
class ScaleModule(torch.nn.Module):
|
||||
def __init__(self, c: int, init: torch.Tensor):
|
||||
super().__init__()
|
||||
assert init.shape == (c,)
|
||||
self.scale = torch.nn.Parameter(init.clone())
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return x * self.scale.view(1, -1, 1, 1, 1)
|
||||
|
||||
|
||||
def _broadcast_(t: torch.Tensor, src: int = 0) -> None:
|
||||
dist.broadcast(t, src=src)
|
||||
|
||||
|
||||
def _gather_full_vec_allranks(local_shard: torch.Tensor) -> torch.Tensor:
|
||||
gathered = [torch.empty_like(local_shard) for _ in range(dist.get_world_size())]
|
||||
dist.all_gather(gathered, local_shard)
|
||||
return torch.cat(gathered, dim=0)
|
||||
|
||||
|
||||
def test_sp_sharded_loss_matches_full_mse_under_mismatch(dist_setup):
|
||||
from torch.distributed.device_mesh import init_device_mesh
|
||||
from torch.distributed.fsdp import MixedPrecisionPolicy, fully_shard
|
||||
|
||||
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
|
||||
torch.manual_seed(2027)
|
||||
|
||||
b, c, t, h, w = 2, 8, 1, 1, 8 # thw=8 divisible by sp=2
|
||||
|
||||
x = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
|
||||
if _rank() == 0:
|
||||
x.normal_()
|
||||
_broadcast_(x)
|
||||
target = (0.25 * x).contiguous()
|
||||
|
||||
init = torch.empty((c,), device=device, dtype=torch.float32)
|
||||
if _rank() == 0:
|
||||
init.normal_()
|
||||
_broadcast_(init)
|
||||
|
||||
mesh = init_device_mesh(
|
||||
"cuda",
|
||||
mesh_shape=(1, 4),
|
||||
mesh_dim_names=("replicate", "shard"),
|
||||
)
|
||||
|
||||
mp = MixedPrecisionPolicy(
|
||||
param_dtype=torch.float32,
|
||||
reduce_dtype=torch.float32,
|
||||
output_dtype=torch.float32,
|
||||
cast_forward_inputs=False,
|
||||
)
|
||||
|
||||
# (A) Full loss.
|
||||
m_full = ScaleModule(c=c, init=init).to(device)
|
||||
fully_shard(m_full, mesh=mesh, reshard_after_forward=True, mp_policy=mp)
|
||||
pred_full = m_full(x)
|
||||
loss_full = ((pred_full - target) ** 2).mean()
|
||||
loss_full.backward()
|
||||
g_full = m_full.scale.grad
|
||||
g_full_local = g_full.to_local() if hasattr(g_full, "to_local") else g_full
|
||||
g_full_vec = _gather_full_vec_allranks(g_full_local.flatten())
|
||||
|
||||
# (B) SP-sharded loss.
|
||||
m_sp = ScaleModule(c=c, init=init).to(device)
|
||||
fully_shard(m_sp, mesh=mesh, reshard_after_forward=True, mp_policy=mp)
|
||||
pred = m_sp(x)
|
||||
sp = get_sp_group().world_size # 2
|
||||
sharded_pred = shard_latents_across_sp(pred)
|
||||
sharded_target = shard_latents_across_sp(target)
|
||||
local_sse = ((sharded_pred - sharded_target) ** 2).sum()
|
||||
loss_sp = (sp * local_sse) / pred.numel()
|
||||
loss_sp.backward()
|
||||
g_sp = m_sp.scale.grad
|
||||
g_sp_local = g_sp.to_local() if hasattr(g_sp, "to_local") else g_sp
|
||||
g_sp_vec = _gather_full_vec_allranks(g_sp_local.flatten())
|
||||
|
||||
torch.testing.assert_close(g_sp_vec, g_full_vec, rtol=1e-4, atol=1e-5)
|
||||
|
||||
|
||||
@@ -0,0 +1,152 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Mismatch compatibility test (direction A):
|
||||
- WORLD_SIZE=4
|
||||
- SP: sp_size=4
|
||||
- HSDP/FSDP2 mesh: replicate_dim=2, shard_dim=2 (mesh_shape=(2,2))
|
||||
|
||||
We verify gradients from:
|
||||
(A) full MSE: mean((pred-target)^2)
|
||||
(B) SP-sharded loss: sp_size * local_sse / total_numel
|
||||
match under this mismatch.
|
||||
|
||||
Run:
|
||||
source /home/hao_lab/miniconda3/etc/profile.d/conda.sh
|
||||
conda activate alexfv
|
||||
torchrun --standalone --nproc_per_node=4 -m pytest -q fastvideo/tests/distributed/test_hsdp_sp_mismatch_sp4_hsdp_shard2.py
|
||||
"""
|
||||
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.distributed import cleanup_dist_env_and_memory
|
||||
from fastvideo.distributed.parallel_state import (
|
||||
get_sp_group,
|
||||
maybe_init_distributed_environment_and_model_parallel,
|
||||
)
|
||||
from fastvideo.training.training_utils import shard_latents_across_sp
|
||||
|
||||
|
||||
def _world_size() -> int:
|
||||
return int(os.environ.get("WORLD_SIZE", "1"))
|
||||
|
||||
|
||||
def _rank() -> int:
|
||||
return int(os.environ.get("RANK", "0"))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def dist_setup():
|
||||
if _world_size() != 4:
|
||||
pytest.skip("Designed for torchrun WORLD_SIZE=4")
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("Requires CUDA")
|
||||
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
|
||||
torch.cuda.set_device(local_rank)
|
||||
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=4)
|
||||
assert get_sp_group().world_size == 4
|
||||
yield
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
dist.barrier()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
class ScaleModule(torch.nn.Module):
|
||||
def __init__(self, c: int, init: torch.Tensor):
|
||||
super().__init__()
|
||||
assert init.shape == (c,)
|
||||
self.scale = torch.nn.Parameter(init.clone())
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
return x * self.scale.view(1, -1, 1, 1, 1)
|
||||
|
||||
|
||||
def _broadcast_(t: torch.Tensor, src: int = 0) -> None:
|
||||
dist.broadcast(t, src=src)
|
||||
|
||||
|
||||
def _reconstruct_full_vec_from_replicate0(local_shard: torch.Tensor) -> torch.Tensor:
|
||||
"""
|
||||
Mesh layout: [[0,2],[1,3]] with dim_names=("replicate","shard").
|
||||
replicate0 group is ranks [0,2], and shard order inside replicate0 is [0,2].
|
||||
Reconstruct a single full vector from replicate0 only, then broadcast to all ranks.
|
||||
"""
|
||||
group_ranks = [0, 2]
|
||||
pg = dist.new_group(ranks=group_ranks)
|
||||
if dist.get_rank() in group_ranks:
|
||||
gathered = [torch.empty_like(local_shard) for _ in range(len(group_ranks))]
|
||||
dist.all_gather(gathered, local_shard, group=pg)
|
||||
full = torch.cat(gathered, dim=0)
|
||||
else:
|
||||
full = torch.empty((local_shard.numel() * len(group_ranks),),
|
||||
device=local_shard.device,
|
||||
dtype=local_shard.dtype)
|
||||
dist.broadcast(full, src=0)
|
||||
return full
|
||||
|
||||
|
||||
def test_sp_sharded_loss_matches_full_mse_under_mismatch(dist_setup):
|
||||
from torch.distributed.device_mesh import DeviceMesh
|
||||
from torch.distributed.fsdp import MixedPrecisionPolicy, fully_shard
|
||||
|
||||
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
|
||||
torch.manual_seed(2026)
|
||||
|
||||
b, c, t, h, w = 2, 8, 1, 1, 8 # thw=8 divisible by sp=4
|
||||
|
||||
x = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
|
||||
if _rank() == 0:
|
||||
x.normal_()
|
||||
_broadcast_(x)
|
||||
target = (0.5 * x).contiguous()
|
||||
|
||||
init = torch.empty((c,), device=device, dtype=torch.float32)
|
||||
if _rank() == 0:
|
||||
init.normal_()
|
||||
_broadcast_(init)
|
||||
|
||||
mesh = DeviceMesh(
|
||||
"cuda",
|
||||
[[0, 2], [1, 3]],
|
||||
mesh_dim_names=("replicate", "shard"),
|
||||
)
|
||||
|
||||
mp = MixedPrecisionPolicy(
|
||||
param_dtype=torch.float32,
|
||||
reduce_dtype=torch.float32,
|
||||
output_dtype=torch.float32,
|
||||
cast_forward_inputs=False,
|
||||
)
|
||||
|
||||
# (A) Full loss.
|
||||
m_full = ScaleModule(c=c, init=init).to(device)
|
||||
fully_shard(m_full, mesh=mesh, reshard_after_forward=True, mp_policy=mp)
|
||||
pred_full = m_full(x)
|
||||
loss_full = ((pred_full - target) ** 2).mean()
|
||||
loss_full.backward()
|
||||
g_full = m_full.scale.grad
|
||||
g_full_local = g_full.to_local() if hasattr(g_full, "to_local") else g_full
|
||||
g_full_vec = _reconstruct_full_vec_from_replicate0(g_full_local.flatten())
|
||||
|
||||
# (B) SP-sharded loss.
|
||||
m_sp = ScaleModule(c=c, init=init).to(device)
|
||||
fully_shard(m_sp, mesh=mesh, reshard_after_forward=True, mp_policy=mp)
|
||||
pred = m_sp(x)
|
||||
sp = get_sp_group().world_size # 4
|
||||
sharded_pred = shard_latents_across_sp(pred)
|
||||
sharded_target = shard_latents_across_sp(target)
|
||||
local_sse = ((sharded_pred - sharded_target) ** 2).sum()
|
||||
loss_sp = (sp * local_sse) / pred.numel()
|
||||
loss_sp.backward()
|
||||
g_sp = m_sp.scale.grad
|
||||
g_sp_local = g_sp.to_local() if hasattr(g_sp, "to_local") else g_sp
|
||||
g_sp_vec = _reconstruct_full_vec_from_replicate0(g_sp_local.flatten())
|
||||
|
||||
torch.testing.assert_close(g_sp_vec, g_full_vec, rtol=1e-4, atol=1e-5)
|
||||
|
||||
|
||||
@@ -0,0 +1,218 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import os
|
||||
|
||||
import pytest
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
|
||||
from fastvideo.distributed import (cleanup_dist_env_and_memory,
|
||||
get_sp_group,
|
||||
maybe_init_distributed_environment_and_model_parallel)
|
||||
from fastvideo.training.training_utils import shard_latents_across_sp
|
||||
|
||||
|
||||
def _world_size() -> int:
|
||||
return int(os.environ.get("WORLD_SIZE", "1"))
|
||||
|
||||
|
||||
def _rank() -> int:
|
||||
return int(os.environ.get("RANK", "0"))
|
||||
|
||||
|
||||
@pytest.fixture(scope="module")
|
||||
def sp_dist():
|
||||
"""
|
||||
These tests are intended to be run under torchrun with multiple GPUs, e.g.
|
||||
torchrun --nproc_per_node=2 -m pytest -q fastvideo/tests/distributed/test_sp_shard_latents_and_loss.py
|
||||
torchrun --nproc_per_node=4 -m pytest -q fastvideo/tests/distributed/test_sp_shard_latents_and_loss.py
|
||||
torchrun --nproc_per_node=8 -m pytest -q fastvideo/tests/distributed/test_sp_shard_latents_and_loss.py
|
||||
"""
|
||||
ws = _world_size()
|
||||
if ws <= 1:
|
||||
pytest.skip("Requires torchrun with WORLD_SIZE>1")
|
||||
if not torch.cuda.is_available():
|
||||
pytest.skip("Requires CUDA")
|
||||
|
||||
local_rank = int(os.environ.get("LOCAL_RANK", "0"))
|
||||
torch.cuda.set_device(local_rank)
|
||||
|
||||
torch.manual_seed(1234)
|
||||
# Initialize FastVideo dist + model-parallel groups.
|
||||
maybe_init_distributed_environment_and_model_parallel(tp_size=1, sp_size=ws)
|
||||
assert get_sp_group().world_size == ws
|
||||
|
||||
yield
|
||||
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
dist.barrier()
|
||||
cleanup_dist_env_and_memory()
|
||||
|
||||
|
||||
def _broadcast_tensor_(t: torch.Tensor, src: int = 0) -> torch.Tensor:
|
||||
if dist.is_available() and dist.is_initialized():
|
||||
dist.broadcast(t, src=src)
|
||||
return t
|
||||
|
||||
|
||||
def _all_gather_cat(t: torch.Tensor, dim: int) -> torch.Tensor:
|
||||
ws = _world_size()
|
||||
gathered = [torch.empty_like(t) for _ in range(ws)]
|
||||
dist.all_gather(gathered, t)
|
||||
return torch.cat(gathered, dim=dim)
|
||||
|
||||
|
||||
def test_shard_latents_t_not_divisible_but_thw_divisible_no_error(sp_dist):
|
||||
"""
|
||||
Regression test for the commit "fix splitting on t":
|
||||
- t is NOT divisible by sp_size
|
||||
- (t*h*w) IS divisible by sp_size
|
||||
- sharding should NOT raise, and shards should round-trip to the original.
|
||||
"""
|
||||
ws = _world_size()
|
||||
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
|
||||
|
||||
b, c, t, h, w = 2, 3, 3, 2, 4 # t=3 (not divisible by 8), thw=24 (divisible by 8)
|
||||
assert (t * h * w) % ws == 0
|
||||
assert t % ws != 0
|
||||
|
||||
latents = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
|
||||
if _rank() == 0:
|
||||
latents.normal_()
|
||||
_broadcast_tensor_(latents)
|
||||
|
||||
shard = shard_latents_across_sp(latents)
|
||||
assert shard.shape == (b, c, (t * h * w) // ws)
|
||||
|
||||
gathered = _all_gather_cat(shard, dim=2)
|
||||
torch.testing.assert_close(gathered, latents.reshape(b, c, t * h * w))
|
||||
|
||||
dist.barrier()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"shape",
|
||||
[
|
||||
# No padding needed: thw divisible by 8, but t is not.
|
||||
(2, 1, 3, 2, 4), # thw=24
|
||||
# Padding needed: thw not divisible by 8.
|
||||
(2, 1, 3, 2, 5), # thw=30 -> pad to 32
|
||||
],
|
||||
)
|
||||
def test_sharded_loss_gradient_matches_full_mse_after_rank_avg(sp_dist, shape):
|
||||
"""
|
||||
Validates the math behind the new loss computation:
|
||||
|
||||
If each rank computes:
|
||||
loss_rank = sp_world_size * local_sse / total_numel
|
||||
and the training stack averages gradients across ranks, then the resulting
|
||||
gradient should match the full MSE gradient.
|
||||
"""
|
||||
ws = _world_size()
|
||||
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
|
||||
b, c, t, h, w = shape
|
||||
|
||||
# Make inputs identical across ranks for deterministic comparison.
|
||||
init = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
|
||||
x = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
|
||||
y = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
|
||||
if _rank() == 0:
|
||||
init.normal_()
|
||||
x.normal_()
|
||||
y.normal_()
|
||||
_broadcast_tensor_(init)
|
||||
_broadcast_tensor_(x)
|
||||
_broadcast_tensor_(y)
|
||||
|
||||
# Full loss baseline (redundant compute; correct reference).
|
||||
w_full = torch.nn.Parameter(init.clone())
|
||||
pred_full = w_full * x
|
||||
full_loss = ((pred_full - y)**2).mean()
|
||||
full_loss.backward()
|
||||
grad_full = w_full.grad.detach().clone()
|
||||
|
||||
# Sharded loss (compute on a shard only).
|
||||
w_shard = torch.nn.Parameter(init.clone())
|
||||
pred = w_shard * x
|
||||
sharded_pred = shard_latents_across_sp(pred)
|
||||
sharded_target = shard_latents_across_sp(y)
|
||||
local_sse = ((sharded_pred - sharded_target)**2).sum()
|
||||
shard_loss = ws * local_sse / pred.numel()
|
||||
shard_loss.backward()
|
||||
grad_shard = w_shard.grad.detach().clone()
|
||||
|
||||
# Simulate "gradient averaging across ranks" (DDP/FSDP replicate-dim behavior).
|
||||
dist.all_reduce(grad_shard, op=dist.ReduceOp.SUM)
|
||||
grad_shard /= ws
|
||||
|
||||
# (Optional) also average grad_full, for symmetry.
|
||||
dist.all_reduce(grad_full, op=dist.ReduceOp.SUM)
|
||||
grad_full /= ws
|
||||
|
||||
torch.testing.assert_close(grad_shard, grad_full, rtol=1e-5, atol=1e-6)
|
||||
|
||||
dist.barrier()
|
||||
|
||||
|
||||
def test_padding_does_not_change_global_sse(sp_dist):
|
||||
"""
|
||||
Padding-specific correctness test.
|
||||
|
||||
When (t*h*w) is NOT divisible by sp_size, shard_latents_across_sp pads with
|
||||
zeros on the flattened axis. This test validates that:
|
||||
1) Summed SSE across SP ranks equals the full (unpadded) SSE.
|
||||
2) Any padded tokens that land on a rank are exactly zero for both pred/target.
|
||||
"""
|
||||
from fastvideo.distributed.utils import compute_padding_for_sp
|
||||
|
||||
ws = _world_size()
|
||||
device = torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}")
|
||||
|
||||
b, c, t, h = 2, 2, 1, 1
|
||||
# Force padding for any ws>1: seq_len = 2*ws + 1 -> remainder 1
|
||||
w = 2 * ws + 1
|
||||
seq_len = t * h * w
|
||||
assert seq_len % ws != 0
|
||||
|
||||
pred = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
|
||||
target = torch.empty((b, c, t, h, w), device=device, dtype=torch.float32)
|
||||
if _rank() == 0:
|
||||
pred.normal_()
|
||||
target.normal_()
|
||||
_broadcast_tensor_(pred)
|
||||
_broadcast_tensor_(target)
|
||||
|
||||
# Full (unpadded) SSE reference.
|
||||
full_sse = ((pred - target)**2).sum()
|
||||
|
||||
# Sharded local SSE (includes padded region, which should contribute 0).
|
||||
sharded_pred = shard_latents_across_sp(pred)
|
||||
sharded_target = shard_latents_across_sp(target)
|
||||
local_sse = ((sharded_pred - sharded_target)**2).sum()
|
||||
|
||||
global_sse = local_sse.clone()
|
||||
dist.all_reduce(global_sse, op=dist.ReduceOp.SUM)
|
||||
torch.testing.assert_close(global_sse, full_sse, rtol=1e-5, atol=1e-6)
|
||||
|
||||
# Explicitly validate padded tokens are zeros on ranks that cover them.
|
||||
padded_seq_len, padding_amount = compute_padding_for_sp(seq_len, ws)
|
||||
assert padding_amount > 0
|
||||
elements_per_rank = padded_seq_len // ws
|
||||
start = _rank() * elements_per_rank
|
||||
end = (_rank() + 1) * elements_per_rank
|
||||
|
||||
if end > seq_len:
|
||||
# There is padding on this rank: positions [max(seq_len, start), end)
|
||||
pad_start_global = max(seq_len, start)
|
||||
pad_end_global = end
|
||||
pad_start_local = pad_start_global - start
|
||||
pad_end_local = pad_end_global - start
|
||||
assert pad_start_local < pad_end_local
|
||||
|
||||
pad_slice_pred = sharded_pred[:, :, pad_start_local:pad_end_local]
|
||||
pad_slice_target = sharded_target[:, :, pad_start_local:pad_end_local]
|
||||
torch.testing.assert_close(pad_slice_pred, torch.zeros_like(pad_slice_pred))
|
||||
torch.testing.assert_close(pad_slice_target, torch.zeros_like(pad_slice_target))
|
||||
|
||||
dist.barrier()
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ import torch
|
||||
import pytest
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.models.dits.matrix_game.utils import create_action_presets
|
||||
from fastvideo.models.dits.matrixgame.utils import create_action_presets
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.tests.utils import (
|
||||
compute_video_ssim_torchvision,
|
||||
|
||||
@@ -0,0 +1,250 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_matrixgame
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.pipelines.basic.matrixgame.matrixgame_i2v_pipeline import (
|
||||
MatrixGamePipeline)
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class MatrixGameTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
A training pipeline for Matrix-Game-2.0.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def set_schemas(self):
|
||||
self.train_dataset_schema = pyarrow_schema_matrixgame
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
args_copy = deepcopy(training_args)
|
||||
|
||||
args_copy.inference_mode = True
|
||||
args_copy.dit_cpu_offload = True
|
||||
# args_copy.pipeline_config.vae_config.load_encoder = False
|
||||
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
|
||||
self.validation_pipeline = MatrixGamePipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=None,
|
||||
inference_mode=True,
|
||||
loaded_modules={
|
||||
"transformer": self.get_module("transformer"),
|
||||
},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
dit_cpu_offload=True)
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
# Reset iterator for next epoch
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
# Get first batch of new epoch
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
latents = batch['vae_latent']
|
||||
latents = latents[:, :, :self.training_args.num_latent_t]
|
||||
# encoder_hidden_states = batch['text_embedding']
|
||||
# encoder_attention_mask = batch['text_attention_mask']
|
||||
clip_features = batch['clip_feature']
|
||||
image_latents = batch['first_frame_latent']
|
||||
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
|
||||
pil_image = batch['pil_image']
|
||||
infos = batch['info_list']
|
||||
|
||||
training_batch.latents = latents.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = None
|
||||
training_batch.encoder_attention_mask = None
|
||||
# MatrixGame doesn't use text encoder
|
||||
training_batch.preprocessed_image = pil_image.to(
|
||||
get_local_torch_device())
|
||||
training_batch.image_embeds = clip_features.to(get_local_torch_device())
|
||||
training_batch.image_latents = image_latents.to(
|
||||
get_local_torch_device())
|
||||
training_batch.infos = infos
|
||||
|
||||
# Action conditioning
|
||||
if 'mouse_cond' in batch and batch['mouse_cond'].numel() > 0:
|
||||
training_batch.mouse_cond = batch['mouse_cond'].to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
else:
|
||||
training_batch.mouse_cond = None
|
||||
|
||||
if 'keyboard_cond' in batch and batch['keyboard_cond'].numel() > 0:
|
||||
training_batch.keyboard_cond = batch['keyboard_cond'].to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
else:
|
||||
training_batch.keyboard_cond = None
|
||||
|
||||
return training_batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
|
||||
|
||||
# First, call parent method to prepare noise, timesteps, etc. for video latents
|
||||
training_batch = super()._prepare_dit_inputs(training_batch)
|
||||
|
||||
assert isinstance(training_batch.image_latents, torch.Tensor)
|
||||
image_latents = training_batch.image_latents.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
|
||||
temporal_compression_ratio = self.training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (self.training_args.num_latent_t -
|
||||
1) * temporal_compression_ratio + 1
|
||||
batch_size, num_channels, _, latent_height, latent_width = image_latents.shape
|
||||
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
|
||||
latent_width)
|
||||
mask_lat_size[:, :, 1:] = 0
|
||||
|
||||
first_frame_mask = mask_lat_size[:, :, :1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask, dim=2, repeats=temporal_compression_ratio)
|
||||
mask_lat_size = torch.cat([first_frame_mask, mask_lat_size[:, :, 1:]],
|
||||
dim=2)
|
||||
mask_lat_size = mask_lat_size.view(batch_size, -1,
|
||||
temporal_compression_ratio,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
mask_lat_size = mask_lat_size.to(
|
||||
image_latents.device).to(dtype=torch.bfloat16)
|
||||
|
||||
training_batch.noisy_model_input = torch.cat(
|
||||
[training_batch.noisy_model_input, mask_lat_size, image_latents],
|
||||
dim=1)
|
||||
|
||||
return training_batch
|
||||
|
||||
def _build_input_kwargs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
|
||||
# Image Embeds for conditioning
|
||||
image_embeds = training_batch.image_embeds
|
||||
assert torch.isnan(image_embeds).sum() == 0
|
||||
image_embeds = image_embeds.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
encoder_hidden_states_image = image_embeds
|
||||
|
||||
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
|
||||
training_batch.input_kwargs = {
|
||||
"hidden_states":
|
||||
training_batch.noisy_model_input,
|
||||
"encoder_hidden_states":
|
||||
training_batch.encoder_hidden_states, # None for MatrixGame
|
||||
"timestep":
|
||||
training_batch.timesteps.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16),
|
||||
# "encoder_attention_mask":
|
||||
# training_batch.encoder_attention_mask,
|
||||
"encoder_hidden_states_image":
|
||||
encoder_hidden_states_image,
|
||||
# Action conditioning
|
||||
"mouse_cond":
|
||||
training_batch.mouse_cond,
|
||||
"keyboard_cond":
|
||||
training_batch.keyboard_cond,
|
||||
"return_dict":
|
||||
False,
|
||||
}
|
||||
return training_batch
|
||||
|
||||
def _prepare_validation_batch(self, sampling_param: SamplingParam,
|
||||
training_args: TrainingArgs,
|
||||
validation_batch: dict[str, Any],
|
||||
num_inference_steps: int) -> ForwardBatch:
|
||||
sampling_param.prompt = validation_batch['prompt']
|
||||
sampling_param.height = training_args.num_height
|
||||
sampling_param.width = training_args.num_width
|
||||
sampling_param.image_path = validation_batch.get(
|
||||
'image_path') or validation_batch.get('video_path')
|
||||
sampling_param.num_inference_steps = num_inference_steps
|
||||
sampling_param.data_type = "video"
|
||||
assert self.seed is not None
|
||||
sampling_param.seed = self.seed
|
||||
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8, sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (training_args.num_latent_t -
|
||||
1) * temporal_compression_factor + 1
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
VSA_sparsity=training_args.VSA_sparsity,
|
||||
)
|
||||
if "image" in validation_batch and validation_batch["image"] is not None:
|
||||
batch.pil_image = validation_batch["image"]
|
||||
|
||||
if "keyboard_cond" in validation_batch and validation_batch[
|
||||
"keyboard_cond"] is not None:
|
||||
keyboard_cond = validation_batch["keyboard_cond"]
|
||||
keyboard_cond = torch.tensor(keyboard_cond, dtype=torch.bfloat16)
|
||||
keyboard_cond = keyboard_cond.unsqueeze(0)
|
||||
batch.keyboard_cond = keyboard_cond
|
||||
|
||||
if "mouse_cond" in validation_batch and validation_batch[
|
||||
"mouse_cond"] is not None:
|
||||
mouse_cond = validation_batch["mouse_cond"]
|
||||
mouse_cond = torch.tensor(mouse_cond, dtype=torch.bfloat16)
|
||||
mouse_cond = mouse_cond.unsqueeze(0)
|
||||
batch.mouse_cond = mouse_cond
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting training pipeline...")
|
||||
|
||||
pipeline = MatrixGameTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.dit_cpu_offload = False
|
||||
main(args)
|
||||
@@ -443,11 +443,26 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
# make sure no implicit broadcasting happens
|
||||
assert model_pred.shape == target.shape, f"model_pred.shape: {model_pred.shape}, target.shape: {target.shape}"
|
||||
|
||||
sharded_pred = shard_latents_across_sp(model_pred)
|
||||
sharded_target = shard_latents_across_sp(target)
|
||||
loss = (torch.mean(
|
||||
(sharded_pred.float() - sharded_target.float())**2) /
|
||||
self.training_args.gradient_accumulation_steps)
|
||||
# Defensive: avoid NaNs/div0 if an upstream bug ever produces empty tensors.
|
||||
# mean(empty) -> NaN; division by 0 -> inf/NaN. Keep a 0 loss with grad history.
|
||||
if model_pred.numel() == 0:
|
||||
loss = (model_pred.sum() *
|
||||
0.0) / self.training_args.gradient_accumulation_steps
|
||||
else:
|
||||
# Compute MSE on one SP shard. We shard along the flattened token axis
|
||||
# (t*h*w) with optional padding so (t*h*w) need not be divisible by sp_size.
|
||||
sp_world_size = get_sp_group().world_size
|
||||
if sp_world_size > 1:
|
||||
sharded_pred = shard_latents_across_sp(model_pred)
|
||||
sharded_target = shard_latents_across_sp(target)
|
||||
local_sse = ((sharded_pred.float() -
|
||||
sharded_target.float())**2).sum()
|
||||
loss = (sp_world_size * local_sse / model_pred.numel()
|
||||
) / self.training_args.gradient_accumulation_steps
|
||||
else:
|
||||
loss = (torch.mean(
|
||||
(model_pred.float() - target.float())**2) /
|
||||
self.training_args.gradient_accumulation_steps)
|
||||
|
||||
loss.backward()
|
||||
|
||||
@@ -654,7 +669,11 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
|
||||
training_batch.raw_latent_shape[3] //
|
||||
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
|
||||
context_len = int(training_batch.encoder_hidden_states.shape[1])
|
||||
if training_batch.encoder_hidden_states is not None:
|
||||
context_len = int(
|
||||
training_batch.encoder_hidden_states.shape[1])
|
||||
else:
|
||||
context_len = 0
|
||||
|
||||
metrics["dit_seq_len"] = int(seq_len)
|
||||
metrics["context_len"] = context_len
|
||||
|
||||
@@ -20,9 +20,10 @@ from fastvideo.training.checkpointing_utils import (ModelWrapper,
|
||||
RandomStateWrapper,
|
||||
SchedulerWrapper)
|
||||
|
||||
from einops import rearrange
|
||||
from fastvideo.distributed.parallel_state import (get_sp_parallel_rank,
|
||||
get_sp_world_size)
|
||||
from fastvideo.distributed.utils import (compute_padding_for_sp,
|
||||
pad_sequence_tensor)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -867,10 +868,25 @@ def shard_latents_across_sp(latents: torch.Tensor) -> torch.Tensor:
|
||||
sp_world_size = get_sp_world_size()
|
||||
rank_in_sp_group = get_sp_parallel_rank()
|
||||
if sp_world_size > 1:
|
||||
latents = rearrange(latents,
|
||||
"b c (n s) h w -> b c n s h w",
|
||||
n=sp_world_size).contiguous()
|
||||
latents = latents[:, :, rank_in_sp_group, :, :, :]
|
||||
# Shard on the flattened token axis (t*h*w) rather than raw `t`, so we
|
||||
# don't require t % sp_world_size == 0. Pad on the flattened axis if needed.
|
||||
assert latents.ndim == 5, f"Expected latents [b,c,t,h,w], got {latents.shape}"
|
||||
b, c, t, h, w = latents.shape
|
||||
latents = latents.reshape(b, c, t * h * w)
|
||||
|
||||
original_seq_len = latents.shape[2]
|
||||
padded_seq_len, padding_amount = compute_padding_for_sp(
|
||||
original_seq_len, sp_world_size)
|
||||
if padding_amount > 0:
|
||||
latents = pad_sequence_tensor(latents,
|
||||
padded_seq_len,
|
||||
seq_dim=2,
|
||||
pad_value=0.0)
|
||||
|
||||
elements_per_rank = padded_seq_len // sp_world_size
|
||||
start = rank_in_sp_group * elements_per_rank
|
||||
end = (rank_in_sp_group + 1) * elements_per_rank
|
||||
latents = latents[:, :, start:end].contiguous()
|
||||
return latents
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,250 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
import sys
|
||||
from copy import deepcopy
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_wangame
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs, TrainingArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.schedulers.scheduling_flow_unipc_multistep import (
|
||||
FlowUniPCMultistepScheduler)
|
||||
from fastvideo.pipelines.basic.wan.wangame_i2v_pipeline import WanGameActionImageToVideoPipeline
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class WanGameTrainingPipeline(TrainingPipeline):
|
||||
"""
|
||||
A training pipeline for WanGame-2.1-Fun-1.3B-InP.
|
||||
"""
|
||||
_required_config_modules = ["scheduler", "transformer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
self.modules["scheduler"] = FlowUniPCMultistepScheduler(
|
||||
shift=fastvideo_args.pipeline_config.flow_shift)
|
||||
|
||||
def create_training_stages(self, training_args: TrainingArgs):
|
||||
"""
|
||||
May be used in future refactors.
|
||||
"""
|
||||
pass
|
||||
|
||||
def set_schemas(self):
|
||||
self.train_dataset_schema = pyarrow_schema_wangame
|
||||
|
||||
def initialize_validation_pipeline(self, training_args: TrainingArgs):
|
||||
logger.info("Initializing validation pipeline...")
|
||||
# args_copy.pipeline_config.vae_config.load_encoder = False
|
||||
# validation_pipeline = WanImageToVideoValidationPipeline.from_pretrained(
|
||||
self.validation_pipeline = WanGameActionImageToVideoPipeline.from_pretrained(
|
||||
training_args.model_path,
|
||||
args=None,
|
||||
inference_mode=True,
|
||||
loaded_modules={
|
||||
"transformer": self.get_module("transformer"),
|
||||
},
|
||||
tp_size=training_args.tp_size,
|
||||
sp_size=training_args.sp_size,
|
||||
num_gpus=training_args.num_gpus,
|
||||
dit_cpu_offload=False)
|
||||
|
||||
def _get_next_batch(self, training_batch: TrainingBatch) -> TrainingBatch:
|
||||
batch = next(self.train_loader_iter, None) # type: ignore
|
||||
if batch is None:
|
||||
self.current_epoch += 1
|
||||
logger.info("Starting epoch %s", self.current_epoch)
|
||||
# Reset iterator for next epoch
|
||||
self.train_loader_iter = iter(self.train_dataloader)
|
||||
# Get first batch of new epoch
|
||||
batch = next(self.train_loader_iter)
|
||||
|
||||
latents = batch['vae_latent']
|
||||
latents = latents[:, :, :self.training_args.num_latent_t]
|
||||
# encoder_hidden_states = batch['text_embedding']
|
||||
# encoder_attention_mask = batch['text_attention_mask']
|
||||
clip_features = batch['clip_feature']
|
||||
image_latents = batch['first_frame_latent']
|
||||
image_latents = image_latents[:, :, :self.training_args.num_latent_t]
|
||||
pil_image = batch['pil_image']
|
||||
infos = batch['info_list']
|
||||
|
||||
training_batch.latents = latents.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
training_batch.encoder_hidden_states = None
|
||||
training_batch.encoder_attention_mask = None
|
||||
# MatrixGame doesn't use text encoder
|
||||
training_batch.preprocessed_image = pil_image.to(
|
||||
get_local_torch_device())
|
||||
training_batch.image_embeds = clip_features.to(get_local_torch_device())
|
||||
training_batch.image_latents = image_latents.to(
|
||||
get_local_torch_device())
|
||||
training_batch.infos = infos
|
||||
|
||||
# Action conditioning
|
||||
if 'mouse_cond' in batch and batch['mouse_cond'].numel() > 0:
|
||||
training_batch.mouse_cond = batch['mouse_cond'].to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
else:
|
||||
training_batch.mouse_cond = None
|
||||
|
||||
if 'keyboard_cond' in batch and batch['keyboard_cond'].numel() > 0:
|
||||
training_batch.keyboard_cond = batch['keyboard_cond'].to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
else:
|
||||
training_batch.keyboard_cond = None
|
||||
|
||||
return training_batch
|
||||
|
||||
def _prepare_dit_inputs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
"""Override to properly handle I2V concatenation - call parent first, then concatenate image conditioning."""
|
||||
|
||||
# First, call parent method to prepare noise, timesteps, etc. for video latents
|
||||
training_batch = super()._prepare_dit_inputs(training_batch)
|
||||
|
||||
assert isinstance(training_batch.image_latents, torch.Tensor)
|
||||
image_latents = training_batch.image_latents.to(
|
||||
get_local_torch_device(), dtype=torch.bfloat16)
|
||||
|
||||
temporal_compression_ratio = self.training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (self.training_args.num_latent_t -
|
||||
1) * temporal_compression_ratio + 1
|
||||
batch_size, num_channels, _, latent_height, latent_width = image_latents.shape
|
||||
mask_lat_size = torch.ones(batch_size, 1, num_frames, latent_height,
|
||||
latent_width)
|
||||
mask_lat_size[:, :, 1:] = 0
|
||||
|
||||
first_frame_mask = mask_lat_size[:, :, :1]
|
||||
first_frame_mask = torch.repeat_interleave(
|
||||
first_frame_mask, dim=2, repeats=temporal_compression_ratio)
|
||||
mask_lat_size = torch.cat([first_frame_mask, mask_lat_size[:, :, 1:]],
|
||||
dim=2)
|
||||
mask_lat_size = mask_lat_size.view(batch_size, -1,
|
||||
temporal_compression_ratio,
|
||||
latent_height, latent_width)
|
||||
mask_lat_size = mask_lat_size.transpose(1, 2)
|
||||
mask_lat_size = mask_lat_size.to(
|
||||
image_latents.device).to(dtype=torch.bfloat16)
|
||||
|
||||
training_batch.noisy_model_input = torch.cat(
|
||||
[training_batch.noisy_model_input, mask_lat_size, image_latents],
|
||||
dim=1)
|
||||
|
||||
return training_batch
|
||||
|
||||
def _build_input_kwargs(self,
|
||||
training_batch: TrainingBatch) -> TrainingBatch:
|
||||
|
||||
# Image Embeds for conditioning
|
||||
image_embeds = training_batch.image_embeds
|
||||
assert torch.isnan(image_embeds).sum() == 0
|
||||
image_embeds = image_embeds.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16)
|
||||
encoder_hidden_states_image = image_embeds
|
||||
|
||||
from fastvideo.models.dits.hyworld.pose import process_custom_actions
|
||||
viewmats, intrinsics, action_labels = process_custom_actions(training_batch.keyboard_cond, training_batch.mouse_cond)
|
||||
viewmats = viewmats.unsqueeze(0).to(get_local_torch_device(), dtype=torch.bfloat16)
|
||||
intrinsics = intrinsics.unsqueeze(0).to(get_local_torch_device(), dtype=torch.bfloat16)
|
||||
action_labels = action_labels.unsqueeze(0).to(get_local_torch_device(), dtype=torch.bfloat16)
|
||||
|
||||
# NOTE: noisy_model_input already contains concatenated image_latents from _prepare_dit_inputs
|
||||
training_batch.input_kwargs = {
|
||||
"hidden_states":
|
||||
training_batch.noisy_model_input,
|
||||
"encoder_hidden_states":
|
||||
training_batch.encoder_hidden_states, # None for MatrixGame
|
||||
"timestep":
|
||||
training_batch.timesteps.to(get_local_torch_device(),
|
||||
dtype=torch.bfloat16),
|
||||
# "encoder_attention_mask":
|
||||
# training_batch.encoder_attention_mask,
|
||||
"encoder_hidden_states_image":
|
||||
encoder_hidden_states_image,
|
||||
# Action conditioning
|
||||
"viewmats": viewmats,
|
||||
"Ks": intrinsics,
|
||||
"action": action_labels,
|
||||
"return_dict":
|
||||
False,
|
||||
}
|
||||
return training_batch
|
||||
|
||||
def _prepare_validation_batch(self, sampling_param: SamplingParam,
|
||||
training_args: TrainingArgs,
|
||||
validation_batch: dict[str, Any],
|
||||
num_inference_steps: int) -> ForwardBatch:
|
||||
sampling_param.prompt = validation_batch['prompt']
|
||||
sampling_param.height = training_args.num_height
|
||||
sampling_param.width = training_args.num_width
|
||||
sampling_param.image_path = validation_batch.get(
|
||||
'image_path') or validation_batch.get('video_path')
|
||||
sampling_param.num_inference_steps = num_inference_steps
|
||||
sampling_param.data_type = "video"
|
||||
assert self.seed is not None
|
||||
sampling_param.seed = self.seed
|
||||
|
||||
latents_size = [(sampling_param.num_frames - 1) // 4 + 1,
|
||||
sampling_param.height // 8, sampling_param.width // 8]
|
||||
n_tokens = latents_size[0] * latents_size[1] * latents_size[2]
|
||||
temporal_compression_factor = training_args.pipeline_config.vae_config.arch_config.temporal_compression_ratio
|
||||
num_frames = (training_args.num_latent_t -
|
||||
1) * temporal_compression_factor + 1
|
||||
sampling_param.num_frames = num_frames
|
||||
batch = ForwardBatch(
|
||||
**shallow_asdict(sampling_param),
|
||||
latents=None,
|
||||
generator=torch.Generator(device="cpu").manual_seed(self.seed),
|
||||
n_tokens=n_tokens,
|
||||
eta=0.0,
|
||||
VSA_sparsity=training_args.VSA_sparsity,
|
||||
)
|
||||
if "image" in validation_batch and validation_batch["image"] is not None:
|
||||
batch.pil_image = validation_batch["image"]
|
||||
|
||||
if "keyboard_cond" in validation_batch and validation_batch[
|
||||
"keyboard_cond"] is not None:
|
||||
keyboard_cond = validation_batch["keyboard_cond"]
|
||||
keyboard_cond = torch.tensor(keyboard_cond, dtype=torch.bfloat16)
|
||||
keyboard_cond = keyboard_cond.unsqueeze(0)
|
||||
batch.keyboard_cond = keyboard_cond
|
||||
|
||||
if "mouse_cond" in validation_batch and validation_batch[
|
||||
"mouse_cond"] is not None:
|
||||
mouse_cond = validation_batch["mouse_cond"]
|
||||
mouse_cond = torch.tensor(mouse_cond, dtype=torch.bfloat16)
|
||||
mouse_cond = mouse_cond.unsqueeze(0)
|
||||
batch.mouse_cond = mouse_cond
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
def main(args) -> None:
|
||||
logger.info("Starting training pipeline...")
|
||||
|
||||
pipeline = WanGameTrainingPipeline.from_pretrained(
|
||||
args.pretrained_model_name_or_path, args=args)
|
||||
args = pipeline.training_args
|
||||
pipeline.train()
|
||||
logger.info("Training pipeline done")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
argv = sys.argv
|
||||
from fastvideo.fastvideo_args import TrainingArgs
|
||||
from fastvideo.utils import FlexibleArgumentParser
|
||||
parser = FlexibleArgumentParser()
|
||||
parser = TrainingArgs.add_cli_args(parser)
|
||||
parser = FastVideoArgs.add_cli_args(parser)
|
||||
args = parser.parse_args()
|
||||
args.dit_cpu_offload = False
|
||||
main(args)
|
||||
@@ -152,9 +152,10 @@ nav:
|
||||
- Overview: attention/index.md
|
||||
- Video Sparse Attention: attention/vsa/index.md
|
||||
- Sliding Tile Attention: attention/sta/index.md
|
||||
- Adding a New Attention Backend: attention/developer/index.md
|
||||
- Backend Development: contributing/attention_backend.md
|
||||
- Utilities:
|
||||
- LoRA: utilities/lora.md
|
||||
- Debugging: utilities/debugging.md
|
||||
- Design:
|
||||
- Overview: design/overview.md
|
||||
- Developer Guide:
|
||||
@@ -165,7 +166,8 @@ nav:
|
||||
- RunPod: contributing/developer_env/runpod.md
|
||||
- Testing: contributing/testing.md
|
||||
- Profiling: contributing/profiling.md
|
||||
- Adding a New Attention Backend: attention/developer/index.md
|
||||
- Coding Agents: contributing/coding_agents.md
|
||||
- Attention Backend Development: contributing/attention_backend.md
|
||||
- API Reference:
|
||||
- FastVideo: api/fastvideo.md
|
||||
|
||||
|
||||