Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
32711250ad | ||
|
|
3d6cac57fc | ||
|
|
57fd3d8159 | ||
|
|
1eacdd80de | ||
|
|
2ca24b3288 |
@@ -0,0 +1,139 @@
|
||||
# Basic Info
|
||||
export NCCL_P2P_DISABLE=1
|
||||
export TORCH_NCCL_ENABLE_MONITORING=0
|
||||
# different cache dir for different processes
|
||||
export TRITON_CACHE_DIR=/tmp/triton_cache_${SLURM_PROCID}
|
||||
export MASTER_PORT=29500
|
||||
export TOKENIZERS_PARALLELISM=false
|
||||
export WANDB_BASE_URL="https://api.wandb.ai"
|
||||
export WANDB_MODE=offline
|
||||
# export FASTVIDEO_ATTENTION_BACKEND=TORCH_SDPA
|
||||
export WANDB_API_KEY='8d9f4b39abd68eb4e29f6fc010b7ee71a2207cde'
|
||||
|
||||
# Configs
|
||||
NUM_GPUS=8
|
||||
|
||||
# Model paths for DMD distillation:
|
||||
GENERATOR_MODEL_PATH="wlsaidhi/SFWan2.1-T2V-1.3B-Diffusers"
|
||||
REAL_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Teacher model
|
||||
FAKE_SCORE_MODEL_PATH="Wan-AI/Wan2.1-T2V-1.3B-Diffusers" # Critic model
|
||||
|
||||
DATA_DIR="data/crush-smol_processed_t2v/combined_parquet_dataset/"
|
||||
VALIDATION_DATASET_FILE="/mnt/weka/home/hao.zhang/wl/FastVideo/examples/distill/SFWan2.1-I2V/validation_better.json"
|
||||
# export CUDA_VISIBLE_DEVICES=4,5
|
||||
# IP=[MASTER NODE IP]
|
||||
|
||||
# Training arguments
|
||||
training_args=(
|
||||
--tracker_project_name SFwan_t2v_distill_self_forcing_dmd
|
||||
--output_dir "/mnt/sharefs/users/hao.zhang/wl/sf_checkpoints/ode0_SFwan_t2v_finetune_sf_${lr}_c${critic_lr}"
|
||||
--wandb_run_name "DEBUG${lr}_c${critic_lr}"
|
||||
--max_train_steps 4000
|
||||
--train_batch_size 1
|
||||
--train_sp_batch_size 1
|
||||
--gradient_accumulation_steps 1
|
||||
--num_latent_t 21
|
||||
--num_height 480
|
||||
--num_width 832
|
||||
--num_frames 81
|
||||
--warp_denoising_step
|
||||
--enable_gradient_checkpointing_type "full"
|
||||
--log_visualization
|
||||
--simulate_generator_forward
|
||||
--num_frame_per_block 3
|
||||
--enable_gradient_masking
|
||||
--gradient_mask_last_n_frames 21
|
||||
)
|
||||
|
||||
# Parallel arguments
|
||||
parallel_args=(
|
||||
--num_gpus 8 # 64
|
||||
--sp_size 1
|
||||
--tp_size 1
|
||||
--hsdp_replicate_dim 1 # 64
|
||||
--hsdp_shard_dim 8
|
||||
)
|
||||
|
||||
# Model arguments
|
||||
model_args=(
|
||||
--model_path $GENERATOR_MODEL_PATH # TODO: check if you can remove this in this script
|
||||
--pretrained_model_name_or_path $GENERATOR_MODEL_PATH
|
||||
--generator_model_path $GENERATOR_MODEL_PATH
|
||||
--real_score_model_path $REAL_SCORE_MODEL_PATH
|
||||
--fake_score_model_path $FAKE_SCORE_MODEL_PATH
|
||||
)
|
||||
|
||||
# Dataset arguments
|
||||
dataset_args=(
|
||||
--data_path "$DATA_DIR"
|
||||
--dataloader_num_workers 4
|
||||
)
|
||||
|
||||
# Validation arguments
|
||||
validation_args=(
|
||||
# --log_validation
|
||||
--validation_dataset_file "$VALIDATION_DATASET_FILE"
|
||||
--validation_steps 50
|
||||
--validation_sampling_steps "3"
|
||||
--validation_guidance_scale "6.0" # not used for dmd inference
|
||||
)
|
||||
|
||||
# Optimizer arguments
|
||||
optimizer_args=(
|
||||
--learning_rate 1e-5
|
||||
--mixed_precision "bf16"
|
||||
--training_state_checkpointing_steps 50
|
||||
--weight_only_checkpointing_steps 50
|
||||
--weight_decay 0.01
|
||||
--betas '0.0,0.999'
|
||||
--max_grad_norm 1.0
|
||||
)
|
||||
|
||||
# Miscellaneous arguments
|
||||
miscellaneous_args=(
|
||||
--inference_mode False
|
||||
--checkpoints_total_limit 3
|
||||
--training_cfg_rate 0.0
|
||||
--dit_precision "fp32"
|
||||
--flow_shift 5
|
||||
--seed 1000
|
||||
--use_ema True
|
||||
--ema_decay 0.99
|
||||
--ema_start_step 100
|
||||
--init_weights_from_safetensors "/mnt/weka/home/hao.zhang/wl/Self-Forcing/diffusers_ode_init/model.safetensors"
|
||||
)
|
||||
|
||||
# Self-forcing DMD arguments
|
||||
dmd_args=(
|
||||
--dmd_denoising_steps '1000,750,500,250'
|
||||
--min_timestep_ratio 0.02
|
||||
--max_timestep_ratio 0.98
|
||||
--dfake_gen_update_ratio 5
|
||||
--real_score_guidance_scale 3.0
|
||||
--fake_score_learning_rate 8e-6
|
||||
--fake_score_betas '0.0,0.999'
|
||||
)
|
||||
|
||||
# Self-forcing specific arguments
|
||||
self_forcing_args=(
|
||||
--independent_first_frame False # Whether to treat first frame independently
|
||||
--same_step_across_blocks False # Whether to use same denoising step across all blocks
|
||||
--last_step_only False # Whether to only use the last denoising step
|
||||
--context_noise 0 # Amount of noise to add during context caching (0 = no noise)
|
||||
--validate_cache_structure False # Set to True for debugging KV cache issues
|
||||
)
|
||||
|
||||
torchrun \
|
||||
--nnodes 1 \
|
||||
--nproc_per_node $NUM_GPUS \
|
||||
--master_port $MASTER_PORT \
|
||||
fastvideo/training/wan_self_forcing_distillation_pipeline.py \
|
||||
"${parallel_args[@]}" \
|
||||
"${model_args[@]}" \
|
||||
"${dataset_args[@]}" \
|
||||
"${training_args[@]}" \
|
||||
"${optimizer_args[@]}" \
|
||||
"${validation_args[@]}" \
|
||||
"${miscellaneous_args[@]}" \
|
||||
"${dmd_args[@]}" \
|
||||
"${self_forcing_args[@]}"
|
||||
@@ -57,7 +57,7 @@ class WanT2V480PConfig(PipelineConfig):
|
||||
# WanConfig-specific added parameters
|
||||
|
||||
def __post_init__(self):
|
||||
self.vae_config.load_encoder = True
|
||||
self.vae_config.load_encoder = False
|
||||
self.vae_config.load_decoder = True
|
||||
|
||||
|
||||
@@ -148,4 +148,4 @@ class SelfForcingWanT2V480PConfig(WanT2V480PConfig):
|
||||
is_causal: bool = True
|
||||
flow_shift: int = 5
|
||||
dmd_denoising_steps: list[int] | None = field(
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
default_factory=lambda: [1000, 750, 500, 250])
|
||||
|
||||
@@ -147,10 +147,8 @@ class CausalWanSelfAttention(nn.Module):
|
||||
# Assign new keys/values directly up to current_end
|
||||
local_end_index = kv_cache["local_end_index"].item() + current_end - kv_cache["global_end_index"].item()
|
||||
local_start_index = local_end_index - num_new_tokens
|
||||
|
||||
kv_cache["k"] = kv_cache["k"].clone()
|
||||
kv_cache["v"] = kv_cache["v"].clone()
|
||||
|
||||
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
|
||||
kv_cache["v"][:, local_start_index:local_end_index] = v
|
||||
x = self.attn(
|
||||
@@ -248,27 +246,27 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
current_start: int = 0,
|
||||
cache_start: int | None = None,
|
||||
) -> torch.Tensor:
|
||||
logger.info("temb.shape: %s", temb.shape)
|
||||
# logger.info("temb.shape: %s", temb.shape)
|
||||
num_frames = temb.shape[1]
|
||||
logger.info("first hidden_states.shape: %s", hidden_states.shape)
|
||||
logger.info("num_frames: %s", num_frames)
|
||||
# logger.info("first hidden_states.shape: %s", hidden_states.shape)
|
||||
# logger.info("num_frames: %s", num_frames)
|
||||
if hidden_states.dim() == 4:
|
||||
hidden_states = hidden_states.squeeze(1)
|
||||
frame_seqlen = hidden_states.shape[1] // temb.shape[1]
|
||||
logger.info("frame_seqlen: %s", frame_seqlen)
|
||||
# logger.info("frame_seqlen: %s", frame_seqlen)
|
||||
bs, seq_length, _ = hidden_states.shape
|
||||
orig_dtype = hidden_states.dtype
|
||||
# assert orig_dtype != torch.float32
|
||||
e = self.scale_shift_table + temb.float()
|
||||
logger.info("e.shape: %s", e.shape)
|
||||
# logger.info("e.shape: %s", e.shape)
|
||||
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
|
||||
6, dim=2)
|
||||
assert shift_msa.dtype == torch.float32
|
||||
|
||||
# 1. Self-attention
|
||||
logger.info("hidden_states.shape: %s", hidden_states.shape)
|
||||
logger.info("scale_msa.shape: %s", scale_msa.shape)
|
||||
logger.info("shift_msa.shape: %s", shift_msa.shape)
|
||||
# logger.info("hidden_states.shape: %s", hidden_states.shape)
|
||||
# logger.info("scale_msa.shape: %s", scale_msa.shape)
|
||||
# logger.info("shift_msa.shape: %s", shift_msa.shape)
|
||||
|
||||
norm_hidden_states_unflattened = self.norm1(hidden_states.float()).unflatten(dim=1, sizes=(num_frames, frame_seqlen))
|
||||
# logger.info("norm_hidden_states_unflattened.shape: %s", norm_hidden_states_unflattened.shape)
|
||||
@@ -320,7 +318,7 @@ class CausalWanTransformerBlock(nn.Module):
|
||||
ff_output = self.ffn(norm_hidden_states)
|
||||
hidden_states = self.mlp_residual(hidden_states, ff_output, c_gate_msa)
|
||||
# logger.info("after mlp_residual norm_hidden_states.shape: %s", norm_hidden_states.shape)
|
||||
logger.info("after mlp_residual hidden_states.shape: %s", hidden_states.shape)
|
||||
# logger.info("after mlp_residual hidden_states.shape: %s", hidden_states.shape)
|
||||
hidden_states = hidden_states.to(orig_dtype)
|
||||
|
||||
return hidden_states
|
||||
@@ -517,7 +515,7 @@ class CausalWanTransformer3DModel(BaseDiT):
|
||||
hidden_states = hidden_states.flatten(2).transpose(1, 2)
|
||||
# logger.info("forward inference flattened and transposed hidden_states.shape: %s", hidden_states.shape)
|
||||
|
||||
logger.info("timestep shape: %s", timestep.shape)
|
||||
# logger.info("timestep shape: %s", timestep.shape)
|
||||
|
||||
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
|
||||
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
|
||||
|
||||
@@ -635,15 +635,32 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
noise: torch.Tensor,
|
||||
timestep: torch.IntTensor,
|
||||
) -> torch.Tensor:
|
||||
"""
|
||||
Args:
|
||||
clean_latent: the clean latent with shape [B, C, H, W],
|
||||
where B is batch_size or batch_size * num_frames
|
||||
noise: the noise with shape [B, C, H, W]
|
||||
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
|
||||
|
||||
Returns:
|
||||
the corrupted latent with shape [B, C, H, W]
|
||||
"""
|
||||
# If timestep is [bs, num_frames]
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
elif timestep.ndim == 1:
|
||||
# If timestep is [1]
|
||||
if timestep.shape[0] == 1:
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
else:
|
||||
assert timestep.numel() == clean_latent.shape[0]
|
||||
else:
|
||||
raise ValueError(f"[add_noise] Invalid timestep shape: {timestep.shape}")
|
||||
# timestep shape should be [B]
|
||||
|
||||
self.sigmas = self.sigmas.to(noise.device)
|
||||
# TODO: hack
|
||||
if timestep.ndim == 2 and timestep.shape[1] == 1:
|
||||
timestep = timestep.expand(clean_latent.shape[0])
|
||||
self.timesteps = self.timesteps.to(noise.device)
|
||||
logger.info("self.timesteps shape: %s", self.timesteps.shape)
|
||||
logger.info("timestep shape: %s", timestep.shape)
|
||||
if timestep.ndim > 1:
|
||||
timestep = timestep.squeeze(0)
|
||||
timestep_id = torch.argmin(
|
||||
(self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1)
|
||||
sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1)
|
||||
@@ -656,4 +673,4 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
|
||||
return sample
|
||||
|
||||
def __len__(self) -> int:
|
||||
return self.config.num_train_timesteps
|
||||
return self.config.num_train_timesteps
|
||||
@@ -148,11 +148,30 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
|
||||
scheduler: Any) -> torch.Tensor:
|
||||
"""
|
||||
Convert predicted noise to clean latent.
|
||||
|
||||
Args:
|
||||
pred_noise: the predicted noise with shape [B, C, H, W]
|
||||
where B is batch_size or batch_size * num_frames
|
||||
noise_input_latent: the noisy latent with shape [B, C, H, W],
|
||||
timestep: the timestep with shape [1] or [bs * num_frames] or [bs, num_frames]
|
||||
scheduler: the scheduler
|
||||
|
||||
Returns:
|
||||
the predicted video with shape [B, C, H, W]
|
||||
"""
|
||||
logger.info(f"timestep: {timestep.shape}")
|
||||
logger.info(f"noise_input_latent: {noise_input_latent.shape}")
|
||||
logger.info(f"pred_noise: {pred_noise.shape}")
|
||||
timestep = timestep.expand(noise_input_latent.shape[0])
|
||||
# If timestep is [bs, num_frames]
|
||||
if timestep.ndim == 2:
|
||||
timestep = timestep.flatten(0, 1)
|
||||
assert timestep.numel() == noise_input_latent.shape[0]
|
||||
elif timestep.ndim == 1:
|
||||
# If timestep is [1]
|
||||
if timestep.shape[0] == 1:
|
||||
timestep = timestep.expand(noise_input_latent.shape[0])
|
||||
else:
|
||||
assert timestep.numel() == noise_input_latent.shape[0]
|
||||
else:
|
||||
raise ValueError(f"[pred_noise_to_pred_video] Invalid timestep shape: {timestep.shape}")
|
||||
# timestep shape should be [B]
|
||||
dtype = pred_noise.dtype
|
||||
device = pred_noise.device
|
||||
pred_noise = pred_noise.float().to(device)
|
||||
|
||||
@@ -27,7 +27,6 @@ from fastvideo.layers.activation import get_act_fn
|
||||
from fastvideo.models.vaes.common import (DiagonalGaussianDistribution,
|
||||
ParallelTiledVAE)
|
||||
from fastvideo.platforms import current_platform
|
||||
from fastvideo.logger import init_logger
|
||||
|
||||
CACHE_T = 2
|
||||
|
||||
@@ -36,7 +35,6 @@ feat_cache = contextvars.ContextVar("feat_cache", default=None)
|
||||
feat_idx = contextvars.ContextVar("feat_idx", default=0)
|
||||
first_chunk = contextvars.ContextVar("first_chunk", default=None)
|
||||
|
||||
logger = init_logger(__name__)
|
||||
|
||||
@contextmanager
|
||||
def forward_context(first_frame_arg=False,
|
||||
@@ -1131,7 +1129,6 @@ class AutoencoderKLWan(nn.Module, ParallelTiledVAE):
|
||||
self._conv_idx = 0
|
||||
self._feat_map = [None] * self._conv_num
|
||||
# cache encode
|
||||
logger.info("self.config.load_encoder: %s", self.config.load_encoder)
|
||||
if self.config.load_encoder:
|
||||
self._enc_conv_num = _count_conv3d(self.encoder)
|
||||
self._enc_conv_idx = 0
|
||||
|
||||
@@ -12,6 +12,7 @@ from typing import Any
|
||||
import imageio
|
||||
import numpy as np
|
||||
import torch
|
||||
import torch.distributed as dist
|
||||
import torch.nn.functional as F
|
||||
import torchvision
|
||||
from einops import rearrange
|
||||
@@ -728,16 +729,17 @@ class DistillationPipeline(TrainingPipeline):
|
||||
"encoder_hidden_states": training_batch.encoder_hidden_states,
|
||||
"encoder_attention_mask": training_batch.encoder_attention_mask,
|
||||
}
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
if getattr(self, "negative_prompt_embeds", None) is not None:
|
||||
unconditional_dict = {
|
||||
"encoder_hidden_states": self.negative_prompt_embeds,
|
||||
"encoder_attention_mask": self.negative_prompt_attention_mask,
|
||||
}
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
|
||||
training_batch.dmd_latent_vis_dict = {}
|
||||
training_batch.fake_score_latent_vis_dict = {}
|
||||
|
||||
training_batch.conditional_dict = conditional_dict
|
||||
training_batch.unconditional_dict = unconditional_dict
|
||||
training_batch.raw_latent_shape = training_batch.latents.shape
|
||||
training_batch.latents = training_batch.latents.permute(0, 2, 1, 3, 4)
|
||||
self.video_latent_shape = training_batch.latents.shape
|
||||
|
||||
@@ -52,14 +52,17 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
super().initialize_training_pipeline(training_args)
|
||||
self.dfake_gen_update_ratio = getattr(training_args, 'dfake_gen_update_ratio', 5)
|
||||
|
||||
# Self-forcing specific properties
|
||||
self.num_frame_per_block = getattr(training_args, 'num_frame_per_block', 3)
|
||||
self.independent_first_frame = getattr(training_args, 'independent_first_frame', False)
|
||||
self.same_step_across_blocks = getattr(training_args, 'same_step_across_blocks', False)
|
||||
self.last_step_only = getattr(training_args, 'last_step_only', False)
|
||||
self.context_noise = getattr(training_args, 'context_noise', 0)
|
||||
|
||||
# Calculate frame sequence length - this will be set properly in _prepare_dit_inputs
|
||||
self.frame_seq_length = 1560 # TODO: Calculate this dynamically based on patch size
|
||||
|
||||
# Cache references (will be initialized per forward pass)
|
||||
self.kv_cache1 = None
|
||||
self.crossattn_cache = None
|
||||
|
||||
@@ -155,9 +158,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
noise = noise[:, self.rank_in_sp_group, :, :, :, :]
|
||||
noisy_latent = self.noise_scheduler.add_noise(latents.flatten(0, 1),
|
||||
noise.flatten(0, 1),
|
||||
timestep).unflatten(
|
||||
0,
|
||||
(1, latents.shape[1]))
|
||||
torch.tensor([timestep], device=noise.device))
|
||||
|
||||
# Step 4: Build input kwargs with KV cache support
|
||||
training_batch = self._build_distill_input_kwargs(
|
||||
@@ -186,7 +187,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
pred_video = pred_noise_to_pred_video(
|
||||
pred_noise=pred_noise.flatten(0, 1),
|
||||
noise_input_latent=noisy_latent.flatten(0, 1),
|
||||
timestep=timestep,
|
||||
timestep=torch.tensor([timestep], device=noisy_latent.device),
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_noise.shape[:2])
|
||||
|
||||
self._reset_simulation_caches(kv_cache, crossattn_cache)
|
||||
@@ -206,7 +207,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
initial_latent = getattr(training_batch, 'image_latent', None)
|
||||
|
||||
# Dynamic frame generation logic (adapted from _run_generator)
|
||||
num_training_frames = getattr(self.training_args, 'num_frames', 21)
|
||||
num_training_frames = getattr(self.training_args, 'num_latent_t', 21)
|
||||
|
||||
# During training, the number of generated frames should be uniformly sampled from
|
||||
# [21, self.num_training_frames], but still being a multiple of self.num_frame_per_block
|
||||
@@ -287,7 +288,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
all_num_frames = [1] + all_num_frames
|
||||
num_denoising_steps = len(self.denoising_step_list)
|
||||
exit_flags = self.generate_and_sync_list(len(all_num_frames), num_denoising_steps, device=noise.device)
|
||||
start_gradient_frame_index = num_output_frames - 21
|
||||
start_gradient_frame_index = max(0, num_output_frames - 21)
|
||||
|
||||
for block_index, current_num_frames in enumerate(all_num_frames):
|
||||
noisy_input = noise[:, current_start_frame - num_input_frames:current_start_frame + current_num_frames - num_input_frames]
|
||||
@@ -300,7 +301,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
exit_flag = (index == exit_flags[block_index])
|
||||
|
||||
timestep = torch.ones([batch_size, current_num_frames], device=noise.device, dtype=torch.int64) * current_timestep
|
||||
logger.info("timestep shape at initalization: %s", timestep.shape)
|
||||
|
||||
if not exit_flag:
|
||||
with torch.no_grad():
|
||||
@@ -308,7 +308,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
noisy_input, timestep, training_batch.conditional_dict, training_batch)
|
||||
|
||||
logger.info("timestep shape in generator: %s", timestep.shape)
|
||||
pred_flow = self.transformer(
|
||||
hidden_states=training_batch_temp.input_kwargs['hidden_states'],
|
||||
encoder_hidden_states=training_batch_temp.input_kwargs['encoder_hidden_states'],
|
||||
@@ -322,21 +321,18 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep.flatten(),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
|
||||
|
||||
next_timestep = self.denoising_step_list[index + 1]
|
||||
logger.info("denoised_pred shape before: %s", denoised_pred.shape)
|
||||
noisy_input = self.noise_scheduler.add_noise(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
next_timestep * torch.ones([batch_size * current_num_frames], device=noise.device, dtype=torch.long)
|
||||
).unflatten(0, denoised_pred.shape[:2])
|
||||
logger.info("denoised_pred shape after: %s", denoised_pred.shape)
|
||||
else:
|
||||
# Final prediction with gradient control
|
||||
if current_start_frame < start_gradient_frame_index:
|
||||
logger.info("timestep shape in generator: %s", timestep.shape)
|
||||
with torch.no_grad():
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
noisy_input, timestep, training_batch.conditional_dict, training_batch)
|
||||
@@ -351,7 +347,6 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
current_start=current_start_frame * self.frame_seq_length
|
||||
).permute(0, 2, 1, 3, 4)
|
||||
else:
|
||||
logger.info("timestep shape in generator: %s", timestep.shape)
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
noisy_input, timestep, training_batch.conditional_dict, training_batch)
|
||||
|
||||
@@ -368,7 +363,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
denoised_pred = pred_noise_to_pred_video(
|
||||
pred_noise=pred_flow.flatten(0, 1),
|
||||
noise_input_latent=noisy_input.flatten(0, 1),
|
||||
timestep=timestep.flatten(),
|
||||
timestep=timestep,
|
||||
scheduler=self.noise_scheduler).unflatten(0, pred_flow.shape[:2])
|
||||
break
|
||||
|
||||
@@ -377,13 +372,11 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
|
||||
# Step 3.3: rerun with timestep zero to update the cache
|
||||
context_timestep = torch.ones_like(timestep) * self.context_noise
|
||||
logger.info("denoised_pred shape before: %s", denoised_pred.shape)
|
||||
denoised_pred = self.noise_scheduler.add_noise(
|
||||
denoised_pred.flatten(0, 1),
|
||||
torch.randn_like(denoised_pred.flatten(0, 1)),
|
||||
context_timestep * torch.ones([batch_size * current_num_frames], device=noise.device, dtype=torch.long)
|
||||
context_timestep
|
||||
).unflatten(0, denoised_pred.shape[:2])
|
||||
logger.info("denoised_pred shape after: %s", denoised_pred.shape)
|
||||
|
||||
with torch.no_grad():
|
||||
training_batch_temp = self._build_distill_input_kwargs(
|
||||
@@ -433,7 +426,7 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
|
||||
frame = pixels[:, :, -1:, :, :].to(dtype) # Last frame [B, C, 1, H, W]
|
||||
|
||||
# Encode frame back to get image latent
|
||||
image_latent = self.vae.encode(frame).mean.to(dtype)
|
||||
image_latent = self.vae.encode(frame).to(dtype)
|
||||
image_latent = image_latent.permute(0, 2, 1, 3, 4) # [B, F, C, H, W]
|
||||
|
||||
pred_image_or_video_last_21 = torch.cat([image_latent, pred_image_or_video[:, -20:, ...]], dim=1)
|
||||
|
||||
+2
-1
@@ -2,7 +2,8 @@
|
||||
|
||||
counter=0
|
||||
# for pair in "1e-5 1e-5" "1e-5 8e-6" "1e-5 6e-6" "1e-5 4e-6" "1e-5 2e-6" "1e-5 1e-6"; do
|
||||
for pair in "1e-5 8e-6" "1e-5 4e-6" "1e-5 2e-6" "8e-6 8e-6" "8e-6 6e-6" "8e-6 2e-6"; do
|
||||
# for pair in "1e-5 8e-6" "1e-5 4e-6" "1e-5 2e-6" "8e-6 8e-6" "8e-6 6e-6" "8e-6 2e-6"; do
|
||||
for pair in "1e-5 8e-6"; do
|
||||
# for pair in "2e-6 4e-7" "2e-6 6e-7" "4e-6 6e-7" "4e-6 8e-7" "4e-6 1e-6" "6e-6 4e-7" "6e-6 6e-7" "6e-6 8e-7" "6e-6 1e-6"; do
|
||||
# for pair in "2e-6 4e-7" "2e-6 6e-7" "4e-6 6e-7" "4e-6 8e-7" "4e-6 1e-6" "6e-6 4e-7" "6e-6 6e-7" "6e-6 8e-7" "6e-6 1e-6"; do
|
||||
port=$((29500 + counter))
|
||||
|
||||
Reference in New Issue
Block a user