Compare commits
1
Commits
will/f_0
...
wei/issues
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
8cb2ae9d27 |
@@ -4,7 +4,7 @@ export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=online
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
|
||||
MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers"
|
||||
MODEL_PATH="wlsaidhi/SFWan2.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,7 +14,6 @@ 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
|
||||
@@ -34,7 +33,7 @@ parallel_args=(
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1
|
||||
--hsdp_shard_dim 1
|
||||
--hsdp_shard_dim $NUM_GPUS
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
@@ -51,20 +50,17 @@ dataset_args=(
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
--log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "50"
|
||||
--validation_guidance_scale "6.0"
|
||||
--log-visualization
|
||||
--visualization-steps 100
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 6e-6
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--weight_only_checkpointing_steps 1000
|
||||
--training_state_checkpointing_steps 1000
|
||||
--weight_decay 1e-4
|
||||
--weight_decay 0.01
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
|
||||
@@ -916,6 +916,7 @@ 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
|
||||
@@ -1079,6 +1080,9 @@ 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")
|
||||
|
||||
@@ -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
|
||||
from fastvideo.platforms import AttentionBackendEnum, current_platform
|
||||
|
||||
logger = init_logger(__name__)
|
||||
class CausalWanSelfAttention(nn.Module):
|
||||
@@ -286,6 +286,8 @@ 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,
|
||||
@@ -452,7 +454,6 @@ 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):
|
||||
|
||||
@@ -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(
|
||||
|
||||
@@ -314,6 +314,7 @@ 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)
|
||||
|
||||
@@ -43,7 +43,10 @@ from fastvideo.training.training_utils import (
|
||||
from fastvideo.utils import (is_vsa_available, maybe_download_model,
|
||||
set_random_seed, verify_model_config_and_directory)
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -18,7 +18,10 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -17,6 +17,7 @@ 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 (
|
||||
@@ -57,15 +58,17 @@ 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],
|
||||
self.dmd_denoising_steps = torch.tensor([1000, 750, 500, 250, 0],
|
||||
dtype=torch.long,
|
||||
device=get_local_torch_device())
|
||||
if training_args.warp_denoising_step: # Warp the denoising step according to the scheduler time shift
|
||||
@@ -161,27 +164,12 @@ class ODEInitTrainingPipeline(TrainingPipeline):
|
||||
device, dtype=torch.bfloat16)
|
||||
training_batch.infos = infos
|
||||
|
||||
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)
|
||||
"""
|
||||
return training_batch, trajectory_latents[:, :, :self.training_args.
|
||||
num_latent_t].to(
|
||||
device,
|
||||
dtype=torch.bfloat16
|
||||
), trajectory_timesteps.to(
|
||||
device)
|
||||
|
||||
def _get_timestep(self,
|
||||
min_timestep: int,
|
||||
@@ -225,7 +213,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], dtype=torch.long).cpu()
|
||||
[0, 12, 24, 36, S - 1], 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)
|
||||
@@ -367,8 +355,7 @@ 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.validation_pipeline.decoding_stage.decode(
|
||||
latent, training_args)
|
||||
pixel_latent = self.decoding_stage.decode(latent, training_args)
|
||||
|
||||
video = pixel_latent.cpu().float()
|
||||
video = video.permute(0, 2, 1, 3, 4)
|
||||
|
||||
@@ -30,7 +30,10 @@ from fastvideo.profiler import profile_region
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
|
||||
class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
|
||||
@@ -13,16 +13,18 @@ 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
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionMetadataBuilder)
|
||||
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
|
||||
try:
|
||||
from fastvideo.attention.backends.video_sparse_attn import (
|
||||
VideoSparseAttentionMetadataBuilder)
|
||||
from fastvideo.attention.backends.vmoba import VideoMobaAttentionMetadataBuilder
|
||||
except Exception:
|
||||
pass
|
||||
from fastvideo.configs.sample import SamplingParam
|
||||
from fastvideo.dataset import build_parquet_map_style_dataloader
|
||||
from fastvideo.dataset.dataloader.schema import pyarrow_schema_t2v
|
||||
@@ -48,8 +50,12 @@ from fastvideo.training.training_utils import (
|
||||
from fastvideo.utils import (is_vmoba_available, is_vsa_available,
|
||||
set_random_seed, shallow_asdict)
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
vmoba_available = is_vmoba_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
vmoba_available = is_vmoba_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
vmoba_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@@ -108,7 +114,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)
|
||||
set_random_seed(self.seed + self.global_rank)
|
||||
self.transformer.train()
|
||||
if training_args.enable_gradient_checkpointing_type is not None:
|
||||
self.transformer = apply_activation_checkpointing(
|
||||
@@ -588,15 +594,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.noise_random_generator = torch.Generator(
|
||||
device="cpu").manual_seed(self.seed + self.global_rank)
|
||||
self.noise_gen_cuda = torch.Generator(
|
||||
device=current_platform.device_name).manual_seed(self.seed)
|
||||
device=current_platform.device_name).manual_seed(self.seed +
|
||||
self.global_rank)
|
||||
self.validation_random_generator = torch.Generator(
|
||||
device="cpu").manual_seed(self.seed)
|
||||
logger.info("Initialized random seeds with seed: %s", self.seed)
|
||||
|
||||
self.noise_scheduler = FlowMatchEulerDiscreteScheduler()
|
||||
device="cpu").manual_seed(self.seed + self.global_rank)
|
||||
logger.info("Initialized random seeds with seed: %s",
|
||||
self.seed + self.global_rank)
|
||||
|
||||
if self.training_args.resume_from_checkpoint:
|
||||
self._resume_from_checkpoint()
|
||||
@@ -661,26 +667,31 @@ class TrainingPipeline(LoRAPipeline, ABC):
|
||||
"grad_norm": grad_norm,
|
||||
"vsa_sparsity": current_vsa_sparsity,
|
||||
}
|
||||
metrics["batch_size"] = int(training_batch.raw_latent_shape[0])
|
||||
try:
|
||||
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
|
||||
metrics["hidden_dim"] = arch_config.hidden_size
|
||||
metrics["num_layers"] = arch_config.num_layers
|
||||
metrics["ffn_dim"] = arch_config.ffn_dim
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
self.tracker.log(metrics, step)
|
||||
if step % self.training_args.training_state_checkpointing_steps == 0:
|
||||
@@ -693,12 +704,14 @@ 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(
|
||||
|
||||
@@ -257,8 +257,6 @@ 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")
|
||||
@@ -290,8 +288,6 @@ 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",
|
||||
@@ -417,6 +413,67 @@ 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)
|
||||
@@ -454,46 +511,45 @@ 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,
|
||||
@@ -644,18 +700,37 @@ def load_distillation_checkpoint(
|
||||
end_time - begin_time,
|
||||
local_main_process_only=False)
|
||||
|
||||
# Load EMA state if available and generator_ema is provided
|
||||
# Load EMA separately if saved in rank0_full mode
|
||||
if generator_ema is not None:
|
||||
try:
|
||||
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)
|
||||
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)
|
||||
except Exception as e:
|
||||
logger.warning("rank: %s, failed to load EMA state: %s", rank,
|
||||
logger.warning("rank: %s, failed to load generator EMA: %s", rank,
|
||||
str(e))
|
||||
|
||||
# Load generator_2 distributed checkpoint (MoE support)
|
||||
@@ -850,7 +925,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 * 0.476986
|
||||
return latents * vae.config.scaling_factor
|
||||
elif model_type == "wan":
|
||||
latents_mean = torch.tensor(vae.latents_mean)
|
||||
latents_std = 1.0 / torch.tensor(vae.latents_std)
|
||||
@@ -1169,6 +1244,71 @@ 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:
|
||||
@@ -1795,5 +1935,5 @@ class EMA_FSDP:
|
||||
self.saved.clear()
|
||||
return False
|
||||
|
||||
def apply_to_model(self, module):
|
||||
def apply_to_model(self, module: torch.nn.Module) -> _ApplyEMACtx:
|
||||
return EMA_FSDP._ApplyEMACtx(self, module)
|
||||
|
||||
@@ -10,7 +10,10 @@ from fastvideo.pipelines.basic.wan.wan_dmd_pipeline import WanDMDPipeline
|
||||
from fastvideo.training.distillation_pipeline import DistillationPipeline
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -19,7 +19,10 @@ 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
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -18,7 +18,10 @@ from fastvideo.pipelines.pipeline_batch_info import ForwardBatch, TrainingBatch
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available, shallow_asdict
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -10,7 +10,10 @@ from fastvideo.training.self_forcing_distillation_pipeline import (
|
||||
SelfForcingDistillationPipeline)
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
@@ -10,7 +10,10 @@ from fastvideo.pipelines.basic.wan.wan_pipeline import WanPipeline
|
||||
from fastvideo.training.training_pipeline import TrainingPipeline
|
||||
from fastvideo.utils import is_vsa_available
|
||||
|
||||
vsa_available = is_vsa_available()
|
||||
try:
|
||||
vsa_available = is_vsa_available()
|
||||
except Exception:
|
||||
vsa_available = False
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user