Compare commits
8
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
4058e92fb4 | ||
|
|
de6f92b896 | ||
|
|
aaeee0b3ac | ||
|
|
ee2c06890b | ||
|
|
7186408be9 | ||
|
|
878524bc28 | ||
|
|
44da53f2a9 | ||
|
|
a59b7b04d9 |
@@ -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
|
||||
}
|
||||
]
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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),
|
||||
|
||||
Reference in New Issue
Block a user