Compare commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
d03ffcb5f7 | ||
|
|
31c0f1b341 | ||
|
|
4bee0fa199 | ||
|
|
530e6b8363 | ||
|
|
9ab2725db1 | ||
|
|
0aff68f51d | ||
|
|
ad58f802f3 | ||
|
|
04fa356ee3 | ||
|
|
f9c076fe2b |
@@ -32,6 +32,11 @@ env
|
||||
**.txt
|
||||
*.log
|
||||
weights/
|
||||
official_weights/
|
||||
converted_weights/
|
||||
|
||||
# SSIM test outputs
|
||||
fastvideo/tests/ssim/generated_videos/
|
||||
|
||||
# Distribution / packaging
|
||||
build/
|
||||
@@ -69,6 +74,11 @@ docs/distillation/examples/
|
||||
!docs/assets/images/**/*.png
|
||||
!comfyui/assets/**/*.png
|
||||
!comfyui/assets/**/*.gif
|
||||
!assets/images/**/*.png
|
||||
!assets/images/**/*.jpg
|
||||
!assets/images/**/*.jpeg
|
||||
!assets/images/**/*.gif
|
||||
!assets/videos/**/*.mp4
|
||||
|
||||
dmd_t2v_output/
|
||||
preprocess_output_text/
|
||||
|
||||
@@ -10,7 +10,7 @@ exclude: |
|
||||
demo/.*|
|
||||
predict\.py|
|
||||
scripts/.*|
|
||||
prompts/.*|
|
||||
assets/prompts/.*|
|
||||
fastvideo/data_preprocess/.*|
|
||||
fastvideo/dataset/.*|
|
||||
fastvideo/models/.*|
|
||||
|
||||
@@ -0,0 +1,42 @@
|
||||
# Repository Guidelines
|
||||
|
||||
## Project Structure & Module Organization
|
||||
- Core Python package: `fastvideo/` (models, pipelines, training, distributed runtime, CLI entrypoints).
|
||||
- CUDA/custom kernels: `fastvideo-kernel/` (separate build/test flow).
|
||||
- Tests:
|
||||
- `fastvideo/tests/` for package-level tests (dataset, encoders, inference, training, SSIM, workflow).
|
||||
- `tests/local_tests/` for additional local/component checks.
|
||||
- Docs and guides: `docs/` (MkDocs source), with contributor docs in `docs/contributing/`.
|
||||
- Runnable examples and scripts: `examples/` and `scripts/`.
|
||||
- Static assets: `assets/` (including `assets/images/`, `assets/videos/`, and `assets/prompts/`) and `comfyui/assets/`.
|
||||
|
||||
## Build, Test, and Development Commands
|
||||
- `uv pip install -e .[dev]`: editable install with lint/test extras.
|
||||
- `pre-commit install --hook-type pre-commit --hook-type commit-msg`: enable local hooks.
|
||||
- `pre-commit run --all-files`: run formatter/lint/type/spelling checks.
|
||||
- `pytest tests/`: run top-level test suite.
|
||||
- `pytest fastvideo/tests/ -v`: run package tests.
|
||||
- `pytest fastvideo/tests/ssim/ -vs`: run SSIM regression tests (GPU-heavy).
|
||||
- `cd fastvideo-kernel && ./build.sh`: build kernel extensions.
|
||||
|
||||
## Coding Style & Naming Conventions
|
||||
- Python 3.10+; 4-space indentation; keep code and imports readable and explicit.
|
||||
- Style tools are configured in `pyproject.toml` and `.pre-commit-config.yaml`:
|
||||
- `yapf` (format), `ruff` (lint, auto-fix), `mypy` (typing), `codespell`.
|
||||
- Target line length is 80.
|
||||
- Naming: `snake_case` for functions/files, `PascalCase` for classes, `UPPER_SNAKE_CASE` for constants.
|
||||
|
||||
## Testing Guidelines
|
||||
- Use `pytest` and place tests near relevant domains (e.g., `fastvideo/tests/encoders/`).
|
||||
- Prefer descriptive names like `test_<feature>_<expected_behavior>.py`.
|
||||
- For new pipelines/backends, include at least one regression-oriented test; add SSIM coverage when output quality must be preserved.
|
||||
- Document GPU assumptions in tests that require specific hardware.
|
||||
|
||||
## Commit & Pull Request Guidelines
|
||||
- Follow existing commit style: short subject with optional tag prefix, e.g. `[bugfix]: ...`, `[feat]: ...`, `[misc]: ...`, and include PR reference like `(#1234)` when applicable.
|
||||
- Keep commits focused by concern (feature, refactor, fix).
|
||||
- PRs should include:
|
||||
- clear problem/solution summary,
|
||||
- test evidence (`pytest`/SSIM outputs or rationale if skipped),
|
||||
- linked issue/PR context,
|
||||
- screenshots or sample outputs for UI/demo/docs changes.
|
||||
@@ -1,5 +1,10 @@
|
||||
<div align="center">
|
||||
<img src=assets/logos/logo.svg width="30%"/>
|
||||
</div>
|
||||
|
||||
| **[Documentation](https://hao-ai-lab.github.io/FastVideo)** | **[Quick Start](https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/)** | **[Weekly Dev Meeting](https://github.com/hao-ai-lab/FastVideo/discussions/982)** | 🟣💬 **[Slack](https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ)** |
|
||||
<p align="center">
|
||||
| <a href="https://hao-ai-lab.github.io/FastVideo"><b>Documentation</b></a> | <a href="https://hao-ai-lab.github.io/FastVideo/inference/inference_quick_start/"><b> Quick Start</b></a> | <a href="https://github.com/hao-ai-lab/FastVideo/discussions/982" target="_blank"><b>Weekly Dev Meeting</b></a> | 🟣💬 <a href="https://join.slack.com/t/fastvideo/shared_invite/zt-3f4lao1uq-u~Ipx6Lt4J27AlD2y~IdLQ" target="_blank"> <b>Slack</b> </a> | 🟣💬 <a href="https://github.com/hao-ai-lab/FastVideo/discussions/1097" target="_blank"> <b> WeChat </b> </a> |
|
||||
</p>
|
||||
|
||||
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
|
||||
|
||||
|
||||
|
Before Width: | Height: | Size: 113 KiB After Width: | Height: | Size: 113 KiB |
|
Before Width: | Height: | Size: 1.2 MiB After Width: | Height: | Size: 1.2 MiB |
|
Before Width: | Height: | Size: 229 KiB After Width: | Height: | Size: 229 KiB |
|
Before Width: | Height: | Size: 168 KiB After Width: | Height: | Size: 168 KiB |
|
Before Width: | Height: | Size: 103 KiB After Width: | Height: | Size: 103 KiB |
|
Before Width: | Height: | Size: 148 KiB After Width: | Height: | Size: 148 KiB |
|
Before Width: | Height: | Size: 155 KiB After Width: | Height: | Size: 155 KiB |
|
Before Width: | Height: | Size: 723 KiB After Width: | Height: | Size: 723 KiB |
|
Before Width: | Height: | Size: 723 KiB After Width: | Height: | Size: 723 KiB |
|
Before Width: | Height: | Size: 875 KiB After Width: | Height: | Size: 875 KiB |
|
Before Width: | Height: | Size: 664 KiB After Width: | Height: | Size: 664 KiB |
|
Before Width: | Height: | Size: 62 KiB After Width: | Height: | Size: 62 KiB |
|
Before Width: | Height: | Size: 686 KiB After Width: | Height: | Size: 686 KiB |
|
Before Width: | Height: | Size: 957 KiB After Width: | Height: | Size: 957 KiB |
|
Before Width: | Height: | Size: 585 KiB After Width: | Height: | Size: 585 KiB |
|
Before Width: | Height: | Size: 558 KiB After Width: | Height: | Size: 558 KiB |
|
Before Width: | Height: | Size: 942 KiB After Width: | Height: | Size: 942 KiB |
|
Before Width: | Height: | Size: 890 KiB After Width: | Height: | Size: 890 KiB |
|
Before Width: | Height: | Size: 433 KiB After Width: | Height: | Size: 433 KiB |
|
Before Width: | Height: | Size: 595 KiB After Width: | Height: | Size: 595 KiB |
|
Before Width: | Height: | Size: 781 KiB After Width: | Height: | Size: 781 KiB |
|
Before Width: | Height: | Size: 783 KiB After Width: | Height: | Size: 783 KiB |
|
Before Width: | Height: | Size: 762 KiB After Width: | Height: | Size: 762 KiB |
|
Before Width: | Height: | Size: 68 KiB After Width: | Height: | Size: 68 KiB |
|
Before Width: | Height: | Size: 147 KiB After Width: | Height: | Size: 147 KiB |
|
Before Width: | Height: | Size: 89 KiB After Width: | Height: | Size: 89 KiB |
|
Before Width: | Height: | Size: 133 KiB After Width: | Height: | Size: 133 KiB |
|
Before Width: | Height: | Size: 213 KiB After Width: | Height: | Size: 213 KiB |
@@ -1,110 +1,110 @@
|
||||
[
|
||||
{
|
||||
"prompt": "Young man skating with a skateboard on the ramps with graffiti of a park with trees, on a sunny day.",
|
||||
"image_path": "images/mixkit-boy-skating-with-a-skateboard-in-a-park-with-ramps-34389.png"
|
||||
"image_path": "assets/images/mixkit-boy-skating-with-a-skateboard-in-a-park-with-ramps-34389.png"
|
||||
},
|
||||
{
|
||||
"prompt": "In the midst of the joyous New Year's Eve celebration, the cheerful group of friends, their spirits lifted by the festivities, decides to immortalize the moment with a vibrant snapshot",
|
||||
"image_path": "images/mixkit-a-cheerful-group-of-friends-celebrate-new-years-eve-and-51525.png"
|
||||
"image_path": "assets/images/mixkit-a-cheerful-group-of-friends-celebrate-new-years-eve-and-51525.png"
|
||||
},
|
||||
{
|
||||
"prompt": "A man and a woman playing in a field with grass, during a bright afternoon, while cars pass by in the distance.",
|
||||
"image_path": "images/mixkit-a-cute-couple-playing-on-the-grass-4688.png"
|
||||
"image_path": "assets/images/mixkit-a-cute-couple-playing-on-the-grass-4688.png"
|
||||
},
|
||||
{
|
||||
"prompt": "Aerial view of a rocky mountain in the forest at a sunny day drone flight footage",
|
||||
"image_path": "images/mixkit-aerial-view-of-a-rocky-mountain-in-the-forest-50589.png"
|
||||
"image_path": "assets/images/mixkit-aerial-view-of-a-rocky-mountain-in-the-forest-50589.png"
|
||||
},
|
||||
{
|
||||
"prompt": "A little girl wearing a pink security helmet and denim overall discovers the art of cycling amidst the serene park, as the camera captures her graceful progress.",
|
||||
"image_path": "images/mixkit-a-little-girl-cruises-through-the-forest-path-on-her-50088.png"
|
||||
"image_path": "assets/images/mixkit-a-little-girl-cruises-through-the-forest-path-on-her-50088.png"
|
||||
},
|
||||
{
|
||||
"prompt": "Aerial shot of a beach shore with sea waves. Big rocks on the sand at an alone beach.",
|
||||
"image_path": "images/mixkit-aerial-shot-of-a-beach-with-sea-waves-1087.png"
|
||||
"image_path": "assets/images/mixkit-aerial-shot-of-a-beach-with-sea-waves-1087.png"
|
||||
},
|
||||
{
|
||||
"prompt": "Young woman cleaning her house decorated with plants and decorations, while dancing happily to music in her headphones.",
|
||||
"image_path": "images/mixkit-woman-cleaning-her-house-dancing-happy-43379.png"
|
||||
"image_path": "assets/images/mixkit-woman-cleaning-her-house-dancing-happy-43379.png"
|
||||
},
|
||||
{
|
||||
"prompt": "Aerial tour in a meadow surrounded by hills on the horizon, while some birds fly low over a lake.",
|
||||
"image_path": "images/mixkit-birds-flying-low-over-a-lake-in-a-meadow-41417.png"
|
||||
"image_path": "assets/images/mixkit-birds-flying-low-over-a-lake-in-a-meadow-41417.png"
|
||||
},
|
||||
{
|
||||
"prompt": "In the video, a lone rider guides a majestic horse across an expansive, open field as the sun sets in the background. The rider, dressed in a classic blue shirt and wide-brimmed hat, sits confidently in the saddle, silhouetted against the warm glow of the evening sky. The horse moves gracefully, its mane and tail flowing with each step, creating a sense of harmony between horse and rider. Surrounding the pair, towering trees form a natural border, their leaves gently rustling in the breeze. The shadows lengthen on the ground, accentuating the serene and timeless feel of the scene. The distant hills and wooden fences frame the horizon, adding depth to the tranquil landscape. A few horses graze peacefully in the background, blending into the pastoral setting. The overall ambiance evokes a sense of calmness and quietude, capturing a perfect moment in the golden light of dusk.",
|
||||
"image_path": "images/mixkit-a-rancher-riding-a-horse-at-sunset-1143.png"
|
||||
"image_path": "assets/images/mixkit-a-rancher-riding-a-horse-at-sunset-1143.png"
|
||||
},
|
||||
{
|
||||
"prompt": "In the video, a martial artist dressed in a traditional white uniform with a black belt demonstrates a series of precise movements against a stark black background. The individual gracefully transitions between stances, embodying a sense of focused discipline and control. Each motion is executed with a deliberate pace, showcasing the fluidity of martial arts techniques. The soft lighting creates subtle highlights on the uniform, adding depth to the figure as it moves. The practitioner begins with an open-hand pose, feet firmly grounded, gradually shifting to a powerful forward punch. The fluidity of the sequence displays a mastery of balance and poise. Every trajectory of the limbs is precise and deliberate, capturing the elegance and strength of martial arts. The serene, isolated setting enhances the intensity and concentration of the practitioner. This visual presentation is an elegant interplay of motion and stillness, displaying the art form's discipline and grace.",
|
||||
"image_path": "images/mixkit-a-young-man-practicing-his-karate-moves-49635.png"
|
||||
"image_path": "assets/images/mixkit-a-young-man-practicing-his-karate-moves-49635.png"
|
||||
},
|
||||
{
|
||||
"prompt": "In a serene and softly lit yoga studio, three individuals engage in a yoga session, each performing an upward-facing stretch. The central figure is a woman with shoulder-length brown hair, dressed in a light cropped top and green leggings, her posture reflecting grace and concentration. To her right, another participant, a woman in a purple outfit, mirrors the pose with equal poise. On her left, a person with a bun focuses intently, supported slightly by yoga blocks beneath their hands. The warm-colored wooden floor contrasts soothingly with the soft pastel mural on the back wall, featuring an abstract design and partial visage of a serene face. Natural light floods the space from a large window on the right, where lush greens peek through, adding an element of tranquility. In the corner of the room, a collection of meditation instruments, including a gong and a Buddha statue, subtly frame the peaceful setting. The mood is calm yet focused, as all three participants are deeply engaged in their practice. The scene combines elements of balance, harmony, and a shared journey towards mindfulness. This depiction captures the essence of a yoga session that blends personal growth with collective experience.",
|
||||
"image_path": "images/mixkit-small-group-of-people-doing-yoga-together-43730.png"
|
||||
"image_path": "assets/images/mixkit-small-group-of-people-doing-yoga-together-43730.png"
|
||||
},
|
||||
{
|
||||
"prompt": "In the deep blue expanse of the ocean, two dolphins glide effortlessly, their sleek bodies reflecting the sunlight filtering through the water. The prominent shadows and caustics create a shimmering effect on their skin, capturing the beauty of their natural habitat. Each dolphin moves with a fluid grace, occasionally interacting with gentle nudges, showcasing their playful and social nature. The scene is vibrant and dynamic, with the clear blue background accentuating the dolphins' movements, making it an ideal subject for AI recreation.",
|
||||
"image_path": "images/mixkit-dolphins-underwater-4133.png"
|
||||
"image_path": "assets/images/mixkit-dolphins-underwater-4133.png"
|
||||
},
|
||||
{
|
||||
"prompt": "A bustling ski slope comes alive with skiers descending a pristine, snow-covered hill, surrounded by towering, snow-draped evergreens. Several figures stand atop the slope, silhouetted against a clear blue sky, preparing to embark on their ski run. The chair lift on the right continuously drops off eager adventurers, adding to the excitement at the hilltop. Each skier, clad in colorful winter gear, carves distinct paths into the textured snow as they weave their way down. The interplay of sunlight and shadows accentuates the myriad tracks etched into the slope, creating a dynamic visual rhythm. The scene captures a vibrant winter wonderland, full of action and the thrill of a perfect ski day.",
|
||||
"image_path": "images/mixkit-skiers-on-a-snowy-slope-3327.png"
|
||||
"image_path": "assets/images/mixkit-skiers-on-a-snowy-slope-3327.png"
|
||||
},
|
||||
{
|
||||
"prompt": "A determined climber is scaling a massive rock face, showcasing exceptional strength and skill. The person, clad in a teal shirt and dark pants, climbs with precision, their movements measured and deliberate. They are secured by climbing gear, which includes ropes and a harness, emphasizing their commitment to safety. The rugged texture of the sandy-colored rock provides an imposing backdrop, adding drama and scale to the climb. In the distance, other large rock formations and sparse vegetation can be seen under a bright, overcast sky, contributing to the natural and adventurous atmosphere. The scene captures a moment of focus and challenge, highlighting the climber's tenacity and the breathtaking environment.",
|
||||
"image_path": "images/mixkit-alpinist-climbing-a-huge-rock-in-a-desert-43306.png"
|
||||
"image_path": "assets/images/mixkit-alpinist-climbing-a-huge-rock-in-a-desert-43306.png"
|
||||
},
|
||||
{
|
||||
"prompt": "A silver SUV drives along a winding, snow-covered mountain road, with dense pine trees blanketed in snow lining both sides. The scene is serene, with the vehicle moving smoothly, possibly on a winter journey or vacation. As the SUV disappears around the bend, another, darker SUV follows, creating a sense of motion and perspective on the snow-dusted asphalt. The towering, snow-laden rock formation to the right contrasts with the dark green of the pines, highlighting the peacefulness of the wintry landscape.",
|
||||
"image_path": "images/mixkit-curve-on-a-snowy-forest-road-3317.png"
|
||||
"image_path": "assets/images/mixkit-curve-on-a-snowy-forest-road-3317.png"
|
||||
},
|
||||
{
|
||||
"prompt": "A solitary boat glides across the expansive, tranquil expanse of a serene lake. The vessel leaves a gentle wake behind, creating delicate ripples across the mirror-like surface. The water appears a rich shade of teal, seamlessly blending with the sky at the horizon. Silhouettes of distant trees are faintly visible, creating a picturesque backdrop that enhances the solitary journey of the boat. The sky is a calm gradient, shifting from soft oranges near the shore to the pale blues above. In the distance, a few slender poles emerge from the water, remnants of an old structure or natural formation. The mood of the scene is one of peace and solitude, with the boat journeying steadily through the quiet landscape. There is a sense of endless possibilities as the boat moves toward the unseen beyond the frame. The simplicity and stillness of the scene invite contemplation and reflection, encapsulating a perfect moment of quietude on the water.",
|
||||
"image_path": "images/mixkit-motorboat-on-a-large-lake-with-turquoise-blue-waters-4996.png"
|
||||
"image_path": "assets/images/mixkit-motorboat-on-a-large-lake-with-turquoise-blue-waters-4996.png"
|
||||
},
|
||||
{
|
||||
"prompt": "A man wearing grey shorts jumps rope in a gym, weights and gym equipment in the background.",
|
||||
"image_path": "images/gray_short_man.jpg"
|
||||
"image_path": "assets/images/gray_short_man.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "Flying over a peninsula covered in bushy trees, while discovering the sea around it, painted a beautiful turquoise blue, on a sunny day.",
|
||||
"image_path": "images/peninsula.jpg"
|
||||
"image_path": "assets/images/peninsula.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "Skillful cyclist doing a wheelie on a bike while riding through a forest, on a dirt road, surrounded by many trees, in the morning.",
|
||||
"image_path": "images/cyclist.jpg"
|
||||
"image_path": "assets/images/cyclist.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "Some friends dancing and having fun together in circles, at a party surrounded by colored lights at a party, in a fancy old place, in a view from below them.",
|
||||
"image_path": "images/friends.jpg"
|
||||
"image_path": "assets/images/friends.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "A saxophonist wearing a blazer dances while playing a song in a park.",
|
||||
"image_path": "images/saxophonist.jpg"
|
||||
"image_path": "assets/images/saxophonist.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "Romantic couple embracing and looking at each other in the middle of a forest, during a break on a road trip through nature.",
|
||||
"image_path": "images/romance.jpg"
|
||||
"image_path": "assets/images/romance.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "Man dressed in 80's style dances very happily in his kitchen while listening to music on his radio and drinking wine.",
|
||||
"image_path": "images/80s_dance.jpg"
|
||||
"image_path": "assets/images/80s_dance.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "Pair of jazz musicians performing a song with their saxophone and trombone on an abandoned train.",
|
||||
"image_path": "images/jazz.jpg"
|
||||
"image_path": "assets/images/jazz.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "A young woman with short hair wearing pink sunglasses chews gum and makes a bubble gum with the city in the background.",
|
||||
"image_path": "images/pink.jpg"
|
||||
"image_path": "assets/images/pink.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "Natural aerial landscape with a relief covered with abundant trees and vegetation and a thick layer of mist.",
|
||||
"image_path": "images/natural.jpg"
|
||||
"image_path": "assets/images/natural.jpg"
|
||||
},
|
||||
{
|
||||
"prompt": "Loving couple sitting on a log on the shore of a lake outside, sharing an affectionate hug.",
|
||||
"image_path": "images/couple.jpg"
|
||||
"image_path": "assets/images/couple.jpg"
|
||||
}
|
||||
]
|
||||
@@ -1,4 +1,4 @@
|
||||
# FastVideo/videos
|
||||
# FastVideo/assets/videos
|
||||
|
||||
This folder is used to store **video assets for examples**, primarily **input videos** consumed by scripts under `FastVideo/examples/`.
|
||||
|
||||
@@ -1,195 +0,0 @@
|
||||
import argparse
|
||||
import os
|
||||
import tempfile
|
||||
|
||||
import gradio as gr
|
||||
import torch
|
||||
from diffusers import FlowMatchEulerDiscreteScheduler
|
||||
from diffusers.utils import export_to_video
|
||||
|
||||
from fastvideo.distill.solver import PCMFMScheduler
|
||||
from fastvideo.models.mochi_hf.modeling_mochi import MochiTransformer3DModel
|
||||
from fastvideo.models.mochi_hf.pipeline_mochi import MochiPipeline
|
||||
|
||||
|
||||
def init_args():
|
||||
parser = argparse.ArgumentParser()
|
||||
parser.add_argument("--prompts", nargs="+", default=[])
|
||||
parser.add_argument("--num_frames", type=int, default=25)
|
||||
parser.add_argument("--height", type=int, default=480)
|
||||
parser.add_argument("--width", type=int, default=848)
|
||||
parser.add_argument("--num_inference_steps", type=int, default=8)
|
||||
parser.add_argument("--guidance_scale", type=float, default=4.5)
|
||||
parser.add_argument("--model_path", type=str, default="data/mochi")
|
||||
parser.add_argument("--seed", type=int, default=12345)
|
||||
parser.add_argument("--transformer_path", type=str, default=None)
|
||||
parser.add_argument("--scheduler_type", type=str, default="pcm_linear_quadratic")
|
||||
parser.add_argument("--lora_checkpoint_dir", type=str, default=None)
|
||||
parser.add_argument("--shift", type=float, default=8.0)
|
||||
parser.add_argument("--num_euler_timesteps", type=int, default=50)
|
||||
parser.add_argument("--linear_threshold", type=float, default=0.1)
|
||||
parser.add_argument("--linear_range", type=float, default=0.75)
|
||||
parser.add_argument("--cpu_offload", action="store_true")
|
||||
return parser.parse_args()
|
||||
|
||||
|
||||
def load_model(args):
|
||||
if args.scheduler_type == "euler":
|
||||
scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
else:
|
||||
linear_quadratic = True if "linear_quadratic" in args.scheduler_type else False
|
||||
scheduler = PCMFMScheduler(
|
||||
1000,
|
||||
args.shift,
|
||||
args.num_euler_timesteps,
|
||||
linear_quadratic,
|
||||
args.linear_threshold,
|
||||
args.linear_range,
|
||||
)
|
||||
|
||||
if args.transformer_path:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.transformer_path)
|
||||
else:
|
||||
transformer = MochiTransformer3DModel.from_pretrained(args.model_path, subfolder="transformer/")
|
||||
|
||||
pipe = MochiPipeline.from_pretrained(args.model_path, transformer=transformer, scheduler=scheduler)
|
||||
pipe.enable_vae_tiling()
|
||||
# pipe.to(device)
|
||||
# if args.cpu_offload:
|
||||
pipe.enable_sequential_cpu_offload()
|
||||
return pipe
|
||||
|
||||
|
||||
def generate_video(
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed=False,
|
||||
):
|
||||
if randomize_seed:
|
||||
seed = torch.randint(0, 1000000, (1, )).item()
|
||||
|
||||
generator = torch.Generator(device="cuda").manual_seed(seed)
|
||||
|
||||
if not use_negative_prompt:
|
||||
negative_prompt = None
|
||||
|
||||
with torch.autocast("cuda", dtype=torch.bfloat16):
|
||||
output = pipe(
|
||||
prompt=[prompt],
|
||||
negative_prompt=negative_prompt,
|
||||
height=height,
|
||||
width=width,
|
||||
num_frames=num_frames,
|
||||
num_inference_steps=num_inference_steps,
|
||||
guidance_scale=guidance_scale,
|
||||
generator=generator,
|
||||
).frames[0]
|
||||
|
||||
output_path = os.path.join(tempfile.mkdtemp(), "output.mp4")
|
||||
export_to_video(output, output_path, fps=30)
|
||||
return output_path, seed
|
||||
|
||||
|
||||
examples = [
|
||||
"A hand enters the frame, pulling a sheet of plastic wrap over three balls of dough placed on a wooden surface. The plastic wrap is stretched to cover the dough more securely. The hand adjusts the wrap, ensuring that it is tight and smooth over the dough. The scene focuses on the hand’s movements as it secures the edges of the plastic wrap. No new objects appear, and the camera remains stationary, focusing on the action of covering the dough.",
|
||||
"A vintage train snakes through the mountains, its plume of white steam rising dramatically against the jagged peaks. The cars glint in the late afternoon sun, their deep crimson and gold accents lending a touch of elegance. The tracks carve a precarious path along the cliffside, revealing glimpses of a roaring river far below. Inside, passengers peer out the large windows, their faces lit with awe as the landscape unfolds.",
|
||||
"A crowded rooftop bar buzzes with energy, the city skyline twinkling like a field of stars in the background. Strings of fairy lights hang above, casting a warm, golden glow over the scene. Groups of people gather around high tables, their laughter blending with the soft rhythm of live jazz. The aroma of freshly mixed cocktails and charred appetizers wafts through the air, mingling with the cool night breeze.",
|
||||
]
|
||||
|
||||
args = init_args()
|
||||
pipe = load_model(args)
|
||||
print("load model successfully")
|
||||
with gr.Blocks() as demo:
|
||||
gr.Markdown("# Fastvideo Mochi Video Generation Demo")
|
||||
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
prompt = gr.Text(
|
||||
label="Prompt",
|
||||
show_label=False,
|
||||
max_lines=1,
|
||||
placeholder="Enter your prompt",
|
||||
container=False,
|
||||
)
|
||||
run_button = gr.Button("Run", scale=0)
|
||||
result = gr.Video(label="Result", show_label=False)
|
||||
|
||||
with gr.Accordion("Advanced options", open=False):
|
||||
with gr.Group():
|
||||
with gr.Row():
|
||||
height = gr.Slider(
|
||||
label="Height",
|
||||
minimum=256,
|
||||
maximum=1024,
|
||||
step=32,
|
||||
value=args.height,
|
||||
)
|
||||
width = gr.Slider(label="Width", minimum=256, maximum=1024, step=32, value=args.width)
|
||||
|
||||
with gr.Row():
|
||||
num_frames = gr.Slider(
|
||||
label="Number of Frames",
|
||||
minimum=21,
|
||||
maximum=163,
|
||||
value=args.num_frames,
|
||||
)
|
||||
guidance_scale = gr.Slider(
|
||||
label="Guidance Scale",
|
||||
minimum=1,
|
||||
maximum=12,
|
||||
value=args.guidance_scale,
|
||||
)
|
||||
num_inference_steps = gr.Slider(
|
||||
label="Inference Steps",
|
||||
minimum=4,
|
||||
maximum=100,
|
||||
value=args.num_inference_steps,
|
||||
)
|
||||
|
||||
with gr.Row():
|
||||
use_negative_prompt = gr.Checkbox(label="Use negative prompt", value=False)
|
||||
negative_prompt = gr.Text(
|
||||
label="Negative prompt",
|
||||
max_lines=1,
|
||||
placeholder="Enter a negative prompt",
|
||||
visible=False,
|
||||
)
|
||||
|
||||
seed = gr.Slider(label="Seed", minimum=0, maximum=1000000, step=1, value=args.seed)
|
||||
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
|
||||
seed_output = gr.Number(label="Used Seed")
|
||||
|
||||
gr.Examples(examples=examples, inputs=prompt)
|
||||
|
||||
use_negative_prompt.change(
|
||||
fn=lambda x: gr.update(visible=x),
|
||||
inputs=use_negative_prompt,
|
||||
outputs=negative_prompt,
|
||||
)
|
||||
|
||||
run_button.click(
|
||||
fn=generate_video,
|
||||
inputs=[
|
||||
prompt,
|
||||
negative_prompt,
|
||||
use_negative_prompt,
|
||||
seed,
|
||||
guidance_scale,
|
||||
num_frames,
|
||||
height,
|
||||
width,
|
||||
num_inference_steps,
|
||||
randomize_seed,
|
||||
],
|
||||
outputs=[result, seed_output],
|
||||
)
|
||||
|
||||
if __name__ == "__main__":
|
||||
demo.queue(max_size=20).launch(server_name="0.0.0.0", server_port=7860)
|
||||
@@ -1,15 +0,0 @@
|
||||
Fast-Hunyuan comparison with original Hunyuan, achieving an 8X diffusion speed boost with the FastVideo framework.
|
||||
|
||||
https://github.com/user-attachments/assets/064ac1d2-11ed-4a0c-955b-4d412a96ef30
|
||||
|
||||
Fast-Mochi comparison with original Mochi, achieving an 8X diffusion speed boost with the FastVideo framework.
|
||||
|
||||
https://github.com/user-attachments/assets/5fbc4596-56d6-43aa-98e0-da472cf8e26c
|
||||
|
||||
Comparison between OpenAI Sora, original Hunyuan and FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/d323b712-3f68-42b2-952b-94f6a49c4836
|
||||
|
||||
Comparison between original FastHunyuan, LLM-INT8 quantized FastHunyuan and NF4 quantized FastHunyuan
|
||||
|
||||
https://github.com/user-attachments/assets/cf89efb5-5f68-4949-a085-f41c1ef26c94
|
||||
@@ -20,7 +20,7 @@ def main():
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
|
||||
# image2world example from official repo
|
||||
image_path = "images/bus_terminal.jpg"
|
||||
image_path = "assets/images/bus_terminal.jpg"
|
||||
|
||||
prompt = (
|
||||
"A nighttime city bus terminal gradually shifts from stillness to subtle movement. "
|
||||
@@ -48,4 +48,3 @@ def main():
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -20,7 +20,7 @@ def main():
|
||||
sampling_param = SamplingParam.from_pretrained(model_path)
|
||||
|
||||
# video2world example from official repo
|
||||
video_path = "videos/robot_pouring.mp4"
|
||||
video_path = "assets/videos/robot_pouring.mp4"
|
||||
|
||||
prompt = (
|
||||
"A robotic arm, primarily white with black joints and cables, is shown in a clean, modern indoor setting with a white tabletop. "
|
||||
@@ -51,4 +51,3 @@ def main():
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,49 @@
|
||||
from fastvideo import VideoGenerator
|
||||
from fastvideo.models.dits.lingbotworld.cam_utils import prepare_camera_embedding
|
||||
|
||||
# from fastvideo.configs.sample import SamplingParam
|
||||
OUTPUT_PATH = "video_samples_lingbotworld"
|
||||
def main():
|
||||
# FastVideo will automatically use the optimal default arguments for the
|
||||
# model.
|
||||
# If a local path is provided, FastVideo will make a best effort
|
||||
# attempt to identify the optimal arguments.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
# FastVideo will automatically handle distributed setup
|
||||
num_gpus=1,
|
||||
use_fsdp_inference=False, # set to True if GPU is out of memory
|
||||
dit_cpu_offload=True, # DiT need to be offloaded for MoE
|
||||
vae_cpu_offload=False,
|
||||
text_encoder_cpu_offload=True,
|
||||
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
|
||||
pin_cpu_memory=True,
|
||||
# image_encoder_cpu_offload=False,
|
||||
)
|
||||
|
||||
num_frames = 81
|
||||
prompt = "The video presents a soaring journey through a fantasy jungle. The wind whips past the rider's blue hands gripping the reins, causing the leather straps to vibrate. The ancient gothic castle approaches steadily, its stone details becoming clearer against the backdrop of floating islands and distant waterfalls."
|
||||
image_path = "https://raw.githubusercontent.com/Robbyant/lingbot-world/main/examples/00/image.jpg"
|
||||
action_path = "examples/inference/basic/lingbotworld_examples/00"
|
||||
c2ws_plucker_emb, num_frames = prepare_camera_embedding(
|
||||
action_path=action_path,
|
||||
num_frames=num_frames,
|
||||
height=480,
|
||||
width=832,
|
||||
spatial_scale=8,
|
||||
)
|
||||
|
||||
generator.generate_video(
|
||||
prompt,
|
||||
image_path=image_path,
|
||||
output_path=OUTPUT_PATH,
|
||||
save_video=True,
|
||||
num_frames=num_frames,
|
||||
height=480,
|
||||
width=832,
|
||||
c2ws_plucker_emb=c2ws_plucker_emb,
|
||||
)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -1,5 +1,6 @@
|
||||
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 "
|
||||
@@ -16,16 +17,20 @@ PROMPT = (
|
||||
|
||||
|
||||
def main() -> None:
|
||||
# Uses FastVideo default sampling settings for LTX2 base.
|
||||
generator = VideoGenerator.from_pretrained(
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
"Davids048/LTX2-Base-Diffusers",
|
||||
num_gpus=1,
|
||||
)
|
||||
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_distilled_t2v.mp4"
|
||||
output_path = "outputs_video/ltx2_basic/output_ltx2_base_t2v_1088_1920_1.1.mp4"
|
||||
generator.generate_video(
|
||||
prompt=PROMPT,
|
||||
output_path=output_path,
|
||||
save_video=True,
|
||||
num_frames=121,
|
||||
height=1088,
|
||||
width=1920,
|
||||
)
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
@@ -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()
|
||||
@@ -0,0 +1,137 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import os
|
||||
import re
|
||||
from typing import List
|
||||
|
||||
|
||||
DEFAULT_PROMPTS = [
|
||||
"a photo of a cat",
|
||||
"a cinematic photo of a red panda wearing a tiny backpack, standing on a rainy neon-lit street at night, shallow depth of field, sharp focus, 35mm, bokeh",
|
||||
]
|
||||
|
||||
|
||||
def _safe_filename(text: str, max_len: int = 100) -> str:
|
||||
"""
|
||||
Make a stable, filesystem-friendly filename base.
|
||||
VideoGenerator uses prompt[:100].strip() internally, so we mirror that,
|
||||
but also remove path separators and other problematic characters.
|
||||
"""
|
||||
s = text[:max_len].strip()
|
||||
s = s.replace(os.sep, "_")
|
||||
if os.altsep:
|
||||
s = s.replace(os.altsep, "_")
|
||||
s = re.sub(r"\s+", " ", s)
|
||||
s = re.sub(r"[^A-Za-z0-9 .,_-]", "_", s)
|
||||
s = s.strip(" .")
|
||||
return s or "prompt"
|
||||
|
||||
|
||||
def _remove_existing_outputs(out_dir: str, filename_base: str) -> None:
|
||||
"""
|
||||
Ensure deterministic naming by deleting any existing mp4s that would
|
||||
cause VideoGenerator to append suffixes like _1, _2, etc.
|
||||
"""
|
||||
if not os.path.isdir(out_dir):
|
||||
return
|
||||
|
||||
pattern = re.compile(rf"^{re.escape(filename_base)}(_\d+)?\.mp4$")
|
||||
for fn in os.listdir(out_dir):
|
||||
if pattern.match(fn):
|
||||
try:
|
||||
os.remove(os.path.join(out_dir, fn))
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
|
||||
|
||||
def parse_args() -> argparse.Namespace:
|
||||
p = argparse.ArgumentParser(description="Run SD3.5 Medium text-to-image with FastVideo VideoGenerator.")
|
||||
p.add_argument("--model-path", default="stabilityai/stable-diffusion-3.5-medium", help="Path to local diffusers-format SD3.5 weights directory.")
|
||||
p.add_argument(
|
||||
"--out-dir",
|
||||
"--outdir",
|
||||
default="outputs/sd35/samples",
|
||||
help="Output directory for generated mp4 files.",
|
||||
)
|
||||
p.add_argument(
|
||||
"--prompt",
|
||||
action="append",
|
||||
default=None,
|
||||
help="Prompt text. Repeat --prompt multiple times to generate multiple samples.",
|
||||
)
|
||||
p.add_argument("--negative", default="lowres, blurry, jpeg artifacts, watermark, text", help="Negative prompt.")
|
||||
p.add_argument(
|
||||
"--backend",
|
||||
default=None,
|
||||
help="Set FASTVIDEO_ATTENTION_BACKEND (e.g. TORCH_SDPA). If omitted, respects the existing env var.",
|
||||
)
|
||||
p.add_argument("--seed", type=int, default=42, help="Base seed. Each prompt uses seed + prompt_idx.")
|
||||
p.add_argument("--height", type=int, default=768, help="Output height.")
|
||||
p.add_argument("--width", type=int, default=768, help="Output width.")
|
||||
p.add_argument("--steps", type=int, default=28, help="Number of inference steps.")
|
||||
p.add_argument("--guidance", type=float, default=6.0, help="Guidance scale.")
|
||||
p.add_argument("--num-gpus", type=int, default=1, help="Number of GPUs to use.")
|
||||
return p.parse_args()
|
||||
|
||||
|
||||
def main() -> None:
|
||||
args = parse_args()
|
||||
prompts: List[str] = args.prompt if args.prompt else DEFAULT_PROMPTS
|
||||
|
||||
if args.backend:
|
||||
os.environ["FASTVIDEO_ATTENTION_BACKEND"] = args.backend
|
||||
|
||||
from fastvideo import VideoGenerator
|
||||
|
||||
os.makedirs(args.out_dir, exist_ok=True)
|
||||
|
||||
init_kwargs = {
|
||||
"num_gpus": args.num_gpus,
|
||||
"workload_type": "t2i",
|
||||
"sp_size": 1,
|
||||
"tp_size": 1,
|
||||
"dit_cpu_offload": False,
|
||||
"dit_layerwise_offload": False,
|
||||
"text_encoder_cpu_offload": False,
|
||||
"vae_cpu_offload": False,
|
||||
"image_encoder_cpu_offload": False,
|
||||
"pin_cpu_memory": False,
|
||||
"use_fsdp_inference": False,
|
||||
}
|
||||
|
||||
generator = VideoGenerator.from_pretrained(model_path=args.model_path, **init_kwargs)
|
||||
try:
|
||||
for i, prompt in enumerate(prompts):
|
||||
seed = args.seed + i
|
||||
|
||||
filename_base = f"sd35_{i:02d}_seed{seed}_{_safe_filename(prompt, max_len=80)}"
|
||||
_remove_existing_outputs(args.out_dir, filename_base)
|
||||
|
||||
output_path = os.path.join(args.out_dir, f"{filename_base}.png")
|
||||
print(f"[sd35] prompt_idx={i} seed={seed} output_path={output_path}")
|
||||
|
||||
generation_kwargs = {
|
||||
"output_path": output_path,
|
||||
"height": args.height,
|
||||
"width": args.width,
|
||||
"num_frames": 1,
|
||||
"fps": 1,
|
||||
"num_inference_steps": args.steps,
|
||||
"guidance_scale": args.guidance,
|
||||
"seed": seed,
|
||||
"negative_prompt": args.negative,
|
||||
"save_video": True,
|
||||
}
|
||||
|
||||
generator.generate_video(prompt, **generation_kwargs)
|
||||
|
||||
print(f"[sd35] done. outputs written to: {args.out_dir}")
|
||||
finally:
|
||||
generator.shutdown()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
@@ -31,7 +31,7 @@ def main():
|
||||
sampling_param.height = 480
|
||||
sampling_param.seed = 1000
|
||||
|
||||
with open("prompts/mixkit_i2v.jsonl", "r") as f:
|
||||
with open("assets/prompts/mixkit_i2v.jsonl", "r") as f:
|
||||
prompt_image_pairs = json.load(f)
|
||||
|
||||
for prompt_image_pair in prompt_image_pairs:
|
||||
|
||||
@@ -188,8 +188,9 @@ def load_example_prompts():
|
||||
prompt_to_image = {}
|
||||
# Try to find the JSON file relative to project root
|
||||
possible_json_paths = [
|
||||
Path("prompts/mixkit_i2v.jsonl"),
|
||||
Path(__file__).parent.parent.parent.parent / "prompts" / "mixkit_i2v.jsonl",
|
||||
Path("assets/prompts/mixkit_i2v.jsonl"),
|
||||
Path(__file__).resolve().parents[4] / "assets" / "prompts" /
|
||||
"mixkit_i2v.jsonl",
|
||||
]
|
||||
json_path = None
|
||||
for path in possible_json_paths:
|
||||
@@ -201,8 +202,8 @@ def load_example_prompts():
|
||||
try:
|
||||
with open(json_path, "r", encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
# Get the project root directory (parent of prompts directory)
|
||||
project_root = json_path.parent.parent
|
||||
# Resolve paths relative to repository root.
|
||||
project_root = Path(__file__).resolve().parents[4]
|
||||
for item in data:
|
||||
prompt_text = item.get("prompt", "").strip()
|
||||
image_path = item.get("image_path", "")
|
||||
@@ -736,8 +737,8 @@ def main():
|
||||
allowed_paths=[
|
||||
os.path.abspath("outputs"),
|
||||
os.path.abspath("fastvideo-logos"),
|
||||
os.path.abspath("prompts"),
|
||||
os.path.abspath("images"),
|
||||
os.path.abspath("assets/prompts"),
|
||||
os.path.abspath("assets/images"),
|
||||
os.path.abspath(tempfile.gettempdir()),
|
||||
os.path.abspath(os.path.join(tempfile.gettempdir(), "gradio")),
|
||||
]
|
||||
@@ -747,4 +748,4 @@ def main():
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
main()
|
||||
main()
|
||||
|
||||
@@ -0,0 +1,23 @@
|
||||
# LTX-2 Crush-Smol Example
|
||||
# TODO: Update this doc.
|
||||
|
||||
These are e2e example scripts for finetuning LTX-2 on the crush-smol dataset.
|
||||
|
||||
## Execute the following commands from `FastVideo/` to run training:
|
||||
|
||||
### Download crush-smol dataset:
|
||||
|
||||
`bash examples/training/finetune/ltx2/overfit/download_dataset.sh`
|
||||
|
||||
### Preprocess the videos and captions into latents:
|
||||
|
||||
`bash examples/training/finetune/ltx2/overfit/preprocess_ltx2_data_t2v_new.sh`
|
||||
|
||||
### Edit the following file and run finetuning:
|
||||
|
||||
`bash examples/training/finetune/ltx2/overfit/finetune_t2v.sh`
|
||||
|
||||
Notes:
|
||||
- Update `DATASET_PATH` in the preprocess script to point to your merged dataset root (`videos/` + `videos2caption.json`).
|
||||
- `MODEL_PATH` should point to a local LTX-2 diffusers-style directory that contains `model_index.json` and `text_encoder/gemma`.
|
||||
|
||||
@@ -0,0 +1,3 @@
|
||||
# #!/bin/bash
|
||||
#
|
||||
python scripts/huggingface/download_hf.py --repo_id "wlsaidhi/crush-smol-merged" --local_dir "data/crush-smol" --repo_type "dataset"
|
||||
@@ -0,0 +1,95 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
|
||||
# Also can use simple 1 video for overfitting experiments.
|
||||
# DATA_DIR="/home/hal-jundas/codes/FastVideo/data/crush-smol"
|
||||
DATA_DIR="<PATH_TO_PROCESSED_DATASET>"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
echo VALIDATION_DATASET_FILE: $VALIDATION_DATASET_FILE
|
||||
NUM_GPUS=4
|
||||
OVERFIT_HEIGHT=480
|
||||
OVERFIT_WIDTH=832
|
||||
OVERFIT_FRAMES=73
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name "ltx2_t2v_finetune"
|
||||
--output_dir "checkpoints/ltx2_t2v_finetune"
|
||||
--max_train_steps 5000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 10
|
||||
--num_height $OVERFIT_HEIGHT
|
||||
--num_width $OVERFIT_WIDTH
|
||||
--num_frames $OVERFIT_FRAMES
|
||||
--ltx2-first-frame-conditioning-p 0.1
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--mode "finetuning"
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "3.0"
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
--lr_scheduler "linear"
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--dit_precision "fp32"
|
||||
--dit_cpu_offload False
|
||||
--dit_layerwise_offload False
|
||||
--text_encoder_cpu_offload False
|
||||
--image_encoder_cpu_offload False
|
||||
--vae_cpu_offload False
|
||||
)
|
||||
|
||||
# NOTE: Setting this environment variable to TORCH_SDPA to avoid the issue of stacking that failed in flash attn.
|
||||
export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ltx2_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,80 @@
|
||||
#!/bin/bash
|
||||
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="/path/to/LTX-2"
|
||||
DATA_DIR="data/crush-smol"
|
||||
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
|
||||
NUM_GPUS=1
|
||||
|
||||
training_args=(
|
||||
--tracker_project_name "ltx2_t2v_lora_finetune"
|
||||
--output_dir "checkpoints/ltx2_t2v_lora_finetune"
|
||||
--max_train_steps 2000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 8
|
||||
--num_latent_t 10
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 77
|
||||
--ltx2-first-frame-conditioning-p 0.1
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
)
|
||||
|
||||
parallel_args=(
|
||||
--num_gpus $NUM_GPUS
|
||||
--sp_size $NUM_GPUS
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
model_args=(
|
||||
--model_path $MODEL_PATH
|
||||
--pretrained_model_name_or_path $MODEL_PATH
|
||||
)
|
||||
|
||||
dataset_args=(
|
||||
--data_path $DATA_DIR
|
||||
--dataloader_num_workers 1
|
||||
)
|
||||
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file $VALIDATION_DATASET_FILE
|
||||
--validation_steps 200
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "3.0"
|
||||
)
|
||||
|
||||
optimizer_args=(
|
||||
--learning_rate 2e-4
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--max_grad_norm 1.0
|
||||
--lora_training True
|
||||
--lora_rank 16
|
||||
--lora_alpha 16
|
||||
)
|
||||
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
fastvideo/training/ltx2_training_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}"
|
||||
@@ -0,0 +1,35 @@
|
||||
#!/bin/bash
|
||||
|
||||
GPU_NUM=1
|
||||
MODEL_PATH="Davids048/LTX2-Base-Diffusers"
|
||||
# DATASET_PATH="data/overfit"
|
||||
DATASET_PATH="data/crush-smol"
|
||||
OUTPUT_DIR="$DATASET_PATH"
|
||||
WITH_AUDIO=true
|
||||
|
||||
# Convert one-file overfit metadata into merged format if needed.
|
||||
if [ ! -f "$DATASET_PATH/videos2caption.json" ] && [ -f "$DATASET_PATH/overfit.json" ]; then
|
||||
python scripts/dataset_preparation/convert_to_merged_dataset.py \
|
||||
--items-json "$DATASET_PATH/overfit.json" \
|
||||
--output-dir "$DATASET_PATH"
|
||||
fi
|
||||
|
||||
|
||||
torchrun --nproc_per_node=$GPU_NUM \
|
||||
--master_port=29513 \
|
||||
-m fastvideo.pipelines.preprocess.v1_preprocessing_new \
|
||||
--model_path $MODEL_PATH \
|
||||
--mode preprocess \
|
||||
--workload_type t2v \
|
||||
--preprocess.video_loader_type torchvision \
|
||||
--preprocess.dataset_type merged \
|
||||
--preprocess.dataset_path $DATASET_PATH \
|
||||
--preprocess.dataset_output_dir $OUTPUT_DIR \
|
||||
--preprocess.with_audio $WITH_AUDIO \
|
||||
--preprocess.preprocess_video_batch_size 1 \
|
||||
--preprocess.dataloader_num_workers 0 \
|
||||
--preprocess.max_height 480 \
|
||||
--preprocess.max_width 832 \
|
||||
--preprocess.num_frames 73 \
|
||||
--preprocess.train_fps 16 \
|
||||
--preprocess.video_length_tolerance_range 5
|
||||
@@ -0,0 +1,13 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "The camera opens in a calm, sunlit frog yoga studio. Warm morning light washes over the wooden floor as incense smoke drifts lazily in the air. The senior frog instructor sits cross-legged at the center, eyes closed, voice deep and calm. “We are one with the pond.” All the frogs answer softly: “Ommm...” “We are one with the mud.” “Ommm...” He smiles faintly. “We are one with the flies.” A quiet pause. The camera slowly pans to the side — one frog twitches, eyes darting. Suddenly — *thwip!* — its tongue snaps out, catching a fly mid-air and pulling it into its mouth. The master exhales slowly, still serene.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 1088,
|
||||
"width": 1920,
|
||||
"num_frames": 121
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
{
|
||||
"data": [
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of Oreo cookies, flattening them as if they were under a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen compressing colorful clay into a compact shape, demonstrating the power of a hydraulic press.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
},
|
||||
{
|
||||
"caption": "A large metal cylinder is seen pressing down on a pile of colorful candies, flattening them as if they were under a hydraulic press. The candies are crushed and broken into small pieces, creating a mess on the table.",
|
||||
"image_path": null,
|
||||
"video_path": null,
|
||||
"num_inference_steps": 50,
|
||||
"height": 480,
|
||||
"width": 832,
|
||||
"num_frames": 77
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -86,6 +86,7 @@ class PreprocessConfig:
|
||||
|
||||
# Model configuration
|
||||
training_cfg_rate: float = 0.0
|
||||
with_audio: bool = False
|
||||
|
||||
# framework configuration
|
||||
seed: int = 42
|
||||
@@ -190,6 +191,10 @@ class PreprocessConfig:
|
||||
type=float,
|
||||
default=PreprocessConfig.training_cfg_rate,
|
||||
help="Training CFG rate")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}with-audio",
|
||||
action=StoreBoolean,
|
||||
default=PreprocessConfig.with_audio,
|
||||
help="Whether to extract and encode audio")
|
||||
preprocess_args.add_argument(f"--{prefix_with_dot}seed",
|
||||
type=int,
|
||||
default=PreprocessConfig.seed,
|
||||
|
||||
@@ -2,6 +2,7 @@ from fastvideo.configs.models.dits.cosmos import CosmosVideoConfig
|
||||
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.lingbotworld import LingBotWorldVideoConfig
|
||||
from fastvideo.configs.models.dits.longcat import LongCatVideoConfig
|
||||
from fastvideo.configs.models.dits.ltx2 import LTX2VideoConfig
|
||||
from fastvideo.configs.models.dits.stepvideo import StepVideoConfig
|
||||
@@ -11,5 +12,6 @@ from fastvideo.configs.models.dits.hyworld import HYWorldConfig
|
||||
__all__ = [
|
||||
"HunyuanVideoConfig", "HunyuanVideo15Config", "WanVideoConfig",
|
||||
"StepVideoConfig", "CosmosVideoConfig", "Cosmos25VideoConfig",
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig"
|
||||
"LongCatVideoConfig", "LTX2VideoConfig", "HYWorldConfig",
|
||||
"LingBotWorldVideoConfig"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,110 @@
|
||||
# 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 LingBotWorldArchConfig(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"^patch_embedding_wancamctrl\.(.*)$":
|
||||
r"patch_embedding_wancamctrl.proj.\1",
|
||||
r"^c2ws_hidden_states_layer1\.(.*)$": r"c2ws_mlp.fc_in.\1",
|
||||
r"^c2ws_hidden_states_layer2\.(.*)$": r"c2ws_mlp.fc_out.\1",
|
||||
r"^text_embedding\.0\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_in.\1",
|
||||
r"^text_embedding\.2\.(.*)$":
|
||||
r"condition_embedder.text_embedder.fc_out.\1",
|
||||
r"^time_embedding\.0\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_in.\1",
|
||||
r"^time_embedding\.2\.(.*)$":
|
||||
r"condition_embedder.time_embedder.mlp.fc_out.\1",
|
||||
r"^time_projection\.1\.(.*)$":
|
||||
r"condition_embedder.time_modulation.linear.\1",
|
||||
r"^blocks\.(\d+)\.modulation$": r"blocks.\1.scale_shift_table",
|
||||
r"^blocks\.(\d+)\.self_attn\.q\.(.*)$": r"blocks.\1.to_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.k\.(.*)$": r"blocks.\1.to_k.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.v\.(.*)$": r"blocks.\1.to_v.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.o\.(.*)$": r"blocks.\1.to_out.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_q\.(.*)$": r"blocks.\1.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.self_attn\.norm_k\.(.*)$": r"blocks.\1.norm_k.\2",
|
||||
r"^blocks\.(\d+)\.norm3\.(.*)$":
|
||||
r"blocks.\1.self_attn_residual_norm.norm.\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.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_q\.(.*)$":
|
||||
r"blocks.\1.attn2.norm_q.\2",
|
||||
r"^blocks\.(\d+)\.cross_attn\.norm_k\.(.*)$":
|
||||
r"blocks.\1.attn2.norm_k.\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",
|
||||
r"^blocks\.(\d+)\.cam_injector_layer1\.(.*)$":
|
||||
r"blocks.\1.cam_conditioner.cam_injector.fc_in.\2",
|
||||
r"^blocks\.(\d+)\.cam_injector_layer2\.(.*)$":
|
||||
r"blocks.\1.cam_conditioner.cam_injector.fc_out.\2",
|
||||
r"^blocks\.(\d+)\.cam_scale_layer\.(.*)$":
|
||||
r"blocks.\1.cam_conditioner.cam_scale_layer.\2",
|
||||
r"^blocks\.(\d+)\.cam_shift_layer\.(.*)$":
|
||||
r"blocks.\1.cam_conditioner.cam_shift_layer.\2",
|
||||
r"^head\.modulation$": r"scale_shift_table",
|
||||
r"^head\.head\.(.*)$": r"proj_out.\1",
|
||||
})
|
||||
|
||||
# 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: {})
|
||||
|
||||
patch_size: tuple[int, int, int] = (1, 2, 2)
|
||||
text_len: int = 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 LingBotWorldVideoConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(default_factory=LingBotWorldArchConfig)
|
||||
|
||||
prefix: str = "Wan"
|
||||
@@ -7,10 +7,12 @@ from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
import re
|
||||
|
||||
|
||||
def is_ltx2_blocks(name: str, _module) -> bool:
|
||||
"""FSDP shard condition for LTX-2 transformer blocks."""
|
||||
return "transformer_blocks" in name
|
||||
res = re.search(r"(?:^|\.)transformer_blocks\.\d+$", name) is not None
|
||||
return res
|
||||
|
||||
|
||||
@dataclass
|
||||
|
||||
@@ -0,0 +1,31 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models.dits.base import DiTArchConfig, DiTConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class SD3Transformer2DArchConfig(DiTArchConfig):
|
||||
# Diffusers SD3Transformer2DModel config fields.
|
||||
sample_size: int = 128
|
||||
patch_size: int = 2
|
||||
num_layers: int = 24
|
||||
attention_head_dim: int = 64
|
||||
joint_attention_dim: int = 4096
|
||||
caption_projection_dim: int = 1536
|
||||
pooled_projection_dim: int = 2048
|
||||
pos_embed_max_size: int = 384
|
||||
dual_attention_layers: list[int] = field(
|
||||
default_factory=lambda: [0, 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12])
|
||||
qk_norm: str = "rms_norm"
|
||||
in_channels: int = 16
|
||||
out_channels: int = 16
|
||||
num_attention_heads = 24
|
||||
|
||||
|
||||
@dataclass
|
||||
class SD3DiTConfig(DiTConfig):
|
||||
arch_config: DiTArchConfig = field(
|
||||
default_factory=SD3Transformer2DArchConfig)
|
||||
prefix: str = "sd3"
|
||||
@@ -0,0 +1,39 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models.vaes.base import VAEArchConfig, VAEConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class AutoencoderKLArchConfig(VAEArchConfig):
|
||||
_name_or_path: str = ""
|
||||
act_fn: str = "silu"
|
||||
block_out_channels: tuple[int, ...] | list[int] = field(
|
||||
default_factory=list)
|
||||
down_block_types: tuple[str, ...] | list[str] = field(default_factory=list)
|
||||
up_block_types: tuple[str, ...] | list[str] = field(default_factory=list)
|
||||
force_upcast: bool = True
|
||||
in_channels: int = 3
|
||||
latent_channels: int = 4
|
||||
latents_mean: tuple[float, ...] | list[float] | None = None
|
||||
latents_std: tuple[float, ...] | list[float] | None = None
|
||||
layers_per_block: int = 1
|
||||
mid_block_add_attention: bool = True
|
||||
norm_num_groups: int = 32
|
||||
out_channels: int = 3
|
||||
sample_size: int = 32
|
||||
scaling_factor: float | torch.Tensor = 0.18215
|
||||
shift_factor: float | None = None
|
||||
use_post_quant_conv: bool = True
|
||||
use_quant_conv: bool = True
|
||||
|
||||
temporal_compression_ratio: int = 1
|
||||
spatial_compression_ratio: int = 8
|
||||
|
||||
|
||||
@dataclass
|
||||
class AutoencoderKLVAEConfig(VAEConfig):
|
||||
arch_config: VAEArchConfig = field(default_factory=AutoencoderKLArchConfig)
|
||||
@@ -5,7 +5,9 @@ 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.lingbotworld import LingBotWorldI2V480PConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.sd35 import SD35Config
|
||||
from fastvideo.registry import get_pipeline_config_cls_from_name
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
from fastvideo.configs.pipelines.wan import (SelfForcingWanT2V480PConfig,
|
||||
@@ -18,5 +20,6 @@ __all__ = [
|
||||
"WanT2V480PConfig", "WanI2V480PConfig", "WanT2V720PConfig",
|
||||
"WanI2V720PConfig", "StepVideoT2VConfig", "SelfForcingWanT2V480PConfig",
|
||||
"CosmosConfig", "Cosmos25Config", "LTX2T2VConfig", "HYWorldConfig",
|
||||
"SD35Config", "LingBotWorldI2V480PConfig",
|
||||
"get_pipeline_config_cls_from_name"
|
||||
]
|
||||
|
||||
@@ -0,0 +1,13 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
from fastvideo.configs.models import DiTConfig
|
||||
from fastvideo.configs.pipelines.wan import Wan2_2_I2V_A14B_Config
|
||||
from fastvideo.configs.models.dits.lingbotworld import LingBotWorldVideoConfig
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorldI2V480PConfig(Wan2_2_I2V_A14B_Config):
|
||||
dit_config: DiTConfig = field(default_factory=LingBotWorldVideoConfig)
|
||||
flow_shift: float | None = 10.0
|
||||
boundary_ratio: float | None = 0.947
|
||||
@@ -0,0 +1,65 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, field
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models import EncoderConfig
|
||||
from fastvideo.configs.models.encoders import (
|
||||
BaseEncoderOutput,
|
||||
CLIPTextConfig,
|
||||
T5Config,
|
||||
)
|
||||
from fastvideo.configs.models.dits.sd3 import SD3DiTConfig
|
||||
from fastvideo.configs.models.vaes.autoencoder_kl import AutoencoderKLVAEConfig
|
||||
from fastvideo.configs.pipelines.base import PipelineConfig, preprocess_text
|
||||
|
||||
|
||||
def _sd35_text_postprocess(outputs: BaseEncoderOutput) -> torch.Tensor:
|
||||
assert outputs.last_hidden_state is not None
|
||||
return outputs.last_hidden_state
|
||||
|
||||
|
||||
@dataclass
|
||||
class SD35Config(PipelineConfig):
|
||||
|
||||
scheduler_arch: str = "FlowMatchEulerDiscreteScheduler"
|
||||
transformer_arch: str = "SD3Transformer2DModel"
|
||||
vae_arch: str = "AutoencoderKL"
|
||||
text_encoder_archs: tuple[str, ...] = (
|
||||
"CLIPTextModelWithProjection",
|
||||
"CLIPTextModelWithProjection",
|
||||
"T5EncoderModel",
|
||||
)
|
||||
tokenizer_archs: tuple[str, ...] = (
|
||||
"CLIPTokenizer",
|
||||
"CLIPTokenizer",
|
||||
"T5TokenizerFast",
|
||||
)
|
||||
|
||||
dit_config: SD3DiTConfig = field(default_factory=SD3DiTConfig)
|
||||
vae_config: AutoencoderKLVAEConfig = field(
|
||||
default_factory=AutoencoderKLVAEConfig)
|
||||
|
||||
embedded_cfg_scale: float = 0.0
|
||||
flow_shift: float | None = None
|
||||
|
||||
text_encoder_configs: tuple[EncoderConfig, ...] = field(
|
||||
default_factory=lambda:
|
||||
(CLIPTextConfig(), CLIPTextConfig(), T5Config()))
|
||||
preprocess_text_funcs: tuple[Callable[[str], str], ...] = field(
|
||||
default_factory=lambda:
|
||||
(preprocess_text, preprocess_text, preprocess_text))
|
||||
postprocess_text_funcs: tuple[
|
||||
Callable[[BaseEncoderOutput], torch.Tensor],
|
||||
...] = field(default_factory=lambda:
|
||||
(_sd35_text_postprocess, _sd35_text_postprocess,
|
||||
_sd35_text_postprocess))
|
||||
|
||||
dit_precision: str = "bf16"
|
||||
vae_precision: str = "fp32"
|
||||
text_encoder_precisions: tuple[str, ...] = field(
|
||||
default_factory=lambda: ("fp32", "fp32", "bf16"))
|
||||
@@ -31,6 +31,9 @@ class SamplingParam:
|
||||
# Camera control inputs (HYWorld)
|
||||
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
|
||||
|
||||
# Camera control inputs (LingBotWorld)
|
||||
c2ws_plucker_emb: Any | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
|
||||
|
||||
# Refine inputs (LongCat 480p->720p upscaling)
|
||||
# Path-based refine (load stage1 video from disk, e.g. MP4)
|
||||
refine_from: str | None = None # Path to stage1 video (480p output from distill)
|
||||
|
||||
@@ -0,0 +1,21 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
from dataclasses import dataclass
|
||||
from fastvideo.configs.sample.wan import Wan2_2_I2V_A14B_SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class LingBotWorld_SamplingParam(Wan2_2_I2V_A14B_SamplingParam):
|
||||
guidance_scale: float = 5.0 # high_noise
|
||||
guidance_scale_2: float = 5.0 # low_noise
|
||||
num_inference_steps: int = 70
|
||||
boundary_ratio: float | None = 0.947
|
||||
negative_prompt: str | None = (
|
||||
"画面突变,色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,"
|
||||
"最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,"
|
||||
"畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走,"
|
||||
"镜头晃动,画面闪烁,模糊,噪点,水印,签名,文字,变形,扭曲,液化,不合逻辑的结构,卡顿,"
|
||||
"PPT幻灯片感,过暗,欠曝,低对比度,霓虹灯光感,过度锐化,3D渲染感,人物,行人,游客,身体,"
|
||||
"皮肤,肢体,面部特征,汽车,电线")
|
||||
fps: int = 16
|
||||
# NOTE(will): default boundary timestep is tracked by PipelineConfig, but
|
||||
# can be overridden during sampling
|
||||
@@ -5,10 +5,44 @@ from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2SamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 distilled T2V.
|
||||
class LTX2BaseSamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 base one-stage T2V.
|
||||
|
||||
Values follow the official LTX-2 one-stage defaults.
|
||||
"""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 512
|
||||
width: int = 768
|
||||
fps: int = 24
|
||||
num_inference_steps: int = 40
|
||||
guidance_scale: float = 3.0
|
||||
# Copied/following official LTX-2 DEFAULT_NEGATIVE_PROMPT.
|
||||
negative_prompt: str = (
|
||||
"blurry, out of focus, overexposed, underexposed, low contrast, "
|
||||
"washed out colors, excessive noise, grainy texture, poor lighting, "
|
||||
"flickering, motion blur, distorted proportions, unnatural skin "
|
||||
"tones, deformed facial features, asymmetrical face, missing facial "
|
||||
"features, extra limbs, disfigured hands, wrong hand count, "
|
||||
"artifacts around text, inconsistent perspective, camera shake, "
|
||||
"incorrect depth of field, background too sharp, background clutter, "
|
||||
"distracting reflections, harsh shadows, inconsistent lighting "
|
||||
"direction, color banding, cartoonish rendering, 3D CGI look, "
|
||||
"unrealistic materials, uncanny valley effect, incorrect ethnicity, "
|
||||
"wrong gender, exaggerated expressions, wrong gaze direction, "
|
||||
"mismatched lip sync, silent or muted audio, distorted voice, "
|
||||
"robotic voice, echo, background noise, off-sync audio, incorrect "
|
||||
"dialogue, added dialogue, repetitive speech, jittery movement, "
|
||||
"awkward pauses, incorrect timing, unnatural transitions, "
|
||||
"inconsistent framing, tilted camera, flat lighting, inconsistent "
|
||||
"tone, cinematic oversaturation, stylized filters, or AI artifacts.")
|
||||
|
||||
|
||||
@dataclass
|
||||
class LTX2DistilledSamplingParam(SamplingParam):
|
||||
"""Default sampling parameters for LTX-2 distilled one-stage T2V."""
|
||||
|
||||
seed: int = 10
|
||||
num_frames: int = 121
|
||||
height: int = 1024
|
||||
@@ -18,3 +52,7 @@ class LTX2SamplingParam(SamplingParam):
|
||||
guidance_scale: float = 1.0
|
||||
# No default negative_prompt for distilled models
|
||||
negative_prompt: str = ""
|
||||
|
||||
|
||||
# Backward compatibility alias.
|
||||
LTX2SamplingParam = LTX2DistilledSamplingParam
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
|
||||
|
||||
@dataclass
|
||||
class SD35SamplingParam(SamplingParam):
|
||||
|
||||
prompt: str | None = "a photo of a cat"
|
||||
negative_prompt: str = ""
|
||||
|
||||
num_videos_per_prompt: int = 1
|
||||
seed: int = 0
|
||||
|
||||
num_frames: int = 1
|
||||
height: int = 512
|
||||
width: int = 512
|
||||
fps: int = 1
|
||||
|
||||
num_inference_steps: int = 28
|
||||
guidance_scale: float = 6.0
|
||||
@@ -4,6 +4,8 @@ from torchvision.transforms import Lambda
|
||||
|
||||
from fastvideo.dataset.parquet_dataset_map_style import (
|
||||
build_parquet_map_style_dataloader)
|
||||
from fastvideo.dataset.ltx2_precomputed_dataset import (
|
||||
build_ltx2_precomputed_dataloader, LTX2PrecomputedDataset)
|
||||
from fastvideo.dataset.preprocessing_datasets import VideoCaptionMergedDataset, TextDataset
|
||||
from fastvideo.dataset.transform import (CenterCropResizeVideo, Normalize255,
|
||||
TemporalRandomCrop)
|
||||
@@ -46,6 +48,10 @@ def gettextdataset(args) -> TextDataset:
|
||||
|
||||
|
||||
__all__ = [
|
||||
"build_parquet_map_style_dataloader", "ValidationDataset",
|
||||
"VideoCaptionMergedDataset", "TextDataset"
|
||||
"build_parquet_map_style_dataloader",
|
||||
"build_ltx2_precomputed_dataloader",
|
||||
"LTX2PrecomputedDataset",
|
||||
"ValidationDataset",
|
||||
"VideoCaptionMergedDataset",
|
||||
"TextDataset",
|
||||
]
|
||||
|
||||
@@ -0,0 +1,210 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Dataset utilities for loading LTX2 precomputed training artifacts.
|
||||
#
|
||||
# Usage:
|
||||
# - Input root can be either `<data_root>/` or `<data_root>/.precomputed/`.
|
||||
# - Required sources are `latents/` and `conditions/` with matching `.pt` files.
|
||||
# - Optional source `audio_latents/` is loaded when provided in `data_sources`.
|
||||
# - `build_ltx2_precomputed_dataloader(...)` is the intended entrypoint used by
|
||||
# `fastvideo/training/ltx2_training_pipeline.py`.
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from pathlib import Path
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
from einops import rearrange
|
||||
from torch.utils.data import Dataset
|
||||
from torchdata.stateful_dataloader import StatefulDataLoader
|
||||
|
||||
from fastvideo.dataset.parquet_dataset_map_style import DP_SP_BatchSampler
|
||||
from fastvideo.distributed import get_sp_world_size, get_world_rank, get_world_size
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
PRECOMPUTED_DIR_NAME = ".precomputed"
|
||||
|
||||
|
||||
class LTX2PrecomputedDataset(Dataset):
|
||||
"""Dataset for LTX-2 precomputed latents and conditions.
|
||||
|
||||
Expected directory structure (data_root):
|
||||
.precomputed/
|
||||
latents/*.pt
|
||||
conditions/*.pt
|
||||
audio_latents/*.pt (optional)
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
data_root: str,
|
||||
data_sources: dict[str, str] | list[str] | None = None,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.data_root = self._setup_data_root(data_root)
|
||||
self.data_sources = self._normalize_data_sources(data_sources)
|
||||
self.source_paths = self._setup_source_paths()
|
||||
self.sample_files = self._discover_samples()
|
||||
self._validate_setup()
|
||||
|
||||
@staticmethod
|
||||
def _setup_data_root(data_root: str) -> Path:
|
||||
data_root_path = Path(data_root).expanduser().resolve()
|
||||
if not data_root_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Data root directory does not exist: {data_root_path}")
|
||||
if (data_root_path / PRECOMPUTED_DIR_NAME).exists():
|
||||
data_root_path = data_root_path / PRECOMPUTED_DIR_NAME
|
||||
return data_root_path
|
||||
|
||||
@staticmethod
|
||||
def _normalize_data_sources(
|
||||
data_sources: dict[str, str] | list[str] | None,
|
||||
) -> dict[str, str]:
|
||||
if data_sources is None:
|
||||
return {"latents": "latents", "conditions": "conditions"}
|
||||
if isinstance(data_sources, list):
|
||||
return {source: source for source in data_sources}
|
||||
if isinstance(data_sources, dict):
|
||||
return data_sources.copy()
|
||||
raise TypeError(
|
||||
f"data_sources must be dict, list, or None, got {type(data_sources)}")
|
||||
|
||||
def _setup_source_paths(self) -> dict[str, Path]:
|
||||
source_paths: dict[str, Path] = {}
|
||||
for dir_name in self.data_sources:
|
||||
source_path = self.data_root / dir_name
|
||||
if not source_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Required {dir_name} directory does not exist: {source_path}")
|
||||
source_paths[dir_name] = source_path
|
||||
return source_paths
|
||||
|
||||
def _discover_samples(self) -> dict[str, list[Path]]:
|
||||
data_key = ("latents"
|
||||
if "latents" in self.data_sources else next(iter(
|
||||
self.data_sources.keys())))
|
||||
data_path = self.source_paths[data_key]
|
||||
data_files = list(data_path.glob("**/*.pt"))
|
||||
if not data_files:
|
||||
raise ValueError(f"No data files found in {data_path}")
|
||||
|
||||
sample_files = {output_key: [] for output_key in self.data_sources.values()}
|
||||
for data_file in data_files:
|
||||
rel_path = data_file.relative_to(data_path)
|
||||
if self._all_source_files_exist(data_file, rel_path):
|
||||
self._fill_sample_data_files(data_file, rel_path, sample_files)
|
||||
return sample_files
|
||||
|
||||
def _all_source_files_exist(self, data_file: Path, rel_path: Path) -> bool:
|
||||
for dir_name in self.data_sources:
|
||||
expected_path = self._get_expected_file_path(dir_name, data_file,
|
||||
rel_path)
|
||||
if not expected_path.exists():
|
||||
logger.warning(
|
||||
"No matching %s file found for: %s (expected in: %s)",
|
||||
dir_name,
|
||||
data_file.name,
|
||||
expected_path,
|
||||
)
|
||||
return False
|
||||
return True
|
||||
|
||||
def _get_expected_file_path(self, dir_name: str, data_file: Path,
|
||||
rel_path: Path) -> Path:
|
||||
source_path = self.source_paths[dir_name]
|
||||
if dir_name == "conditions" and data_file.name.startswith("latent_"):
|
||||
return source_path / f"condition_{data_file.stem[7:]}.pt"
|
||||
return source_path / rel_path
|
||||
|
||||
def _fill_sample_data_files(self, data_file: Path, rel_path: Path,
|
||||
sample_files: dict[str, list[Path]]) -> None:
|
||||
for dir_name, output_key in self.data_sources.items():
|
||||
expected_path = self._get_expected_file_path(dir_name, data_file,
|
||||
rel_path)
|
||||
sample_files[output_key].append(
|
||||
expected_path.relative_to(self.source_paths[dir_name]))
|
||||
|
||||
def _validate_setup(self) -> None:
|
||||
if not self.sample_files:
|
||||
raise ValueError(
|
||||
"No valid samples found - all data sources must have matching files"
|
||||
)
|
||||
sample_counts = {
|
||||
key: len(files)
|
||||
for key, files in self.sample_files.items()
|
||||
}
|
||||
if len(set(sample_counts.values())) > 1:
|
||||
raise ValueError(
|
||||
f"Mismatched sample counts across sources: {sample_counts}")
|
||||
|
||||
def __len__(self) -> int:
|
||||
first_key = next(iter(self.sample_files.keys()))
|
||||
return len(self.sample_files[first_key])
|
||||
|
||||
def __getitem__(self, index: int) -> dict[str, torch.Tensor]:
|
||||
result: dict[str, Any] = {}
|
||||
for dir_name, output_key in self.data_sources.items():
|
||||
source_path = self.source_paths[dir_name]
|
||||
file_rel_path = self.sample_files[output_key][index]
|
||||
file_path = source_path / file_rel_path
|
||||
try:
|
||||
data = torch.load(file_path, map_location="cpu", weights_only=True)
|
||||
if "latent" in dir_name.lower():
|
||||
data = self._normalize_video_latents(data)
|
||||
result[output_key] = data
|
||||
except Exception as e:
|
||||
raise RuntimeError(
|
||||
f"Failed to load {output_key} from {file_path}: {e}") from e
|
||||
result["idx"] = index
|
||||
return result
|
||||
|
||||
@staticmethod
|
||||
def _normalize_video_latents(data: dict) -> dict:
|
||||
latents = data["latents"]
|
||||
if latents.dim() == 2:
|
||||
num_frames = data["num_frames"]
|
||||
height = data["height"]
|
||||
width = data["width"]
|
||||
latents = rearrange(
|
||||
latents,
|
||||
"(f h w) c -> c f h w",
|
||||
f=num_frames,
|
||||
h=height,
|
||||
w=width,
|
||||
)
|
||||
data = data.copy()
|
||||
data["latents"] = latents
|
||||
return data
|
||||
|
||||
|
||||
def build_ltx2_precomputed_dataloader(
|
||||
path: str,
|
||||
batch_size: int,
|
||||
num_data_workers: int,
|
||||
data_sources: dict[str, str] | list[str] | None = None,
|
||||
drop_last: bool = True,
|
||||
seed: int = 42,
|
||||
) -> tuple[LTX2PrecomputedDataset, StatefulDataLoader]:
|
||||
dataset = LTX2PrecomputedDataset(path, data_sources=data_sources)
|
||||
sampler = DP_SP_BatchSampler(
|
||||
batch_size=batch_size,
|
||||
dataset_size=len(dataset),
|
||||
num_sp_groups=get_world_size() // get_sp_world_size(),
|
||||
sp_world_size=get_sp_world_size(),
|
||||
global_rank=get_world_rank(),
|
||||
drop_last=drop_last,
|
||||
drop_first_row=False,
|
||||
seed=seed,
|
||||
)
|
||||
loader = StatefulDataLoader(
|
||||
dataset,
|
||||
batch_sampler=sampler,
|
||||
collate_fn=None,
|
||||
num_workers=num_data_workers,
|
||||
pin_memory=True,
|
||||
persistent_workers=num_data_workers > 0,
|
||||
)
|
||||
return dataset, loader
|
||||
@@ -223,64 +223,82 @@ class VideoGenerator:
|
||||
sampling_param=sampling_param,
|
||||
**kwargs)
|
||||
|
||||
def _is_image_workload(self) -> bool:
|
||||
"""Return True when the workload produces a single image (t2i, i2i …)."""
|
||||
args = getattr(self, "fastvideo_args", None)
|
||||
if args is None:
|
||||
return False
|
||||
return args.workload_type.value.endswith("2i")
|
||||
|
||||
def _prepare_output_path(
|
||||
self,
|
||||
output_path: str,
|
||||
prompt: str,
|
||||
) -> str:
|
||||
"""Build a unique, sanitized .mp4 output file path.
|
||||
"""Build a unique, sanitized output file path.
|
||||
|
||||
- If `output_path` ends with .mp4 (case-insensitive), treat it as a file path.
|
||||
- Otherwise, treat `output_path` as a directory and derive the filename
|
||||
from the prompt.
|
||||
The file extension is chosen automatically based on the workload type:
|
||||
``.png`` for image workloads (``t2i``, ``i2i``, …) and ``.mp4`` for
|
||||
video workloads.
|
||||
|
||||
- If ``output_path`` already carries the correct extension, treat it
|
||||
as a file path.
|
||||
- Otherwise, treat ``output_path`` as a directory and derive the
|
||||
filename from the prompt.
|
||||
- Invalid filename characters are removed; if the name changes, a
|
||||
warning is logged.
|
||||
- If the target path already exists, a numeric suffix is appended.
|
||||
"""
|
||||
target_ext = ".png" if self._is_image_workload() else ".mp4"
|
||||
|
||||
def _sanitize_filename_component(name: str) -> str:
|
||||
# Remove characters invalid on common filesystems, strip spaces/dots
|
||||
sanitized = re.sub(r'[\\/:*?"<>|]', '', name)
|
||||
sanitized = sanitized.strip().strip('.')
|
||||
sanitized = re.sub(r'\s+', ' ', sanitized)
|
||||
return sanitized or "video"
|
||||
return sanitized or "output"
|
||||
|
||||
base_path, extension = os.path.splitext(output_path)
|
||||
extension_lower = extension.lower()
|
||||
|
||||
if extension_lower == ".mp4":
|
||||
if extension_lower == target_ext:
|
||||
output_dir = os.path.dirname(output_path)
|
||||
base_name = os.path.basename(
|
||||
base_path) # filename without extension
|
||||
sanitized_base = _sanitize_filename_component(base_name)
|
||||
if sanitized_base != base_name:
|
||||
logger.warning(
|
||||
"The video name '%s' contained invalid characters. It has been renamed to '%s.mp4'",
|
||||
"The output name '%s' contained invalid characters. "
|
||||
"It has been renamed to '%s%s'",
|
||||
os.path.basename(output_path),
|
||||
sanitized_base,
|
||||
target_ext,
|
||||
)
|
||||
video_name = f"{sanitized_base}.mp4"
|
||||
out_name = f"{sanitized_base}{target_ext}"
|
||||
else:
|
||||
# Treat as directory; inform if an unexpected extension was provided.
|
||||
# Treat as directory; inform if an unexpected extension was
|
||||
# provided.
|
||||
if extension:
|
||||
logger.info(
|
||||
"Output path '%s' has non-mp4 extension '%s'; treating it as a directory and using a .mp4 filename derived from the prompt",
|
||||
"Output path '%s' has extension '%s' which does not "
|
||||
"match the target '%s'; treating it as a directory",
|
||||
output_path,
|
||||
extension,
|
||||
target_ext,
|
||||
)
|
||||
output_dir = output_path
|
||||
prompt_component = _sanitize_filename_component(prompt[:100])
|
||||
video_name = f"{prompt_component}.mp4"
|
||||
out_name = f"{prompt_component}{target_ext}"
|
||||
|
||||
if output_dir:
|
||||
os.makedirs(output_dir, exist_ok=True)
|
||||
|
||||
new_output_path = os.path.join(output_dir, video_name)
|
||||
new_output_path = os.path.join(output_dir, out_name)
|
||||
counter = 1
|
||||
while os.path.exists(new_output_path):
|
||||
name_part, ext_part = os.path.splitext(video_name)
|
||||
new_video_name = f"{name_part}_{counter}{ext_part}"
|
||||
new_output_path = os.path.join(output_dir, new_video_name)
|
||||
name_part, ext_part = os.path.splitext(out_name)
|
||||
new_name = f"{name_part}_{counter}{ext_part}"
|
||||
new_output_path = os.path.join(output_dir, new_name)
|
||||
counter += 1
|
||||
return new_output_path
|
||||
|
||||
@@ -426,15 +444,25 @@ class VideoGenerator:
|
||||
x = x.transpose(0, 1).transpose(1, 2).squeeze(-1)
|
||||
frames.append((x * 255).numpy().astype(np.uint8))
|
||||
|
||||
# Save video if requested
|
||||
# Save output if requested
|
||||
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 self._is_image_workload():
|
||||
# Image workloads (t2i, i2i, …): save the first frame as PNG.
|
||||
imageio.imwrite(output_path, frames[0])
|
||||
logger.info("Saved image to %s", output_path)
|
||||
else:
|
||||
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
|
||||
|
||||
@@ -903,6 +903,7 @@ class TrainingArgs(FastVideoArgs):
|
||||
lora_rank: int | None = None
|
||||
lora_alpha: int | None = None
|
||||
lora_training: bool = False
|
||||
ltx2_first_frame_conditioning_p: float = 0.1
|
||||
|
||||
# distillation args
|
||||
generator_update_interval: int = 5
|
||||
@@ -1257,6 +1258,13 @@ class TrainingArgs(FastVideoArgs):
|
||||
help="Whether to use LoRA training")
|
||||
parser.add_argument("--lora-rank", type=int, help="LoRA rank")
|
||||
parser.add_argument("--lora-alpha", type=int, help="LoRA alpha")
|
||||
parser.add_argument(
|
||||
"--ltx2-first-frame-conditioning-p",
|
||||
type=float,
|
||||
default=TrainingArgs.ltx2_first_frame_conditioning_p,
|
||||
help=
|
||||
"Probability of conditioning on the first frame during LTX-2 training",
|
||||
)
|
||||
|
||||
# V-MoBA parameters
|
||||
parser.add_argument(
|
||||
|
||||
@@ -60,6 +60,61 @@ class PatchEmbed(nn.Module):
|
||||
return x
|
||||
|
||||
|
||||
class WanCamControlPatchEmbedding(nn.Module):
|
||||
"""Lingbot World Patch embedding for Plucker features."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
patch_size=(1, 2, 2),
|
||||
in_chans=384, # 6 * 64
|
||||
embed_dim=2048,
|
||||
bias=True,
|
||||
dtype=None,
|
||||
prefix: str = ""):
|
||||
super().__init__()
|
||||
# must be 3-tuple
|
||||
if isinstance(patch_size, list | tuple):
|
||||
if len(patch_size) != 3:
|
||||
raise ValueError(
|
||||
f"patch_size must have length 3, got {len(patch_size)}")
|
||||
else:
|
||||
raise ValueError(f"Unsupported patch_size type: {type(patch_size)}")
|
||||
|
||||
self.patch_size = patch_size
|
||||
pt, ph, pw = self.patch_size
|
||||
self.in_features = in_chans * pt * ph * pw
|
||||
self.proj = nn.Linear(self.in_features,
|
||||
embed_dim,
|
||||
bias=bias,
|
||||
dtype=dtype)
|
||||
|
||||
def forward(self, x: torch.Tensor) -> torch.Tensor:
|
||||
if x.dim() != 5:
|
||||
raise ValueError(
|
||||
f"Expected camera embedding shape [B, C, F, H, W], got {x.shape}"
|
||||
)
|
||||
bsz, channels, frames, height, width = x.shape
|
||||
pt, ph, pw = self.patch_size
|
||||
if (frames % pt) != 0 or (height % ph) != 0 or (width % pw) != 0:
|
||||
raise ValueError(
|
||||
f"Input shape {x.shape} must be divisible by patch_size {self.patch_size}"
|
||||
)
|
||||
|
||||
# '1 c (f c1) (h c2) (w c3) -> 1 (f h w) (c c1 c2 c3)',
|
||||
x = x.view(
|
||||
bsz,
|
||||
channels,
|
||||
frames // pt,
|
||||
pt,
|
||||
height // ph,
|
||||
ph,
|
||||
width // pw,
|
||||
pw,
|
||||
)
|
||||
x = x.permute(0, 2, 4, 6, 1, 3, 5, 7).reshape(bsz, -1, self.in_features)
|
||||
return self.proj(x)
|
||||
|
||||
|
||||
class TimestepEmbedder(nn.Module):
|
||||
"""
|
||||
Embeds scalar timesteps into vector representations.
|
||||
@@ -252,4 +307,4 @@ class Timesteps(nn.Module):
|
||||
downscale_freq_shift=self.downscale_freq_shift,
|
||||
scale=self.scale,
|
||||
)
|
||||
return t_emb
|
||||
return t_emb
|
||||
|
||||
@@ -0,0 +1,61 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""Audio preprocessing helpers for LTX-2 training."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
from torch import nn
|
||||
|
||||
|
||||
class AudioProcessor(nn.Module):
|
||||
"""Converts audio waveforms to log-mel spectrograms with resampling."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
sample_rate: int,
|
||||
mel_bins: int,
|
||||
mel_hop_length: int,
|
||||
n_fft: int,
|
||||
) -> None:
|
||||
super().__init__()
|
||||
self.sample_rate = sample_rate
|
||||
self.mel_transform = torchaudio.transforms.MelSpectrogram(
|
||||
sample_rate=sample_rate,
|
||||
n_fft=n_fft,
|
||||
win_length=n_fft,
|
||||
hop_length=mel_hop_length,
|
||||
f_min=0.0,
|
||||
f_max=sample_rate / 2.0,
|
||||
n_mels=mel_bins,
|
||||
window_fn=torch.hann_window,
|
||||
center=True,
|
||||
pad_mode="reflect",
|
||||
power=1.0,
|
||||
mel_scale="slaney",
|
||||
norm="slaney",
|
||||
)
|
||||
|
||||
def resample_waveform(
|
||||
self,
|
||||
waveform: torch.Tensor,
|
||||
source_rate: int,
|
||||
target_rate: int,
|
||||
) -> torch.Tensor:
|
||||
if source_rate == target_rate:
|
||||
return waveform
|
||||
resampled = torchaudio.functional.resample(
|
||||
waveform, source_rate, target_rate)
|
||||
return resampled.to(device=waveform.device, dtype=waveform.dtype)
|
||||
|
||||
def waveform_to_mel(
|
||||
self,
|
||||
waveform: torch.Tensor,
|
||||
waveform_sample_rate: int,
|
||||
) -> torch.Tensor:
|
||||
waveform = self.resample_waveform(
|
||||
waveform, waveform_sample_rate, self.sample_rate)
|
||||
mel = self.mel_transform(waveform)
|
||||
mel = torch.log(torch.clamp(mel, min=1e-5))
|
||||
mel = mel.to(device=waveform.device, dtype=waveform.dtype)
|
||||
return mel.permute(0, 1, 3, 2).contiguous()
|
||||
@@ -0,0 +1,8 @@
|
||||
from .model import LingBotWorldTransformer3DModel
|
||||
|
||||
__all__ = [
|
||||
"LingBotWorldTransformer3DModel",
|
||||
]
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = [LingBotWorldTransformer3DModel]
|
||||
@@ -0,0 +1,203 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
# Adapted from LingBot World: https://github.com/Robbyant/lingbot-world/blob/main/wan/utils/cam_utils.py
|
||||
|
||||
import numpy as np
|
||||
import os
|
||||
import torch
|
||||
from scipy.interpolate import interp1d
|
||||
from scipy.spatial.transform import Rotation, Slerp
|
||||
|
||||
|
||||
# --- Official Code (Leave Unchanged) ---
|
||||
|
||||
def interpolate_camera_poses(
|
||||
src_indices: np.ndarray,
|
||||
src_rot_mat: np.ndarray,
|
||||
src_trans_vec: np.ndarray,
|
||||
tgt_indices: np.ndarray,
|
||||
) -> torch.Tensor:
|
||||
# interpolate translation
|
||||
interp_func_trans = interp1d(
|
||||
src_indices,
|
||||
src_trans_vec,
|
||||
axis=0,
|
||||
kind='linear',
|
||||
bounds_error=False,
|
||||
fill_value="extrapolate",
|
||||
)
|
||||
interpolated_trans_vec = interp_func_trans(tgt_indices)
|
||||
|
||||
# interpolate rotation
|
||||
src_quat_vec = Rotation.from_matrix(src_rot_mat)
|
||||
# ensure there is no sudden change in qw
|
||||
quats = src_quat_vec.as_quat().copy() # [N, 4]
|
||||
for i in range(1, len(quats)):
|
||||
if np.dot(quats[i], quats[i-1]) < 0:
|
||||
quats[i] = -quats[i]
|
||||
src_quat_vec = Rotation.from_quat(quats)
|
||||
slerp_func_rot = Slerp(src_indices, src_quat_vec)
|
||||
interpolated_rot_quat = slerp_func_rot(tgt_indices)
|
||||
interpolated_rot_mat = interpolated_rot_quat.as_matrix()
|
||||
|
||||
poses = np.zeros((len(tgt_indices), 4, 4))
|
||||
poses[:, :3, :3] = interpolated_rot_mat
|
||||
poses[:, :3, 3] = interpolated_trans_vec
|
||||
poses[:, 3, 3] = 1.0
|
||||
return torch.from_numpy(poses).float()
|
||||
|
||||
|
||||
def SE3_inverse(T: torch.Tensor) -> torch.Tensor:
|
||||
Rot = T[:, :3, :3] # [B,3,3]
|
||||
trans = T[:, :3, 3:] # [B,3,1]
|
||||
R_inv = Rot.transpose(-1, -2)
|
||||
t_inv = -torch.bmm(R_inv, trans)
|
||||
T_inv = torch.eye(4, device=T.device, dtype=T.dtype)[None, :, :].repeat(T.shape[0], 1, 1)
|
||||
T_inv[:, :3, :3] = R_inv
|
||||
T_inv[:, :3, 3:] = t_inv
|
||||
return T_inv
|
||||
|
||||
|
||||
def compute_relative_poses(
|
||||
c2ws_mat: torch.Tensor,
|
||||
framewise: bool = False,
|
||||
normalize_trans: bool = True,
|
||||
) -> torch.Tensor:
|
||||
ref_w2cs = SE3_inverse(c2ws_mat[0:1])
|
||||
relative_poses = torch.matmul(ref_w2cs, c2ws_mat)
|
||||
# ensure identity matrix for 1st frame
|
||||
relative_poses[0] = torch.eye(4, device=c2ws_mat.device, dtype=c2ws_mat.dtype)
|
||||
if framewise:
|
||||
# compute pose between i and i+1
|
||||
relative_poses_framewise = torch.bmm(SE3_inverse(relative_poses[:-1]), relative_poses[1:])
|
||||
relative_poses[1:] = relative_poses_framewise
|
||||
if normalize_trans: # note refer to camctrl2: "we scale the coordinate inputs to roughly 1 standard deviation to simplify model learning."
|
||||
translations = relative_poses[:, :3, 3] # [f, 3]
|
||||
max_norm = torch.norm(translations, dim=-1).max()
|
||||
# only normlaize when moving
|
||||
if max_norm > 0:
|
||||
relative_poses[:, :3, 3] = translations / max_norm
|
||||
return relative_poses
|
||||
|
||||
|
||||
@torch.no_grad()
|
||||
def create_meshgrid(n_frames: int, height: int, width: int, bias: float = 0.5, device='cuda', dtype=torch.float32) -> torch.Tensor:
|
||||
x_range = torch.arange(width, device=device, dtype=dtype)
|
||||
y_range = torch.arange(height, device=device, dtype=dtype)
|
||||
grid_y, grid_x = torch.meshgrid(y_range, x_range, indexing='ij')
|
||||
grid_xy = torch.stack([grid_x, grid_y], dim=-1).view([-1, 2]) + bias # [h*w, 2]
|
||||
grid_xy = grid_xy[None, ...].repeat(n_frames, 1, 1) # [f, h*w, 2]
|
||||
return grid_xy
|
||||
|
||||
|
||||
def get_plucker_embeddings(
|
||||
c2ws_mat: torch.Tensor,
|
||||
Ks: torch.Tensor,
|
||||
height: int,
|
||||
width: int,
|
||||
):
|
||||
n_frames = c2ws_mat.shape[0]
|
||||
grid_xy = create_meshgrid(n_frames, height, width, device=c2ws_mat.device, dtype=c2ws_mat.dtype) # [f, h*w, 2]
|
||||
fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1]
|
||||
|
||||
i = grid_xy[..., 0] # [f, h*w]
|
||||
j = grid_xy[..., 1] # [f, h*w]
|
||||
zs = torch.ones_like(i) # [f, h*w]
|
||||
xs = (i - cx) / fx * zs
|
||||
ys = (j - cy) / fy * zs
|
||||
|
||||
directions = torch.stack([xs, ys, zs], dim=-1) # [f, h*w, 3]
|
||||
directions = directions / directions.norm(dim=-1, keepdim=True) # [f, h*w, 3]
|
||||
|
||||
rays_d = directions @ c2ws_mat[:, :3, :3].transpose(-1, -2) # [f, h*w, 3]
|
||||
rays_o = c2ws_mat[:, :3, 3] # [f, 3]
|
||||
rays_o = rays_o[:, None, :].expand_as(rays_d) # [f, h*w, 3]
|
||||
# rays_dxo = torch.cross(rays_o, rays_d, dim=-1) # [f, h*w, 3]
|
||||
# note refer to: apt2
|
||||
plucker_embeddings = torch.cat([rays_o, rays_d], dim=-1) # [f, h*w, 6]
|
||||
plucker_embeddings = plucker_embeddings.view([n_frames, height, width, 6]) # [f*h*w, 6]
|
||||
return plucker_embeddings
|
||||
|
||||
|
||||
def get_Ks_transformed(
|
||||
Ks: torch.Tensor,
|
||||
height_org: int,
|
||||
width_org: int,
|
||||
height_resize: int,
|
||||
width_resize: int,
|
||||
height_final: int,
|
||||
width_final: int,
|
||||
):
|
||||
fx, fy, cx, cy = Ks.chunk(4, dim=-1) # [f, 1]
|
||||
|
||||
scale_x = width_resize / width_org
|
||||
scale_y = height_resize / height_org
|
||||
|
||||
fx_resize = fx * scale_x
|
||||
fy_resize = fy * scale_y
|
||||
cx_resize = cx * scale_x
|
||||
cy_resize = cy * scale_y
|
||||
|
||||
crop_offset_x = (width_resize - width_final) / 2
|
||||
crop_offset_y = (height_resize - height_final) / 2
|
||||
|
||||
cx_final = cx_resize - crop_offset_x
|
||||
cy_final = cy_resize - crop_offset_y
|
||||
|
||||
Ks_transformed = torch.zeros_like(Ks)
|
||||
Ks_transformed[:, 0:1] = fx_resize
|
||||
Ks_transformed[:, 1:2] = fy_resize
|
||||
Ks_transformed[:, 2:3] = cx_final
|
||||
Ks_transformed[:, 3:4] = cy_final
|
||||
|
||||
return Ks_transformed
|
||||
|
||||
|
||||
# --- Custom ---
|
||||
|
||||
def prepare_camera_embedding(
|
||||
action_path: str,
|
||||
num_frames: int,
|
||||
height: int,
|
||||
width: int,
|
||||
spatial_scale: int = 8,
|
||||
) -> tuple[torch.Tensor, int]:
|
||||
c2ws = np.load(os.path.join(action_path, "poses.npy"))
|
||||
len_c2ws = ((len(c2ws) - 1) // 4) * 4 + 1
|
||||
num_frames = min(num_frames, len_c2ws)
|
||||
c2ws = c2ws[:num_frames]
|
||||
|
||||
Ks = torch.from_numpy(
|
||||
np.load(os.path.join(action_path, "intrinsics.npy"))
|
||||
).float()
|
||||
Ks = get_Ks_transformed(
|
||||
Ks,
|
||||
height_org=480,
|
||||
width_org=832,
|
||||
height_resize=height,
|
||||
width_resize=width,
|
||||
height_final=height,
|
||||
width_final=width,
|
||||
)
|
||||
Ks = Ks[0] # use first frame
|
||||
|
||||
len_c2ws = len(c2ws)
|
||||
num_latent_frames = (len_c2ws - 1) // 4 + 1
|
||||
c2ws_infer = interpolate_camera_poses(
|
||||
src_indices=np.linspace(0, len_c2ws - 1, len_c2ws),
|
||||
src_rot_mat=c2ws[:, :3, :3],
|
||||
src_trans_vec=c2ws[:, :3, 3],
|
||||
tgt_indices=np.linspace(0, len_c2ws - 1, num_latent_frames),
|
||||
)
|
||||
c2ws_infer = compute_relative_poses(c2ws_infer, framewise=True)
|
||||
Ks = Ks.repeat(num_latent_frames, 1)
|
||||
plucker = get_plucker_embeddings(c2ws_infer, Ks, height, width) # [F, H, W, 6]
|
||||
|
||||
# reshpae
|
||||
latent_height = height // spatial_scale
|
||||
latent_width = width // spatial_scale
|
||||
plucker = plucker.view(num_latent_frames, latent_height, spatial_scale, latent_width, spatial_scale, 6)
|
||||
plucker = plucker.permute(0, 1, 3, 5, 2, 4).contiguous()
|
||||
plucker = plucker.view(num_latent_frames, latent_height, latent_width, 6 * spatial_scale * spatial_scale)
|
||||
c2ws_plucker_emb = plucker.permute(3, 0, 1, 2).contiguous().unsqueeze(0)
|
||||
|
||||
return c2ws_plucker_emb, num_frames
|
||||
@@ -0,0 +1,569 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
import math
|
||||
from contextlib import nullcontext
|
||||
from typing import Any
|
||||
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
|
||||
from fastvideo.attention import DistributedAttention
|
||||
from fastvideo.configs.models.dits.lingbotworld import LingBotWorldVideoConfig
|
||||
from fastvideo.configs.sample.wan import WanTeaCacheParams
|
||||
from fastvideo.distributed.communication_op import (
|
||||
sequence_model_parallel_all_gather_with_unpad,
|
||||
sequence_model_parallel_shard)
|
||||
from fastvideo.forward_context import get_forward_context
|
||||
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
|
||||
RMSNorm, ScaleResidual,
|
||||
ScaleResidualLayerNormScaleShift)
|
||||
from fastvideo.layers.linear import ReplicatedLinear
|
||||
# from torch.nn import RMSNorm
|
||||
# TODO: RMSNorm ....
|
||||
from fastvideo.layers.mlp import MLP
|
||||
from fastvideo.layers.rotary_embedding import get_rotary_pos_embed
|
||||
from fastvideo.layers.visual_embedding import (PatchEmbed, WanCamControlPatchEmbedding)
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.dits.base import CachableDiT
|
||||
from fastvideo.models.dits.wanvideo import (
|
||||
WanI2VCrossAttention,
|
||||
WanT2VCrossAttention,
|
||||
WanTimeTextImageEmbedding,
|
||||
)
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
|
||||
from fastvideo.distributed.parallel_state import get_sp_world_size
|
||||
from fastvideo.distributed.utils import create_attention_mask_for_padding
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LingBotWorldCamConditioner(nn.Module):
|
||||
|
||||
def __init__(self, dim: int) -> None:
|
||||
super().__init__()
|
||||
self.cam_injector = MLP(dim, dim, dim, bias=True, act_type="silu")
|
||||
self.cam_scale_layer = nn.Linear(dim, dim)
|
||||
self.cam_shift_layer = nn.Linear(dim, dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
c2ws_plucker_emb: torch.Tensor | None,
|
||||
) -> torch.Tensor:
|
||||
if c2ws_plucker_emb is None:
|
||||
return hidden_states
|
||||
assert c2ws_plucker_emb.shape == hidden_states.shape, (
|
||||
f"c2ws_plucker_emb shape must match hidden_states shape, got "
|
||||
f"{tuple(c2ws_plucker_emb.shape)} vs {tuple(hidden_states.shape)}"
|
||||
)
|
||||
c2ws_hidden_states = self.cam_injector(c2ws_plucker_emb)
|
||||
c2ws_hidden_states = c2ws_hidden_states + c2ws_plucker_emb
|
||||
cam_scale = self.cam_scale_layer(c2ws_hidden_states)
|
||||
cam_shift = self.cam_shift_layer(c2ws_hidden_states)
|
||||
return (1.0 + cam_scale) * hidden_states + cam_shift
|
||||
|
||||
|
||||
class LingBotWorldTransformerBlock(nn.Module):
|
||||
|
||||
def __init__(self,
|
||||
dim: int,
|
||||
ffn_dim: int,
|
||||
num_heads: int,
|
||||
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 = DistributedAttention(
|
||||
num_heads=num_heads,
|
||||
head_size=dim // num_heads,
|
||||
causal=False,
|
||||
supported_attention_backends=supported_attention_backends,
|
||||
prefix=f"{prefix}.attn1")
|
||||
self.hidden_dim = dim
|
||||
self.num_attention_heads = num_heads
|
||||
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":
|
||||
# LTX applies qk norm across all heads
|
||||
self.norm_q = RMSNorm(dim, eps=eps)
|
||||
self.norm_k = RMSNorm(dim, eps=eps)
|
||||
else:
|
||||
raise NotImplementedError(
|
||||
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,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 2. Cross-attention
|
||||
if added_kv_proj_dim is not None:
|
||||
# I2V
|
||||
self.attn2 = WanI2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
else:
|
||||
# T2V
|
||||
self.attn2 = WanT2VCrossAttention(dim,
|
||||
num_heads,
|
||||
qk_norm=qk_norm,
|
||||
eps=eps)
|
||||
self.cross_attn_residual_norm = ScaleResidualLayerNormScaleShift(
|
||||
dim,
|
||||
norm_type="layer",
|
||||
eps=eps,
|
||||
elementwise_affine=False,
|
||||
dtype=torch.float32,
|
||||
compute_dtype=torch.float32)
|
||||
|
||||
# 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)
|
||||
self.cam_conditioner = LingBotWorldCamConditioner(dim)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
hidden_states: torch.Tensor,
|
||||
encoder_hidden_states: torch.Tensor,
|
||||
temb: torch.Tensor,
|
||||
freqs_cis: tuple[torch.Tensor, torch.Tensor],
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
c2ws_plucker_emb: torch.Tensor | None = None,
|
||||
) -> torch.Tensor:
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
|
||||
if temb.dim() == 4:
|
||||
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
|
||||
self.scale_shift_table.unsqueeze(0) + temb.float()
|
||||
).chunk(6, dim=2)
|
||||
# batch_size, seq_len, 1, inner_dim
|
||||
shift_msa = shift_msa.squeeze(2)
|
||||
scale_msa = scale_msa.squeeze(2)
|
||||
gate_msa = gate_msa.squeeze(2)
|
||||
c_shift_msa = c_shift_msa.squeeze(2)
|
||||
c_scale_msa = c_scale_msa.squeeze(2)
|
||||
c_gate_msa = c_gate_msa.squeeze(2)
|
||||
else:
|
||||
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
|
||||
e = self.scale_shift_table + temb.float()
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=1)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
norm_hidden_states = (self.norm1(hidden_states.float()) *
|
||||
(1 + scale_msa) + shift_msa).to(orig_dtype)
|
||||
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(query)
|
||||
if self.norm_k is not None:
|
||||
key = self.norm_k(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))
|
||||
|
||||
attn_output, _ = self.attn1(query, key, value, freqs_cis=freqs_cis, attention_mask=attention_mask)
|
||||
attn_output = attn_output.flatten(2)
|
||||
attn_output, _ = self.to_out(attn_output)
|
||||
attn_output = attn_output.squeeze(1)
|
||||
|
||||
null_shift = null_scale = torch.tensor([0], device=hidden_states.device)
|
||||
norm_hidden_states, hidden_states = self.self_attn_residual_norm(
|
||||
hidden_states, attn_output, gate_msa, null_shift, null_scale)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
# Inject camera condition
|
||||
# must be applied after the self-attention residual update.
|
||||
hidden_states = self.cam_conditioner(hidden_states, c2ws_plucker_emb)
|
||||
norm_hidden_states = self.self_attn_residual_norm.norm(hidden_states)
|
||||
norm_hidden_states = norm_hidden_states.to(orig_dtype)
|
||||
|
||||
# 2. Cross-attention
|
||||
attn_output = self.attn2(norm_hidden_states,
|
||||
context=encoder_hidden_states,
|
||||
context_lens=None)
|
||||
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
|
||||
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
|
||||
norm_hidden_states, hidden_states = norm_hidden_states.to(
|
||||
orig_dtype), hidden_states.to(orig_dtype)
|
||||
|
||||
# 3. Feed-forward
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
|
||||
|
||||
class LingBotWorldTransformer3DModel(CachableDiT):
|
||||
_fsdp_shard_conditions = LingBotWorldVideoConfig()._fsdp_shard_conditions
|
||||
_compile_conditions = LingBotWorldVideoConfig()._compile_conditions
|
||||
_supported_attention_backends = LingBotWorldVideoConfig(
|
||||
)._supported_attention_backends
|
||||
param_names_mapping = LingBotWorldVideoConfig().param_names_mapping
|
||||
reverse_param_names_mapping = LingBotWorldVideoConfig().reverse_param_names_mapping
|
||||
lora_param_names_mapping = LingBotWorldVideoConfig().lora_param_names_mapping
|
||||
|
||||
def __init__(self, config: LingBotWorldVideoConfig, 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.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.text_len = config.text_len
|
||||
|
||||
assert config.num_attention_heads % get_sp_world_size() == 0, f"The number of attention heads ({config.num_attention_heads}) must be divisible by the sequence parallel size ({get_sp_world_size()})"
|
||||
|
||||
# 1. Patch & position embedding
|
||||
self.patch_embedding = PatchEmbed(in_chans=config.in_channels,
|
||||
embed_dim=inner_dim,
|
||||
patch_size=config.patch_size,
|
||||
flatten=False)
|
||||
self.patch_embedding_wancamctrl = WanCamControlPatchEmbedding(in_chans=6 * 64,
|
||||
embed_dim=inner_dim,
|
||||
patch_size=config.patch_size)
|
||||
self.c2ws_mlp = MLP(inner_dim, inner_dim, inner_dim, bias=True, act_type="silu")
|
||||
|
||||
# 2. Condition embeddings
|
||||
self.condition_embedder = WanTimeTextImageEmbedding(
|
||||
dim=inner_dim,
|
||||
time_freq_dim=config.freq_dim,
|
||||
text_embed_dim=config.text_dim,
|
||||
image_embed_dim=config.image_dim,
|
||||
)
|
||||
|
||||
# 3. Transformer blocks
|
||||
transformer_block = LingBotWorldTransformerBlock
|
||||
self.blocks = nn.ModuleList([
|
||||
transformer_block(inner_dim,
|
||||
config.ffn_dim,
|
||||
config.num_attention_heads,
|
||||
config.qk_norm,
|
||||
config.cross_attn_norm,
|
||||
config.eps,
|
||||
config.added_kv_proj_dim,
|
||||
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,
|
||||
compute_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
|
||||
self._logged_attention_mask = False
|
||||
|
||||
# For type checking
|
||||
self.previous_e0_even = None
|
||||
self.previous_e0_odd = None
|
||||
self.previous_residual_even = None
|
||||
self.previous_residual_odd = None
|
||||
self.is_even = True
|
||||
self.should_calc_even = True
|
||||
self.should_calc_odd = True
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
self.cnt = 0
|
||||
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]
|
||||
| None = None,
|
||||
guidance=None,
|
||||
c2ws_plucker_emb: torch.Tensor | None = None,
|
||||
**kwargs) -> torch.Tensor:
|
||||
forward_batch = get_forward_context().forward_batch
|
||||
enable_teacache = forward_batch is not None and forward_batch.enable_teacache
|
||||
|
||||
orig_dtype = hidden_states.dtype
|
||||
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:
|
||||
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, 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)
|
||||
freqs_cis = (freqs_cos.to(hidden_states.device).float(),
|
||||
freqs_sin.to(hidden_states.device).float())
|
||||
|
||||
hidden_states = self.patch_embedding(hidden_states)
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
c2ws_hidden_states = None
|
||||
if c2ws_plucker_emb is not None:
|
||||
c2ws_plucker_emb = self.patch_embedding_wancamctrl(
|
||||
c2ws_plucker_emb.to(device=hidden_states.device, dtype=hidden_states.dtype)
|
||||
)
|
||||
c2ws_hidden_states = self.c2ws_mlp(c2ws_plucker_emb)
|
||||
c2ws_plucker_emb = c2ws_plucker_emb + c2ws_hidden_states
|
||||
|
||||
# Shard with padding support - returns (sharded_tensor, original_seq_len)
|
||||
hidden_states, original_seq_len = sequence_model_parallel_shard(hidden_states, dim=1)
|
||||
|
||||
# Shard c2ws_plucker_emb
|
||||
if c2ws_plucker_emb is not None:
|
||||
c2ws_plucker_emb, _ = sequence_model_parallel_shard(c2ws_plucker_emb, dim=1)
|
||||
|
||||
# Create attention mask for padded tokens if padding was applied
|
||||
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:
|
||||
if not self._logged_attention_mask:
|
||||
logger.info(f"Padding applied, original seq len: {original_seq_len}, padded seq len: {padded_seq_len}")
|
||||
self._logged_attention_mask = True
|
||||
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:
|
||||
if not self._logged_attention_mask:
|
||||
logger.info(f"Padding not applied")
|
||||
self._logged_attention_mask = True
|
||||
attention_mask = None
|
||||
|
||||
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
|
||||
if timestep.dim() == 2:
|
||||
ts_seq_len = timestep.shape[1]
|
||||
timestep = timestep.flatten() # batch_size * seq_len
|
||||
else:
|
||||
ts_seq_len = None
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
|
||||
if ts_seq_len is not None:
|
||||
# batch_size, seq_len, 6, inner_dim
|
||||
timestep_proj = timestep_proj.unflatten(2, (6, -1))
|
||||
else:
|
||||
# batch_size, 6, inner_dim
|
||||
timestep_proj = timestep_proj.unflatten(1, (6, -1))
|
||||
|
||||
if encoder_hidden_states_image is not None:
|
||||
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)
|
||||
else:
|
||||
encoder_hidden_states = encoder_hidden_states # cast to orig_dtype for MPS & NPU
|
||||
|
||||
assert encoder_hidden_states.dtype == orig_dtype
|
||||
|
||||
# 4. Transformer blocks
|
||||
# if caching is enabled, we might be able to skip the forward pass
|
||||
should_skip_forward = self.should_skip_forward_for_cached_states(
|
||||
timestep_proj=timestep_proj, temb=temb)
|
||||
|
||||
if should_skip_forward:
|
||||
print("skipping forward, cached")
|
||||
hidden_states = self.retrieve_cached_states(hidden_states)
|
||||
else:
|
||||
# if teacache is enabled, we need to cache the original hidden states
|
||||
if enable_teacache:
|
||||
original_hidden_states = hidden_states.clone()
|
||||
|
||||
if torch.is_grad_enabled() and self.gradient_checkpointing:
|
||||
for block in self.blocks:
|
||||
hidden_states = self._gradient_checkpointing_func(
|
||||
block, hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis, attention_mask, c2ws_plucker_emb)
|
||||
else:
|
||||
for block in self.blocks:
|
||||
hidden_states = block(hidden_states, encoder_hidden_states,
|
||||
timestep_proj, freqs_cis, attention_mask, c2ws_plucker_emb)
|
||||
# if teacache is enabled, we need to cache the original hidden states
|
||||
|
||||
if enable_teacache:
|
||||
self.maybe_cache_states(hidden_states, original_hidden_states)
|
||||
# 5. Output norm, projection & unpatchify
|
||||
if temb.dim() == 3:
|
||||
# batch_size, seq_len, inner_dim (wan 2.2 ti2v)
|
||||
shift, scale = (self.scale_shift_table.unsqueeze(0) + temb.unsqueeze(2)).chunk(2, dim=2)
|
||||
shift = shift.squeeze(2)
|
||||
scale = scale.squeeze(2)
|
||||
else:
|
||||
# batch_size, inner_dim
|
||||
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
|
||||
|
||||
|
||||
hidden_states = self.norm_out(hidden_states, shift, scale)
|
||||
|
||||
# Gather and unpad in one operation
|
||||
hidden_states = sequence_model_parallel_all_gather_with_unpad(
|
||||
hidden_states, original_seq_len, dim=1)
|
||||
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
|
||||
|
||||
def maybe_cache_states(self, hidden_states: torch.Tensor,
|
||||
original_hidden_states: torch.Tensor) -> None:
|
||||
if self.is_even:
|
||||
self.previous_residual_even = hidden_states.squeeze(
|
||||
0) - original_hidden_states
|
||||
else:
|
||||
self.previous_residual_odd = hidden_states.squeeze(
|
||||
0) - original_hidden_states
|
||||
|
||||
def should_skip_forward_for_cached_states(self, **kwargs) -> bool:
|
||||
|
||||
forward_context = get_forward_context()
|
||||
forward_batch = forward_context.forward_batch
|
||||
if forward_batch is None or not forward_batch.enable_teacache:
|
||||
return False
|
||||
teacache_params = forward_batch.teacache_params
|
||||
assert teacache_params is not None, "teacache_params is not initialized"
|
||||
assert isinstance(
|
||||
teacache_params,
|
||||
WanTeaCacheParams), "teacache_params is not a WanTeaCacheParams"
|
||||
current_timestep = forward_context.current_timestep
|
||||
num_inference_steps = forward_batch.num_inference_steps
|
||||
|
||||
# initialize the coefficients, cutoff_steps, and ret_steps
|
||||
coefficients = teacache_params.coefficients
|
||||
use_ret_steps = teacache_params.use_ret_steps
|
||||
cutoff_steps = teacache_params.get_cutoff_steps(num_inference_steps)
|
||||
ret_steps = teacache_params.ret_steps
|
||||
teacache_thresh = teacache_params.teacache_thresh
|
||||
|
||||
if current_timestep == 0:
|
||||
self.cnt = 0
|
||||
|
||||
timestep_proj = kwargs["timestep_proj"]
|
||||
temb = kwargs["temb"]
|
||||
modulated_inp = timestep_proj if use_ret_steps else temb
|
||||
|
||||
if self.cnt % 2 == 0: # even -> condition
|
||||
self.is_even = True
|
||||
if self.cnt < ret_steps or self.cnt >= cutoff_steps:
|
||||
self.should_calc_even = True
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
else:
|
||||
assert self.previous_e0_even is not None, "previous_e0_even is not initialized"
|
||||
assert self.accumulated_rel_l1_distance_even is not None, "accumulated_rel_l1_distance_even is not initialized"
|
||||
rescale_func = np.poly1d(coefficients)
|
||||
self.accumulated_rel_l1_distance_even += rescale_func(
|
||||
((modulated_inp - self.previous_e0_even).abs().mean() /
|
||||
self.previous_e0_even.abs().mean()).cpu().item())
|
||||
if self.accumulated_rel_l1_distance_even < teacache_thresh:
|
||||
self.should_calc_even = False
|
||||
else:
|
||||
self.should_calc_even = True
|
||||
self.accumulated_rel_l1_distance_even = 0
|
||||
self.previous_e0_even = modulated_inp.clone()
|
||||
|
||||
else: # odd -> unconditon
|
||||
self.is_even = False
|
||||
if self.cnt < ret_steps or self.cnt >= cutoff_steps:
|
||||
self.should_calc_odd = True
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
else:
|
||||
assert self.previous_e0_odd is not None, "previous_e0_odd is not initialized"
|
||||
assert self.accumulated_rel_l1_distance_odd is not None, "accumulated_rel_l1_distance_odd is not initialized"
|
||||
rescale_func = np.poly1d(coefficients)
|
||||
self.accumulated_rel_l1_distance_odd += rescale_func(
|
||||
((modulated_inp - self.previous_e0_odd).abs().mean() /
|
||||
self.previous_e0_odd.abs().mean()).cpu().item())
|
||||
if self.accumulated_rel_l1_distance_odd < teacache_thresh:
|
||||
self.should_calc_odd = False
|
||||
else:
|
||||
self.should_calc_odd = True
|
||||
self.accumulated_rel_l1_distance_odd = 0
|
||||
self.previous_e0_odd = modulated_inp.clone()
|
||||
self.cnt += 1
|
||||
should_skip_forward = False
|
||||
if self.is_even:
|
||||
if not self.should_calc_even:
|
||||
should_skip_forward = True
|
||||
else:
|
||||
if not self.should_calc_odd:
|
||||
should_skip_forward = True
|
||||
|
||||
return should_skip_forward
|
||||
|
||||
def retrieve_cached_states(self,
|
||||
hidden_states: torch.Tensor) -> torch.Tensor:
|
||||
if self.is_even:
|
||||
return hidden_states + self.previous_residual_even
|
||||
else:
|
||||
return hidden_states + self.previous_residual_odd
|
||||
|
||||
# Entry point for model registry
|
||||
EntryClass = LingBotWorldTransformer3DModel
|
||||
@@ -814,6 +814,8 @@ class TransformerArgsPreprocessor:
|
||||
batch_size = x.shape[0]
|
||||
if context.device != x.device:
|
||||
context = context.to(x.device)
|
||||
if context.dtype != x.dtype:
|
||||
context = context.to(x.dtype)
|
||||
if attention_mask is not None and attention_mask.device != x.device:
|
||||
attention_mask = attention_mask.to(x.device)
|
||||
context = self.caption_projection(context)
|
||||
@@ -1476,6 +1478,26 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
|
||||
self.norm_eps = norm_eps
|
||||
|
||||
def _register_fsdp_backward_hooks_on_output(self, vx, ax):
|
||||
"""Register backward hooks on output tensors to trigger FSDP2 unshard.
|
||||
|
||||
FSDP2's module-level backward hooks don't fire when the module returns
|
||||
dataclass outputs. We must register hooks directly on the output tensors.
|
||||
"""
|
||||
if not hasattr(self, 'unshard'):
|
||||
return # Not wrapped by FSDP2
|
||||
|
||||
def make_unshard_hook():
|
||||
def hook(grad):
|
||||
self.unshard()
|
||||
return grad
|
||||
return hook
|
||||
|
||||
if vx is not None and vx.requires_grad:
|
||||
vx.register_hook(make_unshard_hook())
|
||||
if ax is not None and ax.requires_grad:
|
||||
ax.register_hook(make_unshard_hook())
|
||||
|
||||
def get_ada_values(
|
||||
self, scale_shift_table: torch.Tensor, batch_size: int, timestep: torch.Tensor, indices: slice
|
||||
) -> tuple[torch.Tensor, ...]:
|
||||
@@ -1704,6 +1726,10 @@ class BasicAVTransformerBlock(torch.nn.Module):
|
||||
f"audio_sum={audio_sum:.6f}"
|
||||
)
|
||||
|
||||
# Register FSDP2 backward hooks on output tensors (module-level hooks don't
|
||||
# fire for dataclass outputs, so we must hook the tensors directly)
|
||||
self._register_fsdp_backward_hooks_on_output(vx, ax)
|
||||
|
||||
return (
|
||||
replace(video, x=vx) if video is not None else None,
|
||||
replace(audio, x=ax) if audio is not None else None,
|
||||
@@ -2075,6 +2101,7 @@ class LTX2Transformer3DModel(CachableDiT):
|
||||
param_names_mapping = LTX2VideoConfig().param_names_mapping
|
||||
reverse_param_names_mapping = LTX2VideoConfig().reverse_param_names_mapping
|
||||
lora_param_names_mapping = LTX2VideoConfig().lora_param_names_mapping
|
||||
_fsdp_shard_conditions = LTX2VideoConfig()._fsdp_shard_conditions
|
||||
|
||||
def __init__(self, config: LTX2VideoConfig, hf_config: dict[str, Any]):
|
||||
super().__init__(config=config, hf_config=hf_config)
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from diffusers import SD3Transformer2DModel as _DiffusersSD3Transformer2DModel
|
||||
|
||||
from fastvideo.configs.models import DiTConfig
|
||||
|
||||
|
||||
class SD3Transformer2DModel(_DiffusersSD3Transformer2DModel):
|
||||
|
||||
_fsdp_shard_conditions: list = []
|
||||
_compile_conditions: list = []
|
||||
param_names_mapping: dict[str, Any] = {}
|
||||
reverse_param_names_mapping: dict[str, Any] = {}
|
||||
lora_param_names_mapping: dict[str, Any] = {}
|
||||
|
||||
def __init__(self, config: DiTConfig, hf_config: dict[str, Any], **kwargs):
|
||||
self.fastvideo_config = config
|
||||
self.hf_config = hf_config
|
||||
|
||||
arch = config.arch_config
|
||||
dual_layers = getattr(arch, "dual_attention_layers", ())
|
||||
if isinstance(dual_layers, list):
|
||||
dual_layers = tuple(dual_layers)
|
||||
|
||||
super().__init__(
|
||||
sample_size=arch.sample_size,
|
||||
patch_size=arch.patch_size,
|
||||
in_channels=arch.in_channels,
|
||||
num_layers=arch.num_layers,
|
||||
attention_head_dim=arch.attention_head_dim,
|
||||
num_attention_heads=arch.num_attention_heads,
|
||||
joint_attention_dim=arch.joint_attention_dim,
|
||||
caption_projection_dim=arch.caption_projection_dim,
|
||||
pooled_projection_dim=arch.pooled_projection_dim,
|
||||
out_channels=arch.out_channels,
|
||||
pos_embed_max_size=arch.pos_embed_max_size,
|
||||
dual_attention_layers=dual_layers,
|
||||
qk_norm=arch.qk_norm,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -491,6 +491,43 @@ class CLIPTextModel(TextEncoder):
|
||||
return loaded_params
|
||||
|
||||
|
||||
class CLIPTextModelWithProjection(CLIPTextModel):
|
||||
|
||||
def __init__(self, config: CLIPTextConfig) -> None:
|
||||
super().__init__(config)
|
||||
self.text_projection = nn.Linear(
|
||||
config.hidden_size, config.projection_dim, bias=False
|
||||
)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
outputs = super().forward(
|
||||
input_ids=input_ids,
|
||||
position_ids=position_ids,
|
||||
attention_mask=attention_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
output_hidden_states=output_hidden_states,
|
||||
**kwargs,
|
||||
)
|
||||
pooled = outputs.pooler_output
|
||||
if pooled is not None:
|
||||
pooled = self.text_projection(pooled)
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=outputs.last_hidden_state,
|
||||
pooler_output=pooled,
|
||||
hidden_states=outputs.hidden_states,
|
||||
attentions=outputs.attentions,
|
||||
attention_mask=outputs.attention_mask,
|
||||
)
|
||||
|
||||
|
||||
class CLIPVisionTransformer(nn.Module):
|
||||
|
||||
def __init__(
|
||||
|
||||
@@ -3,11 +3,11 @@ from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
from typing import Iterable
|
||||
from typing import Any, Iterable
|
||||
|
||||
import torch
|
||||
from torch import nn
|
||||
from transformers import Gemma3ForConditionalGeneration
|
||||
from transformers import AutoTokenizer, Gemma3ForConditionalGeneration
|
||||
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, TextEncoderConfig
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
@@ -447,6 +447,60 @@ class LTX2GemmaTextEncoderModel(TextEncoder):
|
||||
|
||||
return encoded, encoded_for_audio, attention_mask.squeeze(-1)
|
||||
|
||||
@torch.no_grad()
|
||||
def preprocess_text_embeddings(
|
||||
self,
|
||||
prompts: str | list[str],
|
||||
tokenizer: AutoTokenizer,
|
||||
tokenizer_kwargs: dict[str, Any] | None = None,
|
||||
padding_side: str | None = None,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
"""Compute pre-connector text embeddings for LTX-2 training preprocessing."""
|
||||
if isinstance(prompts, str):
|
||||
prompts = [prompts]
|
||||
|
||||
model = self.gemma_model
|
||||
kwargs: dict[str, Any] = {
|
||||
"padding": "max_length",
|
||||
"truncation": True,
|
||||
"return_tensors": "pt",
|
||||
}
|
||||
if tokenizer_kwargs is not None:
|
||||
kwargs.update(tokenizer_kwargs)
|
||||
if "max_length" not in kwargs:
|
||||
kwargs["max_length"] = self.config.arch_config.text_len
|
||||
|
||||
original_padding_side = tokenizer.padding_side
|
||||
target_padding_side = padding_side or self.padding_side
|
||||
tokenizer.padding_side = target_padding_side
|
||||
try:
|
||||
text_inputs = tokenizer(prompts, **kwargs)
|
||||
finally:
|
||||
tokenizer.padding_side = original_padding_side
|
||||
|
||||
input_ids = text_inputs["input_ids"].to(device=model.device)
|
||||
attention_mask = text_inputs["attention_mask"].to(device=model.device)
|
||||
outputs = model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
return_dict=True,
|
||||
)
|
||||
prompt_embeds = self._run_feature_extractor(
|
||||
outputs.hidden_states,
|
||||
attention_mask,
|
||||
padding_side=target_padding_side,
|
||||
)
|
||||
return prompt_embeds, attention_mask
|
||||
|
||||
def run_connectors(
|
||||
self,
|
||||
encoded_input: torch.Tensor,
|
||||
attention_mask: torch.Tensor,
|
||||
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
|
||||
"""Apply embedding connectors to precomputed Gemma features."""
|
||||
return self._run_connectors(encoded_input, attention_mask)
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
|
||||
@@ -0,0 +1,75 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
|
||||
from fastvideo.configs.models.encoders import BaseEncoderOutput, T5Config
|
||||
from fastvideo.models.encoders.base import TextEncoder
|
||||
|
||||
|
||||
class T5EncoderModel(TextEncoder):
|
||||
|
||||
supports_hf_from_pretrained: bool = True
|
||||
|
||||
def __init__(self, config: T5Config, hf_model: Any | None = None) -> None:
|
||||
super().__init__(config)
|
||||
self.hf_model = hf_model
|
||||
|
||||
@classmethod
|
||||
def from_pretrained_local(
|
||||
cls,
|
||||
model_path: str,
|
||||
config: T5Config,
|
||||
*,
|
||||
dtype: torch.dtype,
|
||||
device: torch.device,
|
||||
) -> "T5EncoderModel":
|
||||
from transformers import T5EncoderModel as HFT5EncoderModel
|
||||
|
||||
kwargs = {
|
||||
"local_files_only": True,
|
||||
"low_cpu_mem_usage": True,
|
||||
}
|
||||
try:
|
||||
hf = HFT5EncoderModel.from_pretrained(model_path, dtype=dtype, **kwargs)
|
||||
except TypeError:
|
||||
# Backward-compatible fallback for older Transformers versions.
|
||||
hf = HFT5EncoderModel.from_pretrained(
|
||||
model_path, torch_dtype=dtype, **kwargs
|
||||
)
|
||||
|
||||
hf = hf.eval().to(device=device, dtype=dtype)
|
||||
return cls(config=config, hf_model=hf).eval()
|
||||
|
||||
def forward(
|
||||
self,
|
||||
input_ids: torch.Tensor | None,
|
||||
position_ids: torch.Tensor | None = None,
|
||||
attention_mask: torch.Tensor | None = None,
|
||||
inputs_embeds: torch.Tensor | None = None,
|
||||
output_hidden_states: bool | None = None,
|
||||
**kwargs,
|
||||
) -> BaseEncoderOutput:
|
||||
if self.hf_model is None:
|
||||
raise RuntimeError(
|
||||
"T5EncoderModel(HF) is not initialized. Use "
|
||||
"`from_pretrained_local(...)` to construct a loaded instance."
|
||||
)
|
||||
|
||||
out = self.hf_model(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
inputs_embeds=inputs_embeds,
|
||||
output_hidden_states=bool(output_hidden_states)
|
||||
if output_hidden_states is not None
|
||||
else False,
|
||||
return_dict=True,
|
||||
)
|
||||
return BaseEncoderOutput(
|
||||
last_hidden_state=out.last_hidden_state,
|
||||
hidden_states=out.hidden_states if output_hidden_states else None,
|
||||
attention_mask=attention_mask,
|
||||
)
|
||||
@@ -86,8 +86,10 @@ class ComponentLoader(ABC):
|
||||
"vocoder": (VocoderLoader, "diffusers"),
|
||||
"text_encoder": (TextEncoderLoader, "transformers"),
|
||||
"text_encoder_2": (TextEncoderLoader, "transformers"),
|
||||
"text_encoder_3": (TextEncoderLoader, "transformers"),
|
||||
"tokenizer": (TokenizerLoader, "transformers"),
|
||||
"tokenizer_2": (TokenizerLoader, "transformers"),
|
||||
"tokenizer_3": (TokenizerLoader, "transformers"),
|
||||
"image_processor": (ImageProcessorLoader, "transformers"),
|
||||
"feature_extractor": (ImageProcessorLoader, "transformers"),
|
||||
"image_encoder": (ImageEncoderLoader, "transformers"),
|
||||
@@ -292,23 +294,26 @@ class TextEncoderLoader(ComponentLoader):
|
||||
pass
|
||||
logger.info("HF Model config: %s", model_config)
|
||||
|
||||
# @TODO(Wei): Better way to handle this?
|
||||
try:
|
||||
encoder_config = (
|
||||
fastvideo_args.pipeline_config.text_encoder_configs[0]
|
||||
base = os.path.basename(os.path.normpath(model_path))
|
||||
idx = 0
|
||||
if base.startswith("text_encoder_"):
|
||||
try:
|
||||
idx = int(base.split("_")[-1]) - 1
|
||||
except Exception:
|
||||
idx = 0
|
||||
encoder_configs = fastvideo_args.pipeline_config.text_encoder_configs
|
||||
encoder_precisions = fastvideo_args.pipeline_config.text_encoder_precisions
|
||||
if idx < 0 or idx >= len(encoder_configs):
|
||||
raise IndexError(
|
||||
f"text encoder index {idx} out of range for text_encoder_configs (len={len(encoder_configs)}), model_path={model_path}"
|
||||
)
|
||||
encoder_config.update_model_arch(model_config)
|
||||
encoder_precision = (
|
||||
fastvideo_args.pipeline_config.text_encoder_precisions[0]
|
||||
)
|
||||
except Exception:
|
||||
encoder_config = (
|
||||
fastvideo_args.pipeline_config.text_encoder_configs[1]
|
||||
)
|
||||
encoder_config.update_model_arch(model_config)
|
||||
encoder_precision = (
|
||||
fastvideo_args.pipeline_config.text_encoder_precisions[1]
|
||||
encoder_config = encoder_configs[idx]
|
||||
encoder_config.update_model_arch(model_config)
|
||||
if idx < 0 or idx >= len(encoder_precisions):
|
||||
raise IndexError(
|
||||
f"text encoder index {idx} out of range for text_encoder_precisions (len={len(encoder_precisions)}), model_path={model_path}"
|
||||
)
|
||||
encoder_precision = encoder_precisions[idx]
|
||||
|
||||
target_device = get_local_torch_device()
|
||||
# TODO(will): add support for other dtypes
|
||||
@@ -362,7 +367,16 @@ class TextEncoderLoader(ComponentLoader):
|
||||
with target_device:
|
||||
architectures = getattr(model_config, "architectures", [])
|
||||
model_cls, _ = ModelRegistry.resolve_model_cls(architectures)
|
||||
model: TextEncoder = model_cls(model_config) # type: ignore
|
||||
if getattr(model_cls, "supports_hf_from_pretrained", False):
|
||||
model = model_cls.from_pretrained_local( # type: ignore[attr-defined]
|
||||
model_path,
|
||||
model_config, # type: ignore[arg-type]
|
||||
dtype=PRECISION_TO_TYPE[dtype],
|
||||
device=target_device,
|
||||
)
|
||||
return model.eval()
|
||||
|
||||
model = model_cls(model_config) # type: ignore
|
||||
|
||||
weights_to_load = {name for name, _ in model.named_parameters()}
|
||||
if (
|
||||
@@ -659,7 +673,10 @@ class VAELoader(ComponentLoader):
|
||||
break
|
||||
loaded = remapped
|
||||
|
||||
vae.load_state_dict(loaded, strict=False)
|
||||
# Diffusers-format AutoencoderKL checkpoints should match exactly; load
|
||||
# strictly so missing/unexpected keys are surfaced early.
|
||||
strict_load = class_name == "AutoencoderKL"
|
||||
vae.load_state_dict(loaded, strict=strict_load)
|
||||
|
||||
return vae.eval()
|
||||
|
||||
@@ -1005,4 +1022,4 @@ class PipelineComponentLoader:
|
||||
)
|
||||
|
||||
# Load the module
|
||||
return loader.load(component_model_path, fastvideo_args)
|
||||
return loader.load(component_model_path, fastvideo_args)
|
||||
|
||||
@@ -37,6 +37,8 @@ _TEXT_TO_VIDEO_DIT_MODELS = {
|
||||
"LongCatVideoTransformer3DModel": ("dits", "longcat_video_dit", "LongCatVideoTransformer3DModel"), # Wrapper (Phase 1)
|
||||
"LongCatTransformer3DModel": ("dits", "longcat", "LongCatTransformer3DModel"), # Native (Phase 2)
|
||||
"LTX2Transformer3DModel": ("dits", "ltx2", "LTX2Transformer3DModel"),
|
||||
"SD3Transformer2DModel": ("dits", "sd3", "SD3Transformer2DModel"),
|
||||
"LingBotWorldTransformer3DModel": ("dits", "lingbotworld", "LingBotWorldTransformer3DModel"),
|
||||
}
|
||||
|
||||
_IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
@@ -49,9 +51,11 @@ _IMAGE_TO_VIDEO_DIT_MODELS = {
|
||||
|
||||
_TEXT_ENCODER_MODELS = {
|
||||
"CLIPTextModel": ("encoders", "clip", "CLIPTextModel"),
|
||||
"CLIPTextModelWithProjection":
|
||||
("encoders", "clip", "CLIPTextModelWithProjection"),
|
||||
"LlamaModel": ("encoders", "llama", "LlamaModel"),
|
||||
"UMT5EncoderModel": ("encoders", "t5", "UMT5EncoderModel"),
|
||||
"T5EncoderModel": ("encoders", "t5", "T5EncoderModel"),
|
||||
"T5EncoderModel": ("encoders", "t5_hf", "T5EncoderModel"),
|
||||
"STEP1TextEncoder": ("encoders", "stepllm", "STEP1TextEncoder"),
|
||||
"BertModel": ("encoders", "clip", "CLIPTextModel"),
|
||||
"Qwen2_5_VLTextModel": ("encoders", "qwen2_5", "Qwen2_5_VLTextModel"),
|
||||
@@ -75,6 +79,7 @@ _VAE_MODELS = {
|
||||
"AutoencoderKLHunyuanVideo15": ("vaes", "hunyuan15vae", "AutoencoderKLHunyuanVideo15"),
|
||||
"AutoencoderKLWan": ("vaes", "wanvae", "AutoencoderKLWan"),
|
||||
"AutoencoderKLStepvideo": ("vaes", "stepvideovae", "AutoencoderKLStepvideo"),
|
||||
"AutoencoderKL": ("vaes", "autoencoder_kl", "AutoencoderKL"),
|
||||
"CausalVideoAutoencoder": ("vaes", "ltx2vae", "LTX2CausalVideoAutoencoder"),
|
||||
}
|
||||
|
||||
|
||||
@@ -0,0 +1,55 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from diffusers import AutoencoderKL as _DiffusersAutoencoderKL
|
||||
|
||||
from fastvideo.configs.models import VAEConfig
|
||||
|
||||
|
||||
class AutoencoderKL(_DiffusersAutoencoderKL):
|
||||
|
||||
def __init__(self, config: VAEConfig, **kwargs: Any) -> None:
|
||||
self.fastvideo_config = config
|
||||
arch = config.arch_config
|
||||
|
||||
down_block_types = arch.down_block_types
|
||||
if isinstance(down_block_types, list):
|
||||
down_block_types = tuple(down_block_types)
|
||||
up_block_types = arch.up_block_types
|
||||
if isinstance(up_block_types, list):
|
||||
up_block_types = tuple(up_block_types)
|
||||
block_out_channels = arch.block_out_channels
|
||||
if isinstance(block_out_channels, list):
|
||||
block_out_channels = tuple(block_out_channels)
|
||||
|
||||
latents_mean = arch.latents_mean
|
||||
if isinstance(latents_mean, list):
|
||||
latents_mean = tuple(latents_mean)
|
||||
latents_std = arch.latents_std
|
||||
if isinstance(latents_std, list):
|
||||
latents_std = tuple(latents_std)
|
||||
|
||||
super().__init__(
|
||||
in_channels=arch.in_channels,
|
||||
out_channels=arch.out_channels,
|
||||
down_block_types=down_block_types,
|
||||
up_block_types=up_block_types,
|
||||
block_out_channels=block_out_channels,
|
||||
layers_per_block=arch.layers_per_block,
|
||||
act_fn=arch.act_fn,
|
||||
latent_channels=arch.latent_channels,
|
||||
norm_num_groups=arch.norm_num_groups,
|
||||
sample_size=arch.sample_size,
|
||||
scaling_factor=arch.scaling_factor,
|
||||
shift_factor=arch.shift_factor,
|
||||
latents_mean=latents_mean,
|
||||
latents_std=latents_std,
|
||||
force_upcast=arch.force_upcast,
|
||||
use_quant_conv=arch.use_quant_conv,
|
||||
use_post_quant_conv=arch.use_post_quant_conv,
|
||||
mid_block_add_attention=arch.mid_block_add_attention,
|
||||
**kwargs,
|
||||
)
|
||||
@@ -0,0 +1 @@
|
||||
|
||||
@@ -0,0 +1,16 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""
|
||||
Wan video diffusion pipeline implementation.
|
||||
|
||||
This module contains an implementation of the Wan video diffusion pipeline
|
||||
using the modular pipeline architecture.
|
||||
"""
|
||||
|
||||
from fastvideo.pipelines.basic.wan.wan_i2v_pipeline import WanImageToVideoPipeline
|
||||
|
||||
|
||||
class LingBotWorldImageToVideoPipeline(WanImageToVideoPipeline):
|
||||
pass
|
||||
|
||||
|
||||
EntryClass = LingBotWorldImageToVideoPipeline
|
||||
@@ -0,0 +1,118 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.stages.input_validation import InputValidationStage
|
||||
from fastvideo.pipelines.stages.text_encoding import TextEncodingStage
|
||||
from fastvideo.pipelines.stages.timestep_preparation import (
|
||||
TimestepPreparationStage, )
|
||||
from fastvideo.pipelines.stages.sd35_conditioning import (
|
||||
SD35ConditioningStage,
|
||||
SD35DecodingStage,
|
||||
SD35DenoisingStage,
|
||||
SD35LatentPreparationStage,
|
||||
)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class SD35Pipeline(ComposedPipelineBase):
|
||||
"""Minimal SD3.5 Medium text-to-image pipeline (treat as num_frames=1)."""
|
||||
|
||||
_required_config_modules = [
|
||||
"scheduler",
|
||||
"transformer",
|
||||
"vae",
|
||||
"text_encoder",
|
||||
"text_encoder_2",
|
||||
"text_encoder_3",
|
||||
"tokenizer",
|
||||
"tokenizer_2",
|
||||
"tokenizer_3",
|
||||
]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
te_cfgs = list(fastvideo_args.pipeline_config.text_encoder_configs)
|
||||
if len(te_cfgs) >= 2:
|
||||
for i in (0, 1):
|
||||
te_cfgs[i].tokenizer_kwargs.setdefault("padding", "max_length")
|
||||
te_cfgs[i].tokenizer_kwargs.setdefault("max_length", 77)
|
||||
te_cfgs[i].tokenizer_kwargs.setdefault("truncation", True)
|
||||
te_cfgs[i].tokenizer_kwargs.setdefault("return_tensors", "pt")
|
||||
if len(te_cfgs) >= 3:
|
||||
te_cfgs[2].tokenizer_kwargs["max_length"] = min(
|
||||
int(te_cfgs[2].tokenizer_kwargs.get("max_length", 256)), 256)
|
||||
te_cfgs[2].tokenizer_kwargs.setdefault("padding", "max_length")
|
||||
te_cfgs[2].tokenizer_kwargs.setdefault("truncation", True)
|
||||
te_cfgs[2].tokenizer_kwargs.setdefault("return_tensors", "pt")
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs) -> None:
|
||||
self.add_stage(stage_name="input_validation_stage",
|
||||
stage=InputValidationStage())
|
||||
|
||||
self.add_stage(
|
||||
stage_name="text_encoding_stage",
|
||||
stage=TextEncodingStage(
|
||||
text_encoders=[
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder_2"),
|
||||
self.get_module("text_encoder_3"),
|
||||
],
|
||||
tokenizers=[
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer_2"),
|
||||
self.get_module("tokenizer_3"),
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="timestep_preparation_stage",
|
||||
stage=TimestepPreparationStage(
|
||||
scheduler=self.get_module("scheduler")),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="latent_preparation_stage",
|
||||
stage=SD35LatentPreparationStage(
|
||||
scheduler=self.get_module("scheduler"), ),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="sd35_conditioning_stage",
|
||||
stage=SD35ConditioningStage(
|
||||
text_encoders=[
|
||||
self.get_module("text_encoder"),
|
||||
self.get_module("text_encoder_2"),
|
||||
self.get_module("text_encoder_3"),
|
||||
],
|
||||
tokenizers=[
|
||||
self.get_module("tokenizer"),
|
||||
self.get_module("tokenizer_2"),
|
||||
self.get_module("tokenizer_3"),
|
||||
],
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="denoising_stage",
|
||||
stage=SD35DenoisingStage(
|
||||
transformer=self.get_module("transformer"),
|
||||
scheduler=self.get_module("scheduler"),
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="decoding_stage",
|
||||
stage=SD35DecodingStage(vae=self.get_module("vae"), ),
|
||||
)
|
||||
|
||||
|
||||
class StableDiffusion3Pipeline(SD35Pipeline):
|
||||
"""Alias name to match SD3.5 diffusers `model_index.json` _class_name."""
|
||||
|
||||
|
||||
EntryClass = [SD35Pipeline, StableDiffusion3Pipeline]
|
||||
@@ -136,6 +136,9 @@ class ForwardBatch:
|
||||
# Camera control inputs (HYWorld)
|
||||
pose: str | None = None # Camera trajectory: pose string (e.g., 'w-31') or JSON file path
|
||||
|
||||
# Camera control inputs (LingBotWorld)
|
||||
c2ws_plucker_emb: torch.Tensor | None = None # Plucker embedding: [B, C, F_lat, H_lat, W_lat]
|
||||
|
||||
# Latent dimensions
|
||||
height_latents: list[int] | int | None = None
|
||||
width_latents: list[int] | int | None = None
|
||||
@@ -230,6 +233,14 @@ class TrainingBatch:
|
||||
noise_latents: torch.Tensor | None = None
|
||||
encoder_hidden_states: torch.Tensor | None = None
|
||||
encoder_attention_mask: torch.Tensor | None = None
|
||||
# LTX related audio inputs
|
||||
audio_latents: torch.Tensor | None = None
|
||||
audio_noisy_model_input: torch.Tensor | None = None
|
||||
audio_timesteps: torch.Tensor | None = None
|
||||
audio_noise: torch.Tensor | None = None
|
||||
audio_encoder_hidden_states: torch.Tensor | None = None
|
||||
audio_encoder_attention_mask: torch.Tensor | None = None
|
||||
conditioning_mask: torch.Tensor | None = None
|
||||
# i2v
|
||||
preprocessed_image: torch.Tensor | None = None
|
||||
image_embeds: torch.Tensor | None = None
|
||||
|
||||
@@ -0,0 +1,290 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
"""LTX-2 preprocessing pipeline for native FastVideo training data generation.
|
||||
|
||||
This module defines the LTX-2 preprocess pipeline used by FastVideo workflows
|
||||
to build precomputed training artifacts from raw text/video datasets.
|
||||
|
||||
Usage:
|
||||
- Entry is through preprocess workflows that register `PreprocessPipelineT2V`.
|
||||
- Input samples should provide prompt text plus video metadata/loader fields
|
||||
consumed by `TextTransformStage` and `VideoTransformStage`.
|
||||
- Output artifacts are written by the shared preprocessing workflow into
|
||||
`.precomputed/` (latents, conditions, and optional audio_latents).
|
||||
|
||||
Optional audio path:
|
||||
- When audio preprocessing is enabled, this pipeline loads the native
|
||||
LTX-2 audio encoder and stores per-sample audio latents in
|
||||
`batch.extra["ltx2_audio_latents"]`.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import glob
|
||||
import os
|
||||
from typing import Any, cast
|
||||
|
||||
import torch
|
||||
import torchaudio
|
||||
from safetensors.torch import load_file as safetensors_load_file
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.models.audio.ltx2_audio_processing import AudioProcessor
|
||||
from fastvideo.models.audio.ltx2_audio_vae import LTX2AudioEncoder
|
||||
from fastvideo.models.hf_transformer_utils import get_diffusers_config
|
||||
from fastvideo.pipelines.composed_pipeline_base import ComposedPipelineBase
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, PreprocessBatch
|
||||
from fastvideo.pipelines.preprocess.preprocess_stages import (
|
||||
TextTransformStage, VideoTransformStage)
|
||||
from fastvideo.pipelines.stages import EncodingStage, PipelineStage
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
class LTX2TextPrecomputeStage(PipelineStage):
|
||||
"""Compute pre-connector Gemma embeddings for LTX-2 training."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
text_encoder: torch.nn.Module,
|
||||
tokenizer: Any,
|
||||
preprocess_text_fn,
|
||||
tokenizer_kwargs: dict[str, Any],
|
||||
padding_side: str,
|
||||
) -> None:
|
||||
self.text_encoder = text_encoder
|
||||
self.tokenizer = tokenizer
|
||||
self.preprocess_text_fn = preprocess_text_fn
|
||||
self.tokenizer_kwargs = tokenizer_kwargs
|
||||
self.padding_side = padding_side
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
batch = cast(PreprocessBatch, batch)
|
||||
assert isinstance(batch.prompt, list)
|
||||
|
||||
prompts = []
|
||||
for prompt in batch.prompt:
|
||||
if not isinstance(prompt, str):
|
||||
prompt = str(prompt)
|
||||
processed_prompt = self.preprocess_text_fn(prompt)
|
||||
prompts.append(
|
||||
processed_prompt if processed_prompt is not None else "")
|
||||
|
||||
prompt_embeds, prompt_attention_mask = (
|
||||
self.text_encoder.preprocess_text_embeddings(
|
||||
prompts=prompts,
|
||||
tokenizer=self.tokenizer,
|
||||
tokenizer_kwargs=self.tokenizer_kwargs,
|
||||
padding_side=self.padding_side,
|
||||
))
|
||||
batch.prompt_embeds = [prompt_embeds]
|
||||
batch.prompt_attention_mask = [prompt_attention_mask]
|
||||
return batch
|
||||
|
||||
|
||||
class LTX2AudioEncodingStage(PipelineStage):
|
||||
"""Extract audio from input videos and encode into LTX-2 audio latents."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
audio_encoder: torch.nn.Module,
|
||||
audio_processor: AudioProcessor,
|
||||
fallback_fps: int,
|
||||
) -> None:
|
||||
self.audio_encoder = audio_encoder.eval()
|
||||
self.audio_processor = audio_processor
|
||||
self.fallback_fps = fallback_fps
|
||||
self.audio_dtype = next(audio_encoder.parameters()).dtype
|
||||
self.audio_device = next(audio_encoder.parameters()).device
|
||||
|
||||
@staticmethod
|
||||
def _extract_audio(
|
||||
video_path: str,
|
||||
target_duration: float,
|
||||
) -> tuple[torch.Tensor, int] | None:
|
||||
try:
|
||||
waveform, sample_rate = torchaudio.load(video_path)
|
||||
except Exception as e:
|
||||
logger.error("Failed to load audio from %s: %s", video_path, e)
|
||||
raise e
|
||||
|
||||
target_samples = int(target_duration * sample_rate)
|
||||
if target_samples <= 0:
|
||||
return None
|
||||
|
||||
current_samples = waveform.shape[-1]
|
||||
if current_samples > target_samples:
|
||||
waveform = waveform[..., :target_samples]
|
||||
elif current_samples < target_samples:
|
||||
padding = target_samples - current_samples
|
||||
waveform = torch.nn.functional.pad(waveform, (0, padding))
|
||||
|
||||
return waveform, sample_rate
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(
|
||||
self,
|
||||
batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs,
|
||||
) -> ForwardBatch:
|
||||
batch = cast(PreprocessBatch, batch)
|
||||
assert isinstance(batch.video_loader, list)
|
||||
assert isinstance(batch.num_frames, list)
|
||||
assert isinstance(batch.fps, list)
|
||||
|
||||
audio_latents: list[torch.Tensor | None] = []
|
||||
for idx, video_input in enumerate(batch.video_loader):
|
||||
if not isinstance(video_input, str):
|
||||
logger.warning(
|
||||
"Skipping audio for sample %s: video loader is not a path string",
|
||||
idx,
|
||||
)
|
||||
audio_latents.append(None)
|
||||
continue
|
||||
|
||||
fps = float(batch.fps[idx]) if batch.fps[idx] else float(
|
||||
self.fallback_fps)
|
||||
if fps <= 0:
|
||||
fps = float(self.fallback_fps)
|
||||
target_duration = float(batch.num_frames[idx]) / fps
|
||||
|
||||
audio_data = self._extract_audio(video_input, target_duration)
|
||||
if audio_data is None:
|
||||
audio_latents.append(None)
|
||||
continue
|
||||
|
||||
waveform, sample_rate = audio_data
|
||||
waveform = waveform.unsqueeze(0).to(device=self.audio_device,
|
||||
dtype=self.audio_dtype)
|
||||
mel = self.audio_processor.waveform_to_mel(
|
||||
waveform,
|
||||
waveform_sample_rate=sample_rate).to(device=self.audio_device,
|
||||
dtype=self.audio_dtype)
|
||||
latents = self.audio_encoder(mel).squeeze(0).detach().cpu()
|
||||
audio_latents.append(latents)
|
||||
|
||||
batch.extra["ltx2_audio_latents"] = audio_latents
|
||||
return batch
|
||||
|
||||
|
||||
class PreprocessPipelineT2V(ComposedPipelineBase):
|
||||
"""Native LTX-2 preprocessing pipeline (text/video with optional audio)."""
|
||||
|
||||
_required_config_modules = ["text_encoder", "tokenizer", "vae"]
|
||||
|
||||
def initialize_pipeline(self, fastvideo_args: FastVideoArgs):
|
||||
tokenizer = self.get_module("tokenizer")
|
||||
if tokenizer is not None:
|
||||
tokenizer.padding_side = "left"
|
||||
if tokenizer.pad_token is None and tokenizer.eos_token is not None:
|
||||
tokenizer.pad_token = tokenizer.eos_token
|
||||
|
||||
def _load_ltx2_audio_encoder(
|
||||
self) -> tuple[torch.nn.Module, AudioProcessor]:
|
||||
audio_vae_path = os.path.join(self.model_path, "audio_vae")
|
||||
if not os.path.isdir(audio_vae_path):
|
||||
raise FileNotFoundError(
|
||||
f"Expected audio_vae directory for LTX-2 audio preprocessing: {audio_vae_path}"
|
||||
)
|
||||
|
||||
config = get_diffusers_config(model=audio_vae_path)
|
||||
audio_encoder = LTX2AudioEncoder(config).to(
|
||||
device=get_local_torch_device(),
|
||||
dtype=torch.float32,
|
||||
)
|
||||
|
||||
safetensors_list = glob.glob(
|
||||
os.path.join(audio_vae_path, "*.safetensors"))
|
||||
loaded: dict[str, torch.Tensor] = {}
|
||||
for sf_file in safetensors_list:
|
||||
loaded.update(safetensors_load_file(sf_file))
|
||||
|
||||
encoder_state = {}
|
||||
for name, tensor in loaded.items():
|
||||
if name.startswith("encoder."):
|
||||
encoder_state[name.replace("encoder.", "")] = tensor
|
||||
elif name.startswith("per_channel_statistics."):
|
||||
encoder_state[name] = tensor
|
||||
|
||||
target_module = getattr(audio_encoder, "model", audio_encoder)
|
||||
missing, unexpected = target_module.load_state_dict(encoder_state,
|
||||
strict=False)
|
||||
if missing:
|
||||
logger.warning("Missing LTX-2 audio encoder keys: %s", missing[:8])
|
||||
if unexpected:
|
||||
logger.warning("Unexpected LTX-2 audio encoder keys: %s",
|
||||
unexpected[:8])
|
||||
target_module.eval()
|
||||
|
||||
audio_processor = AudioProcessor(
|
||||
sample_rate=target_module.sample_rate,
|
||||
mel_bins=target_module.mel_bins,
|
||||
mel_hop_length=target_module.mel_hop_length,
|
||||
n_fft=target_module.n_fft,
|
||||
).to(next(target_module.parameters()).device)
|
||||
return target_module, audio_processor
|
||||
|
||||
def create_pipeline_stages(self, fastvideo_args: FastVideoArgs):
|
||||
assert fastvideo_args.preprocess_config is not None
|
||||
|
||||
preprocess_cfg = fastvideo_args.preprocess_config
|
||||
self.add_stage(
|
||||
stage_name="text_transform_stage",
|
||||
stage=TextTransformStage(
|
||||
cfg_uncondition_drop_rate=preprocess_cfg.training_cfg_rate,
|
||||
seed=preprocess_cfg.seed,
|
||||
),
|
||||
)
|
||||
|
||||
text_encoder = self.get_module("text_encoder")
|
||||
tokenizer = self.get_module("tokenizer")
|
||||
encoder_config = fastvideo_args.pipeline_config.text_encoder_configs[0]
|
||||
tokenizer_kwargs = dict(encoder_config.tokenizer_kwargs)
|
||||
if "max_length" not in tokenizer_kwargs:
|
||||
tokenizer_kwargs["max_length"] = encoder_config.arch_config.text_len
|
||||
self.add_stage(
|
||||
stage_name="prompt_precompute_stage",
|
||||
stage=LTX2TextPrecomputeStage(
|
||||
text_encoder=text_encoder,
|
||||
tokenizer=tokenizer,
|
||||
preprocess_text_fn=fastvideo_args.pipeline_config.
|
||||
preprocess_text_funcs[0],
|
||||
tokenizer_kwargs=tokenizer_kwargs,
|
||||
padding_side=encoder_config.arch_config.padding_side,
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="video_transform_stage",
|
||||
stage=VideoTransformStage(
|
||||
train_fps=preprocess_cfg.train_fps,
|
||||
num_frames=preprocess_cfg.num_frames,
|
||||
max_height=preprocess_cfg.max_height,
|
||||
max_width=preprocess_cfg.max_width,
|
||||
do_temporal_sample=preprocess_cfg.do_temporal_sample,
|
||||
),
|
||||
)
|
||||
if preprocess_cfg.with_audio:
|
||||
audio_encoder, audio_processor = self._load_ltx2_audio_encoder()
|
||||
self.add_stage(
|
||||
stage_name="audio_encoding_stage",
|
||||
stage=LTX2AudioEncodingStage(
|
||||
audio_encoder=audio_encoder,
|
||||
audio_processor=audio_processor,
|
||||
fallback_fps=preprocess_cfg.train_fps,
|
||||
),
|
||||
)
|
||||
|
||||
self.add_stage(
|
||||
stage_name="video_encoding_stage",
|
||||
stage=EncodingStage(vae=self.get_module("vae")),
|
||||
)
|
||||
|
||||
|
||||
EntryClass = PreprocessPipelineT2V
|
||||
@@ -173,6 +173,7 @@ class DenoisingStage(PipelineStage):
|
||||
{
|
||||
"mouse_cond": batch.mouse_cond,
|
||||
"keyboard_cond": batch.keyboard_cond,
|
||||
"c2ws_plucker_emb": batch.c2ws_plucker_emb,
|
||||
},
|
||||
)
|
||||
|
||||
|
||||
@@ -23,7 +23,6 @@ from fastvideo.models.dits.ltx2 import (
|
||||
DEFAULT_LTX2_AUDIO_DOWNSAMPLE, DEFAULT_LTX2_AUDIO_HOP_LENGTH,
|
||||
DEFAULT_LTX2_AUDIO_MEL_BINS, DEFAULT_LTX2_AUDIO_SAMPLE_RATE,
|
||||
VideoLatentShape)
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
BASE_SHIFT_ANCHOR = 1024
|
||||
MAX_SHIFT_ANCHOR = 4096
|
||||
@@ -46,6 +45,7 @@ def _ltx2_sigmas(
|
||||
stretch: bool = True,
|
||||
terminal: float = 0.1,
|
||||
) -> torch.Tensor:
|
||||
# Copied/following official LTX-2 scheduler (LTX2Scheduler.execute).
|
||||
tokens = math.prod(
|
||||
latent.shape[2:]) if latent is not None else MAX_SHIFT_ANCHOR
|
||||
sigmas = torch.linspace(1.0,
|
||||
@@ -114,12 +114,9 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
if neg_prompt_mask is not None and neg_prompt_mask.device != latents.device:
|
||||
neg_prompt_mask = neg_prompt_mask.to(latents.device)
|
||||
|
||||
target_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.dit_precision]
|
||||
disable_autocast = os.getenv("LTX2_DISABLE_AUTOCAST", "1") == "1"
|
||||
target_dtype = torch.bfloat16
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast and (
|
||||
not disable_autocast)
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
# Use official distilled sigma schedule for 8 steps (distilled models)
|
||||
use_distilled_sigmas = os.getenv("LTX2_USE_DISTILLED_SIGMAS",
|
||||
@@ -137,6 +134,7 @@ class LTX2DenoisingStage(PipelineStage):
|
||||
latent=None,
|
||||
device=latents.device,
|
||||
)
|
||||
logger.info("[LTX2] Using computed sigma schedule")
|
||||
if hasattr(self.transformer, "patchifier"):
|
||||
video_shape = VideoLatentShape.from_torch_shape(latents.shape)
|
||||
token_count = self.transformer.patchifier.get_token_count(
|
||||
|
||||
@@ -0,0 +1,374 @@
|
||||
# SPDX-License-Identifier: Apache-2.0
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import inspect
|
||||
from typing import Any
|
||||
|
||||
import torch
|
||||
import torch.nn.functional as F
|
||||
from diffusers.utils.torch_utils import randn_tensor
|
||||
|
||||
from fastvideo.distributed import get_local_torch_device
|
||||
from fastvideo.fastvideo_args import FastVideoArgs
|
||||
from fastvideo.forward_context import set_forward_context
|
||||
from fastvideo.pipelines.pipeline_batch_info import ForwardBatch
|
||||
from fastvideo.pipelines.stages.base import PipelineStage
|
||||
from fastvideo.utils import PRECISION_TO_TYPE
|
||||
|
||||
|
||||
class SD35LatentPreparationStage(PipelineStage):
|
||||
|
||||
def __init__(self, scheduler) -> None:
|
||||
super().__init__()
|
||||
self.scheduler = scheduler
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.height is None or batch.width is None:
|
||||
raise ValueError(
|
||||
"height/width must be set for SD35LatentPreparationStage")
|
||||
|
||||
if isinstance(batch.prompt, list):
|
||||
batch_size = len(batch.prompt)
|
||||
elif batch.prompt is not None:
|
||||
batch_size = 1
|
||||
else:
|
||||
if not batch.prompt_embeds:
|
||||
raise ValueError("prompt or prompt_embeds must be provided")
|
||||
batch_size = batch.prompt_embeds[0].shape[0]
|
||||
|
||||
batch_size *= batch.num_videos_per_prompt
|
||||
|
||||
dtype = PRECISION_TO_TYPE[fastvideo_args.pipeline_config.dit_precision]
|
||||
device = get_local_torch_device()
|
||||
|
||||
if isinstance(batch.generator, list) and len(
|
||||
batch.generator) != batch_size:
|
||||
raise ValueError(
|
||||
f"generator list length {len(batch.generator)} does not match batch_size {batch_size}"
|
||||
)
|
||||
|
||||
in_channels = fastvideo_args.pipeline_config.dit_config.arch_config.in_channels
|
||||
spatial_ratio = fastvideo_args.pipeline_config.vae_config.arch_config.spatial_compression_ratio
|
||||
h_lat = batch.height // spatial_ratio
|
||||
w_lat = batch.width // spatial_ratio
|
||||
shape = (batch_size, in_channels, 1, h_lat, w_lat)
|
||||
|
||||
latents = batch.latents
|
||||
if latents is None:
|
||||
latents = randn_tensor(
|
||||
shape,
|
||||
generator=batch.generator,
|
||||
device=device,
|
||||
dtype=dtype,
|
||||
)
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
else:
|
||||
latents = latents.to(device=device, dtype=dtype)
|
||||
if hasattr(self.scheduler, "init_noise_sigma"):
|
||||
latents = latents * self.scheduler.init_noise_sigma
|
||||
|
||||
batch.latents = latents
|
||||
batch.raw_latent_shape = shape
|
||||
return batch
|
||||
|
||||
|
||||
class SD35ConditioningStage(PipelineStage):
|
||||
|
||||
def __init__(self, text_encoders, tokenizers) -> None:
|
||||
super().__init__()
|
||||
self.text_encoders = text_encoders
|
||||
self.tokenizers = tokenizers
|
||||
|
||||
@staticmethod
|
||||
def _tokenize(
|
||||
tokenizer: Any,
|
||||
prompts: str | list[str],
|
||||
tok_kwargs: dict[str, Any],
|
||||
device: torch.device,
|
||||
) -> tuple[torch.Tensor, torch.Tensor]:
|
||||
texts = [prompts] if isinstance(prompts, str) else prompts
|
||||
enc = tokenizer(texts, **tok_kwargs)
|
||||
input_ids = enc["input_ids"].to(device)
|
||||
attention_mask = enc["attention_mask"].to(device)
|
||||
return input_ids, attention_mask
|
||||
|
||||
@torch.no_grad()
|
||||
def _clip_pooled(
|
||||
self,
|
||||
text_encoder,
|
||||
tokenizer,
|
||||
prompts: str | list[str],
|
||||
tok_kwargs: dict[str, Any],
|
||||
device: torch.device,
|
||||
dtype: torch.dtype,
|
||||
) -> torch.Tensor:
|
||||
input_ids, attention_mask = self._tokenize(tokenizer, prompts,
|
||||
tok_kwargs, device)
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
out = text_encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=False,
|
||||
)
|
||||
pooled = getattr(out, "pooler_output", None)
|
||||
if pooled is None:
|
||||
raise RuntimeError(
|
||||
"CLIP pooled output is required for SD3.5 conditioning")
|
||||
return pooled.to(dtype=dtype)
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if len(batch.prompt_embeds) < 3:
|
||||
raise ValueError(
|
||||
f"SD35ConditioningStage expects 3 prompt_embeds entries (2x CLIP + 1x T5), got {len(batch.prompt_embeds)}"
|
||||
)
|
||||
|
||||
device = get_local_torch_device()
|
||||
target_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.dit_precision]
|
||||
|
||||
clip_1 = batch.prompt_embeds[0].to(device=device, dtype=target_dtype)
|
||||
clip_2 = batch.prompt_embeds[1].to(device=device, dtype=target_dtype)
|
||||
t5 = batch.prompt_embeds[2].to(device=device, dtype=target_dtype)
|
||||
|
||||
clip_prompt = torch.cat([clip_1, clip_2], dim=-1)
|
||||
if clip_prompt.shape[-1] > t5.shape[-1]:
|
||||
raise ValueError(
|
||||
f"CLIP prompt dim {clip_prompt.shape[-1]} exceeds T5 dim {t5.shape[-1]}"
|
||||
)
|
||||
clip_prompt = F.pad(clip_prompt,
|
||||
(0, t5.shape[-1] - clip_prompt.shape[-1]))
|
||||
prompt_embeds = torch.cat([clip_prompt, t5], dim=-2)
|
||||
|
||||
te_cfgs = fastvideo_args.pipeline_config.text_encoder_configs
|
||||
clip_tok_kwargs_1 = dict(getattr(te_cfgs[0], "tokenizer_kwargs", {}))
|
||||
clip_tok_kwargs_2 = dict(getattr(te_cfgs[1], "tokenizer_kwargs", {}))
|
||||
clip_tok_kwargs_1.setdefault("padding", "max_length")
|
||||
clip_tok_kwargs_2.setdefault("padding", "max_length")
|
||||
clip_tok_kwargs_1.setdefault("max_length", 77)
|
||||
clip_tok_kwargs_2.setdefault("max_length", 77)
|
||||
clip_tok_kwargs_1.setdefault("truncation", True)
|
||||
clip_tok_kwargs_2.setdefault("truncation", True)
|
||||
clip_tok_kwargs_1.setdefault("return_tensors", "pt")
|
||||
clip_tok_kwargs_2.setdefault("return_tensors", "pt")
|
||||
|
||||
pooled_1 = self._clip_pooled(self.text_encoders[0], self.tokenizers[0],
|
||||
batch.prompt, clip_tok_kwargs_1, device,
|
||||
target_dtype)
|
||||
pooled_2 = self._clip_pooled(self.text_encoders[1], self.tokenizers[1],
|
||||
batch.prompt, clip_tok_kwargs_2, device,
|
||||
target_dtype)
|
||||
pooled = torch.cat([pooled_1, pooled_2], dim=-1)
|
||||
|
||||
batch.extra["sd35_encoder_hidden_states"] = prompt_embeds
|
||||
batch.extra["sd35_pooled_projections"] = pooled
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
if batch.negative_prompt_embeds is None or len(
|
||||
batch.negative_prompt_embeds) < 3:
|
||||
raise ValueError(
|
||||
"negative_prompt_embeds must contain 3 entries when CFG is enabled"
|
||||
)
|
||||
|
||||
neg_clip_1 = batch.negative_prompt_embeds[0].to(device=device,
|
||||
dtype=target_dtype)
|
||||
neg_clip_2 = batch.negative_prompt_embeds[1].to(device=device,
|
||||
dtype=target_dtype)
|
||||
neg_t5 = batch.negative_prompt_embeds[2].to(device=device,
|
||||
dtype=target_dtype)
|
||||
|
||||
neg_clip_prompt = torch.cat([neg_clip_1, neg_clip_2], dim=-1)
|
||||
neg_clip_prompt = F.pad(
|
||||
neg_clip_prompt,
|
||||
(0, neg_t5.shape[-1] - neg_clip_prompt.shape[-1]))
|
||||
neg_prompt_embeds = torch.cat([neg_clip_prompt, neg_t5], dim=-2)
|
||||
|
||||
negative_pooled_1 = self._clip_pooled(
|
||||
self.text_encoders[0],
|
||||
self.tokenizers[0],
|
||||
batch.negative_prompt if isinstance(batch.negative_prompt, str
|
||||
| list) else "",
|
||||
clip_tok_kwargs_1,
|
||||
device,
|
||||
target_dtype,
|
||||
)
|
||||
negative_pooled_2 = self._clip_pooled(
|
||||
self.text_encoders[1],
|
||||
self.tokenizers[1],
|
||||
batch.negative_prompt if isinstance(batch.negative_prompt, str
|
||||
| list) else "",
|
||||
clip_tok_kwargs_2,
|
||||
device,
|
||||
target_dtype,
|
||||
)
|
||||
neg_pooled = torch.cat([negative_pooled_1, negative_pooled_2],
|
||||
dim=-1)
|
||||
|
||||
batch.extra[
|
||||
"sd35_negative_encoder_hidden_states"] = neg_prompt_embeds
|
||||
batch.extra["sd35_negative_pooled_projections"] = neg_pooled
|
||||
|
||||
return batch
|
||||
|
||||
|
||||
class SD35DenoisingStage(PipelineStage):
|
||||
"""Denoising loop for SD3.5 (2D transformer + FlowMatch scheduler)."""
|
||||
|
||||
def __init__(self, transformer, scheduler) -> None:
|
||||
super().__init__()
|
||||
self.transformer = transformer
|
||||
self.scheduler = scheduler
|
||||
|
||||
@staticmethod
|
||||
def _prepare_extra_func_kwargs(func, kwargs) -> dict[str, Any]:
|
||||
extra_kwargs: dict[str, Any] = {}
|
||||
sig = inspect.signature(func)
|
||||
for k, v in kwargs.items():
|
||||
if k in sig.parameters:
|
||||
extra_kwargs[k] = v
|
||||
return extra_kwargs
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.timesteps is None:
|
||||
raise ValueError("timesteps must be set before SD35DenoisingStage")
|
||||
if batch.latents is None:
|
||||
raise ValueError("latents must be set before SD35DenoisingStage")
|
||||
|
||||
prompt_embeds: torch.Tensor = batch.extra["sd35_encoder_hidden_states"]
|
||||
pooled: torch.Tensor = batch.extra["sd35_pooled_projections"]
|
||||
|
||||
neg_prompt_embeds: torch.Tensor | None = batch.extra.get(
|
||||
"sd35_negative_encoder_hidden_states")
|
||||
neg_pooled: torch.Tensor | None = batch.extra.get(
|
||||
"sd35_negative_pooled_projections")
|
||||
|
||||
timesteps = batch.timesteps
|
||||
latents = batch.latents
|
||||
guidance_scale = float(batch.guidance_scale)
|
||||
|
||||
target_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.dit_precision]
|
||||
autocast_enabled = (target_dtype != torch.float32
|
||||
) and not fastvideo_args.disable_autocast
|
||||
|
||||
extra_step_kwargs = self._prepare_extra_func_kwargs(
|
||||
self.scheduler.step, {
|
||||
"generator":
|
||||
batch.generator[0]
|
||||
if isinstance(batch.generator, list) else batch.generator
|
||||
})
|
||||
|
||||
for t in timesteps:
|
||||
latents_4d = latents.squeeze(2)
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
if neg_prompt_embeds is None or neg_pooled is None:
|
||||
raise ValueError(
|
||||
"Missing negative conditioning tensors for CFG")
|
||||
latent_model_input = torch.cat([latents_4d] * 2, dim=0)
|
||||
cond_embeds = torch.cat([neg_prompt_embeds, prompt_embeds],
|
||||
dim=0)
|
||||
cond_pooled = torch.cat([neg_pooled, pooled], dim=0)
|
||||
else:
|
||||
latent_model_input = latents_4d
|
||||
cond_embeds = prompt_embeds
|
||||
cond_pooled = pooled
|
||||
|
||||
latent_model_input = self.scheduler.scale_model_input(
|
||||
latent_model_input, t)
|
||||
timestep = t.expand(latent_model_input.shape[0])
|
||||
|
||||
with torch.autocast(
|
||||
device_type="cuda",
|
||||
dtype=target_dtype,
|
||||
enabled=autocast_enabled
|
||||
and (get_local_torch_device().type == "cuda"),
|
||||
):
|
||||
noise_pred = self.transformer(
|
||||
hidden_states=latent_model_input,
|
||||
timestep=timestep,
|
||||
encoder_hidden_states=cond_embeds,
|
||||
pooled_projections=cond_pooled,
|
||||
return_dict=False,
|
||||
)[0]
|
||||
|
||||
if batch.do_classifier_free_guidance:
|
||||
noise_pred_uncond, noise_pred_text = noise_pred.chunk(2)
|
||||
noise_pred = noise_pred_uncond + guidance_scale * (
|
||||
noise_pred_text - noise_pred_uncond)
|
||||
|
||||
noise_pred_5d = noise_pred.unsqueeze(2)
|
||||
latents = self.scheduler.step(noise_pred_5d,
|
||||
t,
|
||||
latents,
|
||||
return_dict=False,
|
||||
**extra_step_kwargs)[0]
|
||||
|
||||
batch.latents = latents
|
||||
return batch
|
||||
|
||||
|
||||
class SD35DecodingStage(PipelineStage):
|
||||
|
||||
def __init__(self, vae) -> None:
|
||||
super().__init__()
|
||||
self.vae = vae
|
||||
|
||||
@staticmethod
|
||||
def _denormalize_latents(latents: torch.Tensor, vae) -> torch.Tensor:
|
||||
# Prefer config fields to avoid diffusers deprecation warnings for direct
|
||||
# attribute access (vae.scaling_factor / vae.shift_factor).
|
||||
cfg = getattr(vae, "config", None)
|
||||
sf = getattr(cfg, "scaling_factor", None) if cfg is not None else None
|
||||
sh = getattr(cfg, "shift_factor", None) if cfg is not None else None
|
||||
|
||||
if sf is None and hasattr(vae, "scaling_factor"):
|
||||
sf = vae.scaling_factor
|
||||
if sh is None and hasattr(vae, "shift_factor"):
|
||||
sh = vae.shift_factor
|
||||
|
||||
if sf is not None:
|
||||
latents = latents / (sf.to(latents.device, latents.dtype)
|
||||
if isinstance(sf, torch.Tensor) else sf)
|
||||
if sh is not None:
|
||||
latents = latents + (sh.to(latents.device, latents.dtype)
|
||||
if isinstance(sh, torch.Tensor) else sh)
|
||||
return latents
|
||||
|
||||
@torch.no_grad()
|
||||
def forward(self, batch: ForwardBatch,
|
||||
fastvideo_args: FastVideoArgs) -> ForwardBatch:
|
||||
if batch.latents is None:
|
||||
raise ValueError("latents must be set before SD35DecodingStage")
|
||||
|
||||
device = get_local_torch_device()
|
||||
latents_5d = batch.latents.to(device)
|
||||
latents_4d = latents_5d.squeeze(2)
|
||||
|
||||
vae_dtype = PRECISION_TO_TYPE[
|
||||
fastvideo_args.pipeline_config.vae_precision]
|
||||
autocast_enabled = (
|
||||
vae_dtype != torch.float32) and not fastvideo_args.disable_autocast
|
||||
|
||||
latents_4d = self._denormalize_latents(latents_4d, self.vae)
|
||||
|
||||
with torch.autocast(
|
||||
device_type="cuda",
|
||||
dtype=vae_dtype,
|
||||
enabled=autocast_enabled and (device.type == "cuda"),
|
||||
):
|
||||
if not autocast_enabled:
|
||||
latents_4d = latents_4d.to(dtype=vae_dtype)
|
||||
dec = self.vae.decode(latents_4d)
|
||||
image = dec.sample if hasattr(dec, "sample") else dec[0]
|
||||
|
||||
image = (image / 2 + 0.5).clamp(0, 1)
|
||||
batch.output = image.unsqueeze(2).detach().float().cpu()
|
||||
return batch
|
||||
@@ -253,11 +253,17 @@ class TextEncodingStage(PipelineStage):
|
||||
input_ids = text_inputs["input_ids"]
|
||||
attention_mask = text_inputs["attention_mask"]
|
||||
|
||||
want_hidden_states = bool(
|
||||
getattr(getattr(encoder_config, "arch_config", None),
|
||||
"output_hidden_states", False))
|
||||
if is_ltx2:
|
||||
want_hidden_states = True
|
||||
|
||||
with set_forward_context(current_timestep=0, attn_metadata=None):
|
||||
outputs = text_encoder(
|
||||
input_ids=input_ids,
|
||||
attention_mask=attention_mask,
|
||||
output_hidden_states=True,
|
||||
output_hidden_states=want_hidden_states,
|
||||
)
|
||||
|
||||
try:
|
||||
|
||||
@@ -22,6 +22,7 @@ from fastvideo.configs.pipelines.hunyuan15 import (
|
||||
Hunyuan15T2V480PConfig, Hunyuan15I2V480PStepDistilledConfig,
|
||||
Hunyuan15T2V720PConfig, Hunyuan15I2V720PConfig, Hunyuan15SR1080PConfig)
|
||||
from fastvideo.configs.pipelines.hyworld import HYWorldConfig
|
||||
from fastvideo.configs.pipelines.lingbotworld import LingBotWorldI2V480PConfig
|
||||
from fastvideo.configs.pipelines.longcat import LongCatT2V480PConfig
|
||||
from fastvideo.configs.pipelines.ltx2 import LTX2T2VConfig
|
||||
from fastvideo.configs.pipelines.stepvideo import StepVideoT2VConfig
|
||||
@@ -45,6 +46,7 @@ from fastvideo.configs.pipelines.wan import (
|
||||
WanT2V480PConfig,
|
||||
WanT2V720PConfig,
|
||||
)
|
||||
from fastvideo.configs.pipelines.sd35 import SD35Config
|
||||
from fastvideo.configs.sample.base import SamplingParam
|
||||
from fastvideo.configs.sample.cosmos import (
|
||||
Cosmos_Predict2_2B_Video2World_SamplingParam, )
|
||||
@@ -57,7 +59,9 @@ from fastvideo.configs.sample.hunyuan15 import (
|
||||
Hunyuan15_720P_SamplingParam, Hunyuan15_720P_Distilled_I2V_SamplingParam,
|
||||
Hunyuan15_SR_1080P_SamplingParam)
|
||||
from fastvideo.configs.sample.hyworld import HYWorld_SamplingParam
|
||||
from fastvideo.configs.sample.ltx2 import LTX2SamplingParam
|
||||
from fastvideo.configs.sample.lingbotworld import LingBotWorld_SamplingParam
|
||||
from fastvideo.configs.sample.ltx2 import (LTX2BaseSamplingParam,
|
||||
LTX2DistilledSamplingParam)
|
||||
from fastvideo.configs.sample.stepvideo import StepVideoT2VSamplingParam
|
||||
from fastvideo.configs.sample.turbodiffusion import (
|
||||
TurboDiffusionI2V_A14B_SamplingParam,
|
||||
@@ -79,6 +83,8 @@ from fastvideo.configs.sample.wan import (
|
||||
WanT2V_14B_SamplingParam,
|
||||
WanT2V_1_3B_SamplingParam,
|
||||
)
|
||||
from fastvideo.configs.sample.sd35 import SD35SamplingParam
|
||||
|
||||
from fastvideo.fastvideo_args import WorkloadType
|
||||
from fastvideo.logger import init_logger
|
||||
from fastvideo.utils import (maybe_download_model_index,
|
||||
@@ -239,17 +245,30 @@ def _get_config_info(
|
||||
|
||||
|
||||
def _register_configs() -> None:
|
||||
# LTX-2
|
||||
# LTX-2 (base)
|
||||
register_configs(
|
||||
sampling_param_cls=LTX2SamplingParam,
|
||||
sampling_param_cls=LTX2BaseSamplingParam,
|
||||
pipeline_config_cls=LTX2T2VConfig,
|
||||
hf_model_paths=[
|
||||
"Lightricks/LTX-2",
|
||||
"converted/ltx2_diffusers",
|
||||
"FastVideo/LTX2-base",
|
||||
"FastVideo/LTX2-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: ("ltx2" in path.lower() or "ltx-2" in path.lower()) and
|
||||
"distilled" not in path.lower(),
|
||||
],
|
||||
)
|
||||
# LTX-2 (distilled)
|
||||
register_configs(
|
||||
sampling_param_cls=LTX2DistilledSamplingParam,
|
||||
pipeline_config_cls=LTX2T2VConfig,
|
||||
hf_model_paths=[
|
||||
"FastVideo/LTX2-Distilled-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: "ltx2" in path.lower() or "ltx-2" in path.lower(),
|
||||
lambda path: ("ltx2" in path.lower() or "ltx-2" in path.lower()) and
|
||||
"distilled" in path.lower(),
|
||||
],
|
||||
)
|
||||
|
||||
@@ -326,6 +345,19 @@ def _register_configs() -> None:
|
||||
model_detectors=[lambda path: "hyworld" in path.lower()],
|
||||
)
|
||||
|
||||
# LingBotWorld
|
||||
register_configs(
|
||||
sampling_param_cls=LingBotWorld_SamplingParam,
|
||||
pipeline_config_cls=LingBotWorldI2V480PConfig,
|
||||
hf_model_paths=[
|
||||
"FastVideo/LingBot-World-Base-Cam-Diffusers",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path:
|
||||
("lingbotworld" in path.lower() or "lingbot-world" in path.lower())
|
||||
],
|
||||
)
|
||||
|
||||
# LongCat
|
||||
register_configs(
|
||||
sampling_param_cls=None,
|
||||
@@ -524,6 +556,22 @@ def _register_configs() -> None:
|
||||
],
|
||||
)
|
||||
|
||||
# SD3.5
|
||||
register_configs(
|
||||
sampling_param_cls=SD35SamplingParam,
|
||||
pipeline_config_cls=SD35Config,
|
||||
hf_model_paths=[
|
||||
"stabilityai/stable-diffusion-3.5-medium",
|
||||
],
|
||||
model_detectors=[
|
||||
lambda path: any(token in path.lower() for token in (
|
||||
"sd35",
|
||||
"stablediffusion3",
|
||||
"stabilityai__stable-diffusion-3.5-medium",
|
||||
)),
|
||||
],
|
||||
)
|
||||
|
||||
|
||||
# --- Part 3: Main Resolver ---
|
||||
|
||||
|
||||