Compare commits

...
Author SHA1 Message Date
SolitaryThinker 4058e92fb4 separate critic lr 2025-07-21 19:26:04 +00:00
SolitaryThinker de6f92b896 update 2025-07-18 23:20:30 +00:00
SolitaryThinker aaeee0b3ac add noisy ground truth for teacher 2025-07-18 22:28:15 +00:00
SolitaryThinker ee2c06890b multi-step simulation 2025-07-18 22:21:41 +00:00
SolitaryThinker 7186408be9 add sample param for wan2.1-fun inp 2025-07-18 21:12:56 +00:00
SolitaryThinker 878524bc28 add more samples to crush-smol val.json 2025-07-18 21:11:25 +00:00
SolitaryThinker 44da53f2a9 start from step 1 2025-07-18 20:58:05 +00:00
SolitaryThinker a59b7b04d9 remove breakpoint 2025-07-18 20:31:09 +00:00
7 changed files with 180 additions and 12 deletions
@@ -26,6 +26,24 @@
"height": 480,
"width": 832,
"num_frames": 77
},
{
"caption": "The video shows a close-up of an orange being flattened as if it were under a hydraulic press, with the press moving down and compressing the fruit until it is completely flattened.",
"image_path": null,
"video_path": "validation_dataset/EJqsC21GSBY-Scene-059.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 98
},
{
"caption": "The video shows a stack of colorful sponges being flattened as if they were under a hydraulic press. The sponges, which are pink, white, blue, and green, are compressed into a smaller size, demonstrating the press's power. The background features a green wall with a yellow and red sign, adding context to the setting.",
"image_path": null,
"video_path": "validation_dataset/EJqsC21GSBY-Scene-013.mp4",
"num_inference_steps": 40,
"height": 480,
"width": 832,
"num_frames": 148
}
]
}
+4 -1
View File
@@ -6,7 +6,8 @@ from typing import Any
from fastvideo.v1.configs.sample.hunyuan import (FastHunyuanSamplingParam,
HunyuanSamplingParam)
from fastvideo.v1.configs.sample.stepvideo import StepVideoT2VSamplingParam
from fastvideo.v1.configs.sample.wan import (WanI2V_14B_480P_SamplingParam,
from fastvideo.v1.configs.sample.wan import (Wan2_1_Fun_1_3B_InP_SamplingParam,
WanI2V_14B_480P_SamplingParam,
WanI2V_14B_720P_SamplingParam,
WanT2V_1_3B_SamplingParam,
WanT2V_14B_SamplingParam)
@@ -23,6 +24,8 @@ SAMPLING_PARAM_REGISTRY: dict[str, Any] = {
"Wan-AI/Wan2.1-T2V-14B-Diffusers": WanT2V_14B_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-480P-Diffusers": WanI2V_14B_480P_SamplingParam,
"Wan-AI/Wan2.1-I2V-14B-720P-Diffusers": WanI2V_14B_720P_SamplingParam,
"weizhou03/Wan2.1-Fun-1.3B-InP-Diffusers":
Wan2_1_Fun_1_3B_InP_SamplingParam,
"FastVideo/stepvideo-t2v-diffusers": StepVideoT2VSamplingParam,
# Add other specific weight variants
}
+17
View File
@@ -94,3 +94,20 @@ class WanI2V_14B_720P_SamplingParam(WanT2V_14B_SamplingParam):
-5784.54975374, 5449.50911966, -1811.16591783, 256.27178429,
-13.02252404
]))
# =============================================
# ============= Wan2.1 Fun Models =============
# =============================================
@dataclass
class Wan2_1_Fun_1_3B_InP_SamplingParam(SamplingParam):
"""Sampling parameters for Wan2.1 Fun 1.3B InP model."""
height: int = 480
width: int = 832
num_frames: int = 81
fps: int = 16
negative_prompt: str | None = "色调艳丽,过曝,静态,细节模糊不清,字幕,风格,作品,画作,画面,静止,整体发灰,最差质量,低质量,JPEG压缩残留,丑陋的,残缺的,多余的手指,画得不好的手部,画得不好的脸部,畸形的,毁容的,形态畸形的肢体,手指融合,静止不动的画面,杂乱的背景,三条腿,背景人很多,倒着走"
guidance_scale: float = 6.0
num_inference_steps: int = 50
+28
View File
@@ -428,6 +428,7 @@ class TrainingArgs(FastVideoArgs):
learning_rate: float = 0.0
scale_lr: bool = False
lr_scheduler: str = "constant"
lr_step_rules: str | None = None
lr_warmup_steps: int = 0
max_grad_norm: float = 0.0
enable_gradient_checkpointing_type: str | None = None
@@ -466,9 +467,14 @@ class TrainingArgs(FastVideoArgs):
# distillation args
student_critic_update_ratio: int = 5
critic_learning_rate: float = 1e-5
critic_lr_scheduler: str = "constant"
critic_lr_step_rules: str | None = None
min_step_ratio: float = 0.2
max_step_ratio: float = 0.98
teacher_guidance_scale: float = 3.5
simulate_student_forward: bool = False
num_teacher_noisy_ground_truth_steps: int = 0
@classmethod
def from_cli_args(cls, args: argparse.Namespace) -> "TrainingArgs":
@@ -628,6 +634,9 @@ class TrainingArgs(FastVideoArgs):
type=str,
default="constant",
help="Learning rate scheduler type")
parser.add_argument("--lr-step-rules",
type=str,
help="Learning rate step rules")
parser.add_argument("--lr-warmup-steps",
type=int,
default=10,
@@ -741,6 +750,17 @@ class TrainingArgs(FastVideoArgs):
type=int,
default=TrainingArgs.student_critic_update_ratio,
help="Ratio of student updates to critic updates.")
parser.add_argument("--critic-learning-rate",
type=float,
default=TrainingArgs.critic_learning_rate,
help="Learning rate for critic")
parser.add_argument("--critic-lr-scheduler",
type=str,
default=TrainingArgs.critic_lr_scheduler,
help="Learning rate scheduler type for critic")
parser.add_argument("--critic-lr-step-rules",
type=str,
help="Learning rate step rules for critic")
parser.add_argument("--min-step-ratio",
type=float,
default=TrainingArgs.min_step_ratio,
@@ -753,6 +773,14 @@ class TrainingArgs(FastVideoArgs):
type=float,
default=TrainingArgs.teacher_guidance_scale,
help="Teacher guidance scale")
parser.add_argument("--simulate-student-forward",
action=StoreBoolean,
default=TrainingArgs.simulate_student_forward,
help="Whether to simulate student forward")
parser.add_argument("--num-teacher-noisy-ground-truth-steps",
type=int,
default=TrainingArgs.num_teacher_noisy_ground_truth_steps,
help="Number of steps to use noisy ground truth for teacher")
return parser
+108 -10
View File
@@ -103,14 +103,18 @@ class DistillationPipeline(TrainingPipeline):
critic_params = list(filter(lambda p: p.requires_grad, self.critic_transformer.parameters()))
self.critic_transformer_optimizer = torch.optim.AdamW(
critic_params,
lr=training_args.learning_rate,
lr=training_args.critic_learning_rate,
betas=(0.9, 0.999),
weight_decay=training_args.weight_decay,
eps=1e-8,
)
if training_args.critic_lr_scheduler == "piecewise_constant":
assert training_args.critic_lr_step_rules is not None, "critic lr step rules is required when using piecewise_constant lr scheduler"
self.critic_lr_scheduler = get_scheduler(
training_args.lr_scheduler,
training_args.critic_lr_scheduler,
step_rules=training_args.critic_lr_step_rules,
optimizer=self.critic_transformer_optimizer,
num_warmup_steps=training_args.lr_warmup_steps * self.world_size,
num_training_steps=training_args.max_train_steps * self.world_size,
@@ -234,25 +238,110 @@ class DistillationPipeline(TrainingPipeline):
pred_video = pred_video.type_as(noisy_input)
return pred_video, timestep.float().detach()
def _multi_step_simulation_student_forward(self, training_batch: TrainingBatch) -> torch.Tensor:
"""Forward pass through student transformer matching inference procedure."""
from fastvideo.v1.training.training_utils import DiffusionWrapper
latents = training_batch.latents
dtype = latents.dtype
# Step 1: Randomly sample a target timestep index from denoising_step_list
target_timestep_idx = torch.randint(0, len(self.denoising_step_list), [
self.video_latent_shape[0], self.video_latent_shape[1]], device=self.device, dtype=torch.long)
target_timestep_idx = self._process_timestep(target_timestep_idx, type=self.distill_task_type)
target_timestep = self.denoising_step_list[target_timestep_idx]
# Step 2: Simulate the multi-step inference process up to the target timestep
# Start from pure noise like in inference
current_latents = torch.randn(self.video_latent_shape, device=self.device, dtype=dtype)
# Only run intermediate steps if target_timestep_idx > 0
max_target_idx = target_timestep_idx.max().item()
if max_target_idx > 0:
# Run student model for all steps before the target timestep
with torch.no_grad():
for step_idx in range(max_target_idx):
current_timestep = self.denoising_step_list[step_idx]
logger.info(f"target_timestep: {target_timestep}, current_timestep: {current_timestep}")
current_timestep_tensor = current_timestep * torch.ones(
self.video_latent_shape[:2], device=self.device, dtype=torch.long)
# Run student model to get flow prediction
training_batch_temp = self._build_input_kwargs(
current_latents, current_timestep_tensor, training_batch.conditional_dict, training_batch)
pred_flow = self.student_transformer.model(**training_batch_temp.input_kwargs).permute(0, 2, 1, 3, 4)
# Convert flow prediction to x0 prediction
pred_clean = DiffusionWrapper._convert_flow_pred_to_x0(
flow_pred=pred_flow.flatten(0, 1),
xt=current_latents.flatten(0, 1),
timestep=current_timestep_tensor.flatten(0, 1),
scheduler=self.noise_scheduler
).unflatten(0, self.video_latent_shape[:2])
# Add noise for the next timestep
next_timestep = self.denoising_step_list[step_idx + 1]
next_timestep_tensor = next_timestep * torch.ones(
self.video_latent_shape[:2], device=self.device, dtype=torch.long)
current_latents = self.noise_scheduler.add_noise(
pred_clean.flatten(0, 1),
torch.randn_like(pred_clean.flatten(0, 1)),
next_timestep_tensor.flatten(0, 1)
).unflatten(0, self.video_latent_shape[:2])
# Step 3: Use the simulated noisy input for the final training step
# For timestep index 0, this is pure noise
# For timestep index k > 0, this is the result after k denoising steps + noise at target level
noisy_input = current_latents
# Step 4: Final student prediction (this is what we train on)
training_batch = self._build_input_kwargs(noisy_input, target_timestep, training_batch.conditional_dict, training_batch)
pred_video = self.student_transformer(training_batch, target_timestep)
pred_video = pred_video.type_as(noisy_input)
return pred_video, target_timestep.float().detach()
def _compute_kl_grad(
self, noisy_video: torch.Tensor,
self,
noisy_video: torch.Tensor,
estimated_clean_video: torch.Tensor,
noise: torch.Tensor,
timestep: torch.Tensor,
training_batch: TrainingBatch,
normalization: bool = True
) -> Tuple[torch.Tensor, dict]:
assert self.training_args is not None
# critic_transformer forward
training_batch = self._build_input_kwargs(noisy_video, timestep, training_batch.conditional_dict, training_batch)
pred_fake_video = self.critic_transformer(training_batch, timestep)
if self.current_trainstep > self.training_args.num_teacher_noisy_ground_truth_steps:
teacher_noisy_input = noisy_video
teacher_timestep = timestep
else:
teacher_timestep = timestep
logger.info(f"Using noisy ground truth for teacher with timestep {teacher_timestep}")
batch_size, num_frame = self.video_latent_shape[:2]
noisy_ground_truth = self.noise_scheduler.add_noise(
training_batch.latents.flatten(0, 1),
noise.flatten(0, 1),
teacher_timestep.flatten(0, 1)
).detach().unflatten(0, (batch_size, num_frame))
teacher_noisy_input = noisy_ground_truth
# teacher_transformer cond forward
training_batch = self._build_input_kwargs(noisy_video, timestep, training_batch.conditional_dict, training_batch)
pred_real_video_cond = self.teacher_transformer(training_batch, timestep)
training_batch = self._build_input_kwargs(teacher_noisy_input, teacher_timestep, training_batch.conditional_dict, training_batch)
pred_real_video_cond = self.teacher_transformer(training_batch, teacher_timestep)
# teacher_transformer uncond forward
training_batch = self._build_input_kwargs(noisy_video, timestep, training_batch.unconditional_dict, training_batch)
pred_real_video_uncond = self.teacher_transformer(training_batch, timestep)
training_batch = self._build_input_kwargs(teacher_noisy_input, teacher_timestep, training_batch.unconditional_dict, training_batch)
pred_real_video_uncond = self.teacher_transformer(training_batch, teacher_timestep)
pred_real_video = pred_real_video_cond + (
pred_real_video_cond - pred_real_video_uncond
@@ -311,6 +400,7 @@ class DistillationPipeline(TrainingPipeline):
grad, dmd_log_dict = self._compute_kl_grad(
noisy_video=noisy_latent,
estimated_clean_video=original_latent,
noise=noise,
timestep=timestep,
training_batch=training_batch
)
@@ -322,12 +412,16 @@ class DistillationPipeline(TrainingPipeline):
def _student_forward_and_compute_dmd_loss(self, training_batch: TrainingBatch) -> Tuple[TrainingBatch, torch.Tensor, dict]:
"""Forward pass through student transformer and compute student losses."""
assert self.training_args is not None
assert training_batch.conditional_dict is not None
assert training_batch.unconditional_dict is not None
assert training_batch.latents is not None
with set_forward_context(
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata_vsa):
pred_video, timestep_dmd = self._student_forward(training_batch)
if self.training_args.simulate_student_forward:
pred_video, timestep_dmd = self._multi_step_simulation_student_forward(training_batch)
else:
pred_video, timestep_dmd = self._student_forward(training_batch)
with set_forward_context(
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata):
@@ -341,6 +435,7 @@ class DistillationPipeline(TrainingPipeline):
return training_batch, dmd_loss, dmd_log_dict
def _critic_forward_and_compute_loss(self, training_batch: TrainingBatch) -> Tuple[TrainingBatch, torch.Tensor, dict]:
assert self.training_args is not None
assert training_batch.conditional_dict is not None
assert training_batch.unconditional_dict is not None
assert training_batch.latents is not None
@@ -348,7 +443,10 @@ class DistillationPipeline(TrainingPipeline):
with torch.no_grad():
with set_forward_context(
current_timestep=training_batch.timesteps, attn_metadata=training_batch.attn_metadata_vsa):
generated_video, timestep_gen = self._student_forward(training_batch)
if self.training_args.simulate_student_forward:
generated_video, timestep_gen = self._multi_step_simulation_student_forward(training_batch)
else:
generated_video, timestep_gen = self._student_forward(training_batch)
critic_timestep = torch.randint(
0,
@@ -801,7 +899,7 @@ class DistillationPipeline(TrainingPipeline):
disable=self.local_rank > 0,
)
for step in range(self.init_steps,
for step in range(self.init_steps+1,
self.training_args.max_train_steps + 1):
start_time = time.perf_counter()
current_vsa_sparsity = self.training_args.VSA_sparsity if vsa_available else 0.0
@@ -109,8 +109,13 @@ class TrainingPipeline(ComposedPipelineBase, ABC):
self.init_steps = 0
logger.info("optimizer: %s", self.optimizer)
if training_args.lr_scheduler == "piecewise_constant":
assert training_args.lr_step_rules is not None, "lr step rules is required when using piecewise_constant lr scheduler"
self.lr_scheduler = get_scheduler(
training_args.lr_scheduler,
step_rules=training_args.lr_step_rules,
optimizer=self.optimizer,
num_warmup_steps=training_args.lr_warmup_steps * self.world_size,
num_training_steps=training_args.max_train_steps * self.world_size,
-1
View File
@@ -575,7 +575,6 @@ class DiffusionWrapper(torch.nn.Module, ABC):
def forward(self, training_batch: TrainingBatch, timestep: torch.Tensor):
pred_noise = self.model(**training_batch.input_kwargs).permute(0, 2, 1, 3, 4)
torch.distributed.breakpoint()
pred_video = self._convert_flow_pred_to_x0(
flow_pred=pred_noise.flatten(0, 1),
xt=training_batch.noise_latents.flatten(0, 1),