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