Compare commits

..
Author SHA1 Message Date
Matthew Noto ec6a7e7d06 edit CLAUDE.md 2026-02-09 00:20:39 -08:00
Matthew Notoandgemini-code-assist[bot] 67e3c6ab42 Apply suggestion from @gemini-code-assist[bot]
Co-authored-by: gemini-code-assist[bot] <176961590+gemini-code-assist[bot]@users.noreply.github.com>
2026-02-09 00:11:58 -08:00
Matthew Noto 7a8167db6a add CLAUDE.md 2026-02-08 23:53:11 -08:00
Matthew Noto 3ec4bd0d70 add FastVideo logo 2026-02-08 23:27:50 -08:00
Matthew Noto 0ee24fd35f add AGENTS.md 2026-02-08 23:11:50 -08:00
19 changed files with 180 additions and 298 deletions
+42
View File
@@ -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/`, `images/`, `videos/`, 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
View File
@@ -0,0 +1 @@
@AGENTS.md
+6 -1
View File
@@ -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://ibb.co/sv3MMKyv" target="_blank"> <b> WeChat </b> </a> |
</p>
**FastVideo is a unified post-training and inference framework for accelerated video generation.**
@@ -4,7 +4,7 @@ export WANDB_BASE_URL="https://api.wandb.ai"
export WANDB_MODE=online
export TOKENIZERS_PARALLELISM=false
MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
DATA_DIR="data/crush-smol_processed_t2v_1_3b_ode_init/"
VALIDATION_DATASET_FILE="$(dirname "$0")/validation.json"
NUM_GPUS=1
@@ -14,6 +14,7 @@ NUM_GPUS=1
training_args=(
--tracker_project_name "wan_ode_init"
--output_dir "wan_ode_init_crush_smol"
--override_transformer_cls_name "CausalWanTransformer3DModel"
--wandb_run_name "wan_ode_init_crush_smol"
--max_train_steps 6000
--train_batch_size 1
@@ -33,7 +34,7 @@ parallel_args=(
--sp_size 1
--tp_size 1
--hsdp_replicate_dim 1
--hsdp_shard_dim $NUM_GPUS
--hsdp_shard_dim 1
)
# Model arguments
@@ -50,17 +51,20 @@ dataset_args=(
# Validation arguments
validation_args=(
--log-visualization
--visualization-steps 100
--log_validation
--validation_dataset_file "$VALIDATION_DATASET_FILE"
--validation_steps 50
--validation_sampling_steps "50"
--validation_guidance_scale "6.0"
)
# Optimizer arguments
optimizer_args=(
--learning_rate 1e-5
--learning_rate 6e-6
--mixed_precision "bf16"
--weight_only_checkpointing_steps 1000
--training_state_checkpointing_steps 1000
--weight_decay 0.01
--weight_decay 1e-4
--max_grad_norm 1.0
)
-4
View File
@@ -916,7 +916,6 @@ class TrainingArgs(FastVideoArgs):
training_state_checkpointing_steps: int = 0 # for resuming training
weight_only_checkpointing_steps: int = 0 # for inference
log_visualization: bool = False
visualization_steps: int = 0
# simulate generator forward to match inference
simulate_generator_forward: bool = False
warp_denoising_step: bool = False
@@ -1080,9 +1079,6 @@ class TrainingArgs(FastVideoArgs):
parser.add_argument("--log-validation",
action=StoreBoolean,
help="Whether to log validation results")
parser.add_argument("--visualization-steps",
type=int,
help="Number of visualization steps")
parser.add_argument("--tracker-project-name",
type=str,
help="Project name for tracking")
+2 -3
View File
@@ -33,7 +33,7 @@ from fastvideo.layers.visual_embedding import (PatchEmbed)
from fastvideo.logger import init_logger
from fastvideo.models.dits.base import BaseDiT
from fastvideo.models.dits.wanvideo import WanT2VCrossAttention, WanTimeTextImageEmbedding
from fastvideo.platforms import AttentionBackendEnum, current_platform
from fastvideo.platforms import AttentionBackendEnum
logger = init_logger(__name__)
class CausalWanSelfAttention(nn.Module):
@@ -286,8 +286,6 @@ class CausalWanTransformerBlock(nn.Module):
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)
# 2. Cross-attention
attn_output = self.attn2(norm_hidden_states,
@@ -454,6 +452,7 @@ class CausalWanTransformer3DModel(BaseDiT):
This function will be run for num_frame times.
Process the latent frames one by one (1560 tokens each)
"""
from fastvideo.platforms import current_platform
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
+1 -1
View File
@@ -277,7 +277,7 @@ def load_video(
if convert_method is not None:
pil_images = convert_method(pil_images)
return (pil_images, original_fps) if return_fps else pil_images
return pil_images, original_fps if return_fps else pil_images
def get_default_height_width(
-1
View File
@@ -314,7 +314,6 @@ class DenoisingStage(PipelineStage):
t_expand = timestep.repeat(latent_model_input.shape[0], 1)
else:
t_expand = t.repeat(latent_model_input.shape[0])
t_expand = t_expand.to(get_local_torch_device())
use_meanflow = getattr(self.transformer.config, "use_meanflow",
False)
+1 -4
View File
@@ -43,10 +43,7 @@ from fastvideo.training.training_utils import (
from fastvideo.utils import (is_vsa_available, maybe_download_model,
set_random_seed, verify_model_config_and_directory)
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
@@ -18,10 +18,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
+26 -13
View File
@@ -17,7 +17,6 @@ from fastvideo.models.schedulers.scheduling_self_forcing_flow_match import (
SelfForcingFlowMatchScheduler)
from fastvideo.pipelines.basic.wan.wan_causal_dmd_pipeline import (
WanCausalDMDPipeline)
from fastvideo.pipelines.stages.decoding import DecodingStage
from fastvideo.pipelines.pipeline_batch_info import TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.training.training_utils import (
@@ -58,17 +57,15 @@ class ODEInitTrainingPipeline(TrainingPipeline):
self.vae.requires_grad_(False)
self.timestep_shift = self.training_args.pipeline_config.flow_shift
assert self.timestep_shift == 5.0, "flow_shift must be 5.0"
self.noise_scheduler = SelfForcingFlowMatchScheduler(
shift=self.timestep_shift, sigma_min=0.0, extra_one_step=True)
self.noise_scheduler.set_timesteps(num_inference_steps=1000,
training=True)
self.add_stage(stage_name="decoding_stage",
stage=DecodingStage(vae=self.get_module("vae")))
logger.info("dmd_denoising_steps: %s",
self.training_args.pipeline_config.dmd_denoising_steps)
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250, 0],
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250],
dtype=torch.long,
device=get_local_torch_device())
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
@@ -164,12 +161,27 @@ class ODEInitTrainingPipeline(TrainingPipeline):
device, dtype=torch.bfloat16)
training_batch.infos = infos
return training_batch, trajectory_latents[:, :, :self.training_args.
num_latent_t].to(
device,
dtype=torch.bfloat16
), trajectory_timesteps.to(
device)
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
## TEMP Used for loading the sf .pt files directly
"""
self.manual_idx = self.manual_idx % 155
path = f"/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_pt_vidprom_1000/{self.manual_idx:05d}.pt"
logger.info("path: %s", path)
self.manual_idx += 1
# path = "/mnt/weka/home/hao.zhang/wl/Self-Forcing/ode_single_full/00000.pt"
b = torch.load(path)
training_batch.encoder_hidden_states = b["text_embedding"][0].unsqueeze(
0).to(device, dtype=torch.bfloat16)
trajectory_latents = b["ode_latent"].to(device, dtype=torch.bfloat16)
logger.info("trajectory_latents: %s", trajectory_latents.shape)
logger.info("encoder_hidden_states: %s",
training_batch.encoder_hidden_states.shape)
assert trajectory_latents.shape[1] <= 10, "trajectory_latents.shape[1] must be <= 10"
return training_batch, trajectory_latents.to(
device, dtype=torch.bfloat16), trajectory_timesteps.to(device)
"""
def _get_timestep(self,
min_timestep: int,
@@ -213,7 +225,7 @@ class ODEInitTrainingPipeline(TrainingPipeline):
# Lazily cache nearest trajectory index per DMD step based on the (fixed) S timesteps
if self._cached_closest_idx_per_dmd is None:
self._cached_closest_idx_per_dmd = torch.tensor(
[0, 12, 24, 36, S - 1], dtype=torch.long).cpu()
[0, 12, 24, 36], dtype=torch.long).cpu()
# [0, 1, 2, 3], dtype=torch.long).cpu()
logger.info("self._cached_closest_idx_per_dmd: %s",
self._cached_closest_idx_per_dmd)
@@ -355,7 +367,8 @@ class ODEInitTrainingPipeline(TrainingPipeline):
assert latent_key in latents_vis_dict and latents_vis_dict[
latent_key] is not None
latent = latents_vis_dict[latent_key]
pixel_latent = self.decoding_stage.decode(latent, training_args)
pixel_latent = self.validation_pipeline.decoding_stage.decode(
latent, training_args)
video = pixel_latent.cpu().float()
video = video.permute(0, 2, 1, 3, 4)
@@ -30,10 +30,7 @@ from fastvideo.profiler import profile_region
logger = init_logger(__name__)
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
class SelfForcingDistillationPipeline(DistillationPipeline):
+33 -46
View File
@@ -13,18 +13,16 @@ import numpy as np
import torch
import torch.distributed as dist
import torchvision
from diffusers import FlowMatchEulerDiscreteScheduler
from einops import rearrange
from torch.utils.data import DataLoader
from torchdata.stateful_dataloader import StatefulDataLoader
from tqdm.auto import tqdm
import fastvideo.envs as envs
try:
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
except Exception:
pass
from fastvideo.attention.backends.video_sparse_attn import (
VideoSparseAttentionMetadataBuilder)
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
from fastvideo.configs.sample import SamplingParam
from fastvideo.dataset import build_parquet_map_style_dataloader
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
@@ -50,12 +48,8 @@ from fastvideo.training.training_utils import (
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
set_random_seed, shallow_asdict)
try:
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
except Exception:
vsa_available = False
vmoba_available = False
vsa_available = is_vsa_available()
vmoba_available = is_vmoba_available()
logger = init_logger(__name__)
@@ -114,7 +108,7 @@ class TrainingPipeline(LoRAPipeline, ABC):
# Set random seeds for deterministic training
assert self.seed is not None, "seed must be set"
set_random_seed(self.seed + self.global_rank)
set_random_seed(self.seed)
self.transformer.train()
if training_args.enable_gradient_checkpointing_type is not None:
self.transformer = apply_activation_checkpointing(
@@ -594,15 +588,15 @@ class TrainingPipeline(LoRAPipeline, ABC):
round(num_trainable_params / 1e9, 3))
# Set random seeds for deterministic training
self.noise_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed + self.global_rank)
self.noise_random_generator = torch.Generator(device="cpu").manual_seed(
self.seed)
self.noise_gen_cuda = torch.Generator(
device=current_platform.device_name).manual_seed(self.seed +
self.global_rank)
device=current_platform.device_name).manual_seed(self.seed)
self.validation_random_generator = torch.Generator(
device="cpu").manual_seed(self.seed + self.global_rank)
logger.info("Initialized random seeds with seed: %s",
self.seed + self.global_rank)
device="cpu").manual_seed(self.seed)
logger.info("Initialized random seeds with seed: %s", self.seed)
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
if self.training_args.resume_from_checkpoint:
self._resume_from_checkpoint()
@@ -667,31 +661,26 @@ class TrainingPipeline(LoRAPipeline, ABC):
"grad_norm": grad_norm,
"vsa_sparsity": current_vsa_sparsity,
}
try:
metrics["batch_size"] = int(
training_batch.raw_latent_shape[0])
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (
training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] // patch_h) * (
training_batch.raw_latent_shape[4] // patch_w)
if training_batch.encoder_hidden_states is not None:
context_len = int(
training_batch.encoder_hidden_states.shape[1])
else:
context_len = 0
patch_t, patch_h, patch_w = self.training_args.pipeline_config.dit_config.patch_size
seq_len = (training_batch.raw_latent_shape[2] // patch_t) * (
training_batch.raw_latent_shape[3] //
patch_h) * (training_batch.raw_latent_shape[4] // patch_w)
if training_batch.encoder_hidden_states is not None:
context_len = int(
training_batch.encoder_hidden_states.shape[1])
else:
context_len = 0
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
metrics["dit_seq_len"] = int(seq_len)
metrics["context_len"] = context_len
arch_config = self.training_args.pipeline_config.dit_config.arch_config
arch_config = self.training_args.pipeline_config.dit_config.arch_config
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
except Exception:
pass
metrics["hidden_dim"] = arch_config.hidden_size
metrics["num_layers"] = arch_config.num_layers
metrics["ffn_dim"] = arch_config.ffn_dim
self.tracker.log(metrics, step)
if step % self.training_args.training_state_checkpointing_steps == 0:
@@ -704,14 +693,12 @@ class TrainingPipeline(LoRAPipeline, ABC):
self.noise_random_generator)
self.transformer.train()
self.sp_group.barrier()
if self.training_args.log_visualization and step % self.training_args.visualization_steps == 0:
self.visualize_intermediate_latents(training_batch,
self.training_args, step)
if self.training_args.log_validation and step % self.training_args.validation_steps == 0:
with self.profiler_controller.region(
"profiler_region_training_validation"):
if self.training_args.log_visualization:
self.visualize_intermediate_latents(
training_batch, self.training_args, step)
self._log_validation(self.transformer, self.training_args,
step)
gpu_memory_usage = current_platform.get_torch_device(
+51 -191
View File
@@ -257,6 +257,8 @@ def save_distillation_checkpoint(
if generator_scheduler is not None:
generator_states["scheduler"] = SchedulerWrapper(
generator_scheduler)
if generator_ema is not None:
generator_states["ema"] = generator_ema.state_dict()
generator_dcp_dir = os.path.join(save_dir, "distributed_checkpoint",
"generator")
@@ -288,6 +290,8 @@ def save_distillation_checkpoint(
if generator_scheduler_2 is not None:
generator_2_states["scheduler"] = SchedulerWrapper(
generator_scheduler_2)
if generator_ema_2 is not None:
generator_2_states["ema"] = generator_ema_2.state_dict()
generator_2_dcp_dir = os.path.join(save_dir,
"distributed_checkpoint",
@@ -413,67 +417,6 @@ def save_distillation_checkpoint(
rank,
local_main_process_only=False)
# Persist EMA separately to avoid shape mismatches across ranks.
# Supports:
# - mode="rank0_full": save consolidated EMA only on rank 0
# - mode="local_shard": save per-rank EMA shard for each rank
try:
if generator_ema is not None and getattr(generator_ema, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema,
generator_transformer, rank,
save_dir, "generator_ema")
elif generator_ema is not None and getattr(generator_ema, "mode",
None) == "local_shard":
# Save per-rank shard
ema_dir_shard = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard, exist_ok=True)
ema_shard_path = os.path.join(ema_dir_shard,
f"generator_ema_rank{rank}.pt")
torch.save(generator_ema.state_dict(), ema_shard_path)
logger.info(
"rank: %s, saved generator EMA shard (local_shard) to %s",
rank,
ema_shard_path,
local_main_process_only=False)
# Also consolidate EMA to a single full-state file on rank 0 by applying EMA to the model and gathering
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema, generator_transformer, rank, save_dir,
"generator_ema")
except Exception as e:
logger.warning("rank: %s, failed saving EMA separately: %s", rank,
str(e))
try:
if generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "rank0_full":
_save_rank0_full_ema_safetensors(generator_ema_2,
generator_transformer_2, rank,
save_dir, "generator_ema_2")
elif generator_ema_2 is not None and getattr(generator_ema_2, "mode",
None) == "local_shard":
# Save per-rank shard for EMA_2
ema_dir_shard_2 = os.path.join(save_dir, "ema_local_shard")
os.makedirs(ema_dir_shard_2, exist_ok=True)
ema2_shard_path = os.path.join(ema_dir_shard_2,
f"generator_ema_2_rank{rank}.pt")
torch.save(generator_ema_2.state_dict(), ema2_shard_path)
logger.info(
"rank: %s, saved generator_2 EMA shard (local_shard) to %s",
rank,
ema2_shard_path,
local_main_process_only=False)
# Also consolidate EMA_2 to a single full-state file on rank 0
_consolidate_local_shard_ema_and_save_safetensors(
generator_ema_2, generator_transformer_2, rank, save_dir,
"generator_ema_2")
except Exception as e:
logger.warning("rank: %s, failed saving EMA_2 separately: %s", rank,
str(e))
# Save generator model weights (consolidated) for inference
cpu_state = gather_state_dict_on_cpu_rank0(generator_transformer,
device=None)
@@ -511,45 +454,46 @@ def save_distillation_checkpoint(
logger.info("--> distillation checkpoint saved at step %s to %s", step,
weight_path)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(generator_transformer_2,
device=None)
# Save generator_2 model weights (consolidated) for inference (MoE support)
if generator_transformer_2 is not None:
inference_save_dir_2 = os.path.join(
save_dir, "generator_2_inference_transformer")
cpu_state_2 = gather_state_dict_on_cpu_rank0(
generator_transformer_2, device=None)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(inference_save_dir_2,
"diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
if rank == 0:
os.makedirs(inference_save_dir_2, exist_ok=True)
weight_path_2 = os.path.join(
inference_save_dir_2, "diffusion_pytorch_model.safetensors")
logger.info(
"rank: %s, saving consolidated generator_2 inference checkpoint to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
# Convert training format to diffusers format and save
diffusers_state_dict_2 = custom_to_hf_state_dict(
cpu_state_2,
generator_transformer_2.reverse_param_names_mapping)
save_file(diffusers_state_dict_2, weight_path_2)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
logger.info(
"rank: %s, consolidated generator_2 inference checkpoint saved to %s",
rank,
weight_path_2,
local_main_process_only=False)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2, "config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
# Save model config
config_dict_2 = generator_transformer_2.hf_config
if "dtype" in config_dict_2:
del config_dict_2["dtype"] # TODO
config_path_2 = os.path.join(inference_save_dir_2,
"config.json")
with open(config_path_2, "w") as f:
json.dump(config_dict_2, f, indent=4)
logger.info(
"--> generator_2 distillation checkpoint saved at step %s to %s",
step, weight_path_2)
def load_checkpoint(transformer,
@@ -700,37 +644,18 @@ def load_distillation_checkpoint(
end_time - begin_time,
local_main_process_only=False)
# Load EMA separately if saved in rank0_full mode
# Load EMA state if available and generator_ema is provided
if generator_ema is not None:
try:
if getattr(generator_ema, "mode", None) == "rank0_full":
ema_path = os.path.join(checkpoint_path, "ema",
"generator_ema.pt")
if rank == 0 and os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA (rank0_full) loaded from %s",
rank, ema_path)
elif rank == 0:
logger.info(
"rank: %s, generator EMA file not found at %s; skipping",
rank, ema_path)
elif getattr(generator_ema, "mode", None) == "local_shard":
ema_path = os.path.join(checkpoint_path, "ema_local_shard",
f"generator_ema_rank{rank}.pt")
if os.path.exists(ema_path):
ema_state = torch.load(ema_path, map_location="cpu")
generator_ema.load_state_dict(ema_state)
logger.info(
"rank: %s, generator EMA shard (local_shard) loaded from %s",
rank, ema_path)
else:
logger.info(
"rank: %s, generator EMA shard file not found at %s; skipping",
rank, ema_path)
ema_state = generator_states.get("ema")
if ema_state is not None:
generator_ema.load_state_dict(ema_state)
logger.info("rank: %s, generator EMA state loaded successfully",
rank)
else:
logger.info("rank: %s, no EMA state found in checkpoint", rank)
except Exception as e:
logger.warning("rank: %s, failed to load generator EMA: %s", rank,
logger.warning("rank: %s, failed to load EMA state: %s", rank,
str(e))
# Load generator_2 distributed checkpoint (MoE support)
@@ -925,7 +850,7 @@ def load_distillation_checkpoint(
def normalize_dit_input(model_type, latents, vae) -> torch.Tensor:
if model_type == "hunyuan_hf" or model_type == "hunyuan":
return latents * vae.config.scaling_factor
return latents * 0.476986
elif model_type == "wan":
latents_mean = torch.tensor(vae.latents_mean)
latents_std = 1.0 / torch.tensor(vae.latents_std)
@@ -1244,71 +1169,6 @@ def custom_to_hf_state_dict(
return new_state_dict
def _save_full_ema_safetensors_from_state(
state_dict: dict[str, Any],
reverse_param_names_mapping: dict[str, tuple[str, int, int]],
output_path: str,
) -> None:
"""
Convert a training-format state_dict to HF format and save as safetensors.
"""
diffusers_state_dict = custom_to_hf_state_dict(state_dict,
reverse_param_names_mapping)
save_file(diffusers_state_dict, output_path)
def _save_rank0_full_ema_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
if rank != 0:
return
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
ema_state = ema.state_dict()
_save_full_ema_safetensors_from_state(ema_state,
module.reverse_param_names_mapping,
output_path)
logger.info("rank: %s, saved %s as consolidated EMA safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
def _consolidate_local_shard_ema_and_save_safetensors(
ema: "EMA_FSDP",
module,
rank: int,
save_dir: str,
base_name: str,
) -> None:
try:
# Temporarily apply EMA to the live (sharded) module and gather full CPU state on rank 0
with ema.apply_to_model(module):
cpu_state_full = gather_state_dict_on_cpu_rank0(module, device=None)
if rank == 0:
ema_dir = os.path.join(save_dir, "ema")
os.makedirs(ema_dir, exist_ok=True)
output_path = os.path.join(ema_dir, f"{base_name}.safetensors")
_save_full_ema_safetensors_from_state(
cpu_state_full, module.reverse_param_names_mapping, output_path)
logger.info(
"rank: %s, saved consolidated %s EMA (from local_shard) as safetensors to %s",
rank,
base_name,
output_path,
local_main_process_only=False)
except Exception as ce:
logger.warning(
"rank: %s, failed consolidating %s EMA (local_shard): %s", rank,
base_name, str(ce))
def shift_timestep(timestep: torch.Tensor, shift: float,
num_train_timestep: float) -> torch.Tensor:
if shift == 1:
@@ -1935,5 +1795,5 @@ class EMA_FSDP:
self.saved.clear()
return False
def apply_to_model(self, module: torch.nn.Module) -> _ApplyEMACtx:
def apply_to_model(self, module):
return EMA_FSDP._ApplyEMACtx(self, module)
@@ -10,10 +10,7 @@ from fastvideo.pipelines.basic.wan.wan_dmd_pipeline import WanDMDPipeline
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.utils import is_vsa_available
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
@@ -19,10 +19,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.distillation_pipeline import DistillationPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
@@ -18,10 +18,7 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available, shallow_asdict
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
@@ -10,10 +10,7 @@ from fastvideo.training.self_forcing_distillation_pipeline import (
SelfForcingDistillationPipeline)
from fastvideo.utils import is_vsa_available
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)
+1 -4
View File
@@ -10,10 +10,7 @@ from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
from fastvideo.training.training_pipeline import TrainingPipeline
from fastvideo.utils import is_vsa_available
try:
vsa_available = is_vsa_available()
except Exception:
vsa_available = False
vsa_available = is_vsa_available()
logger = init_logger(__name__)