Compare commits

...
Author SHA1 Message Date
Matthew Noto 4477dbea67 training refactor done 2025-09-08 03:48:01 +00:00
Matthew Noto 04685a3ecd kv cache debugging 2025-09-08 02:27:19 +00:00
Matthew Noto b8392e9b2a revert 2025-09-07 23:38:02 +00:00
Matthew Noto 60b71c5053 new branch 2025-09-07 23:36:43 +00:00
JerryZhou54 2b7bf88a4d new branch 2025-09-07 08:49:36 +00:00
10 changed files with 210 additions and 40 deletions
+2 -2
View File
@@ -57,7 +57,7 @@ class WanT2V480PConfig(PipelineConfig):
# WanConfig-specific added parameters
def __post_init__(self):
self.vae_config.load_encoder = False
self.vae_config.load_encoder = True
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])
+47 -6
View File
@@ -9,6 +9,9 @@ import torch.nn.functional as F
from fastvideo.layers.custom_op import CustomOp
from fastvideo.platforms import current_platform
from fastvideo.logger import init_logger
logger = init_logger(__name__)
@CustomOp.register("rms_norm")
class RMSNorm(CustomOp):
@@ -100,7 +103,13 @@ class ScaleResidual(nn.Module):
def forward(self, residual: torch.Tensor, x: torch.Tensor,
gate: torch.Tensor) -> torch.Tensor:
"""Apply gated residual connection."""
return residual + x * gate
# logger.info("x.shape: %s", x.shape)
# if isinstance(gate, torch.Tensor):
# logger.info("gate.shape: %s", gate.shape)
num_frames = gate.shape[1]
frame_seqlen = x.shape[1] // num_frames
return residual + (x.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * gate).flatten(1, 2)
# adapted from Diffusers: https://github.com/huggingface/diffusers/blob/main/src/diffusers/models/normalization.py
@@ -172,11 +181,35 @@ class ScaleResidualLayerNormScaleShift(nn.Module):
but before normalization)
"""
# Apply residual connection with gating
residual_output = residual + x * gate
# logger.info("x.shape: %s", x.shape)
if isinstance(gate, int):
# used by cross-attention, should be 1
assert gate == 1
residual_output = residual + x * gate
elif isinstance(gate, torch.Tensor):
# logger.info("gate.shape: %s", gate.shape)
if gate.dim() == 3:
# used by bidirectional self attention
residual_output = residual + x * gate
else:
assert gate.dim() == 4
num_frames = gate.shape[1]
frame_seqlen = x.shape[1] // num_frames
residual_output = residual + (x.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * gate).flatten(1, 2)
# residual_output = residual + x * gate
else:
raise ValueError(f"Gate type {type(gate)} not supported")
# logger.info("residual_output.shape: %s", residual_output.shape)
# Apply normalization
normalized = self.norm(residual_output)
# Apply scale and shift
modulated = normalized * (1.0 + scale) + shift
if isinstance(scale, torch.Tensor) and scale.dim() == 4:
num_frames = scale.shape[1]
frame_seqlen = normalized.shape[1] // num_frames
modulated = (normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale) + shift).flatten(1, 2)
else:
modulated = normalized * (1.0 + scale) + shift
return modulated, residual_output
@@ -219,7 +252,15 @@ class LayerNormScaleShift(nn.Module):
scale: torch.Tensor) -> torch.Tensor:
"""Apply ln followed by scale and shift in a single fused operation."""
normalized = self.norm(x)
if self.compute_dtype == torch.float32:
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
if scale.dim() == 4:
num_frames = scale.shape[1]
frame_seqlen = normalized.shape[1] // num_frames
if self.compute_dtype == torch.float32:
return (normalized.float().unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale) + shift).flatten(1, 2).to(x.dtype)
else:
return (normalized.unflatten(dim=1, sizes=(num_frames, frame_seqlen)) * (1.0 + scale) + shift).flatten(1, 2)
else:
return normalized * (1.0 + scale) + shift
if self.compute_dtype == torch.float32:
return (normalized.float() * (1.0 + scale) + shift).to(x.dtype)
else:
return normalized * (1.0 + scale) + shift
+6 -2
View File
@@ -86,12 +86,16 @@ class TimestepEmbedder(nn.Module):
dtype=dtype)
self.freq_dtype = freq_dtype
def forward(self, t: torch.Tensor) -> torch.Tensor:
def forward(self,
t: torch.Tensor,
timestep_seq_len: int | None = None) -> torch.Tensor:
t_freq = timestep_embedding(t,
self.frequency_embedding_size,
self.max_period,
dtype=self.freq_dtype).to(
self.mlp.fc_in.weight.dtype)
if timestep_seq_len is not None:
t_freq = t_freq.unflatten(0, (1, timestep_seq_len))
# t_freq = t_freq.to(self.mlp.fc_in.weight.dtype)
t_emb = self.mlp(t_freq)
return t_emb
@@ -172,4 +176,4 @@ def unpatchify(x, t, h, w, patch_size, channels) -> torch.Tensor:
x = torch.einsum("nthwcopq->nctohpwq", x)
imgs = x.reshape(shape=(x.shape[0], c, t * pt, h * ph, w * pw))
return imgs
return imgs
+78 -14
View File
@@ -147,6 +147,10 @@ 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(
@@ -244,19 +248,36 @@ class CausalWanTransformerBlock(nn.Module):
current_start: int = 0,
cache_start: int | None = None,
) -> torch.Tensor:
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)
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)
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)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
6, dim=2)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
norm_hidden_states = (self.norm1(hidden_states.float()) *
(1 + scale_msa) + shift_msa).to(orig_dtype)
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)
# norm_hidden_states = (self.norm1(hidden_states.float()) *
# (1 + scale_msa) + shift_msa).to(orig_dtype)
norm_hidden_states = (norm_hidden_states_unflattened *
(1 + scale_msa) + shift_msa).flatten(1, 2).to(orig_dtype)
# logger.info("1 norm_hidden_states.shape: %s", norm_hidden_states.shape)
query, _ = self.to_q(norm_hidden_states)
key, _ = self.to_k(norm_hidden_states)
value, _ = self.to_v(norm_hidden_states)
@@ -278,6 +299,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)
# logger.info("after self_attn_residual_norm norm_hidden_states.shape: %s", norm_hidden_states.shape)
# logger.info("after self_attn_residual_norm hidden_states.shape: %s", hidden_states.shape)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
@@ -288,12 +311,16 @@ class CausalWanTransformerBlock(nn.Module):
crossattn_cache=crossattn_cache)
norm_hidden_states, hidden_states = self.cross_attn_residual_norm(
hidden_states, attn_output, 1, c_shift_msa, c_scale_msa)
# logger.info("after cross_attn_residual_norm norm_hidden_states.shape: %s", norm_hidden_states.shape)
# logger.info("after cross_attn_residual_norm hidden_states.shape: %s", hidden_states.shape)
norm_hidden_states, hidden_states = norm_hidden_states.to(
orig_dtype), hidden_states.to(orig_dtype)
# 3. Feed-forward
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)
hidden_states = hidden_states.to(orig_dtype)
return hidden_states
@@ -359,8 +386,10 @@ class CausalWanTransformer3DModel(BaseDiT):
elementwise_affine=False,
dtype=torch.float32,
compute_dtype=torch.float32)
self.proj_out = nn.Linear(
inner_dim, config.out_channels * math.prod(config.patch_size))
# Debug: Log configuration values
proj_out_dim = config.out_channels * math.prod(config.patch_size)
self.proj_out = nn.Linear(inner_dim, proj_out_dim)
self.scale_shift_table = nn.Parameter(
torch.randn(1, 2, inner_dim) / inner_dim**0.5)
@@ -449,6 +478,7 @@ class CausalWanTransformer3DModel(BaseDiT):
This function will be run for num_frame times.
Process the latent frames one by one (1560 tokens each)
"""
# logger.info("forward inference hidden_states.shape: %s", hidden_states.shape)
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
@@ -485,10 +515,13 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.patch_embedding(hidden_states)
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)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
@@ -526,8 +559,14 @@ class CausalWanTransformer3DModel(BaseDiT):
**causal_kwargs)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
# logger.info("===== INFERENCE 5. Output norm, projection & unpatchify")
# logger.info("hidden_states.shape: %s", hidden_states.shape)
# logger.info("temb.shape: %s", temb.shape)
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
# logger.info("WTFWTF train temb.shape: %s", temb.shape)
# logger.info("WTFWTF train self.scale_shift_table.shape: %s", self.scale_shift_table.shape)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
dim=2)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
@@ -549,6 +588,8 @@ class CausalWanTransformer3DModel(BaseDiT):
start_frame: int = 0,
**kwargs) -> torch.Tensor:
# logger.info("===== forward train hidden_states.shape: %s", hidden_states.shape)
# logger.info("===== forward train timestep.shape: %s", timestep.shape)
orig_dtype = hidden_states.dtype
if not isinstance(encoder_hidden_states, torch.Tensor):
encoder_hidden_states = encoder_hidden_states[0]
@@ -594,10 +635,14 @@ class CausalWanTransformer3DModel(BaseDiT):
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# logger.info("forward train flattened and transposed hidden_states.shape: %s", hidden_states.shape)
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
timestep.flatten(), encoder_hidden_states, encoder_hidden_states_image)
# logger.info("forward train timestep_proj.shape: %s", timestep_proj.shape)
# logger.info("forward train timestep.shape: %s", timestep.shape)
# logger.info("forward train temb.shape: %s", temb.shape)
timestep_proj = timestep_proj.unflatten(1, (6, self.hidden_size)).unflatten(dim=0, sizes=timestep.shape)
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
@@ -617,16 +662,35 @@ class CausalWanTransformer3DModel(BaseDiT):
timestep_proj, freqs_cis,
block_mask=self.block_mask)
else:
for block in self.blocks:
for block_index, block in enumerate(self.blocks):
logger.info("===== TRAIN block %d", block_index)
logger.info("hidden_states.shape: %s", hidden_states.shape)
# logger.info("encoder_hidden_states.shape: %s", encoder_hidden_states.shape)
logger.info("timestep_proj.shape: %s", timestep_proj.shape)
# logger.info("freqs_cis.shape: %s", freqs_cis.shape)
# logger.info("block_mask.shape: %s", self.block_mask.shape)
hidden_states = block(hidden_states, encoder_hidden_states,
timestep_proj, freqs_cis,
block_mask=self.block_mask)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
# logger.info("===== TRAIN 5. Output norm, projection & unpatchify")
# logger.info("hidden_states.shape: %s", hidden_states.shape)
# logger.info("temb.shape: %s", temb.shape)
# shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
temb = temb.unflatten(dim=0, sizes=timestep.shape).unsqueeze(2)
# logger.info("WTFWTF train temb.shape: %s", temb.shape)
# logger.info("WTFWTF train self.scale_shift_table.shape: %s", self.scale_shift_table.shape)
shift, scale = (self.scale_shift_table.unsqueeze(1) + temb).chunk(2,
dim=2)
# logger.info("DEBUG scale.shape: %s", scale.shape)
# logger.info("DEBUG shift.shape: %s", shift.shape)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
# logger.info("DEBUG after proj_out hidden_states.shape: %s", hidden_states.shape)
# logger.info(f"DEBUG reshape dimensions: batch_size={batch_size}, post_patch_num_frames={post_patch_num_frames}")
# logger.info(f"DEBUG reshape dimensions: post_patch_height={post_patch_height}, post_patch_width={post_patch_width}")
# logger.info(f"DEBUG patch dimensions: p_t={p_t}, p_h={p_h}, p_w={p_w}")
hidden_states = hidden_states.reshape(batch_size, post_patch_num_frames,
post_patch_height,
+44 -8
View File
@@ -81,8 +81,9 @@ class WanTimeTextImageEmbedding(nn.Module):
timestep: torch.Tensor,
encoder_hidden_states: torch.Tensor,
encoder_hidden_states_image: torch.Tensor | None = None,
timestep_seq_len: int | None = None,
):
temb = self.time_embedder(timestep)
temb = self.time_embedder(timestep, timestep_seq_len)
timestep_proj = self.time_modulation(temb)
encoder_hidden_states = self.text_embedder(encoder_hidden_states)
@@ -319,9 +320,24 @@ class WanTransformerBlock(nn.Module):
bs, seq_length, _ = hidden_states.shape
orig_dtype = hidden_states.dtype
# assert orig_dtype != torch.float32
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
if temb.dim() == 4:
# temb: batch_size, seq_len, 6, inner_dim (wan2.2 ti2v)
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = (
self.scale_shift_table.unsqueeze(0) + temb.float()
).chunk(6, dim=2)
# batch_size, seq_len, 1, inner_dim
shift_msa = shift_msa.squeeze(2)
scale_msa = scale_msa.squeeze(2)
gate_msa = gate_msa.squeeze(2)
c_shift_msa = c_shift_msa.squeeze(2)
c_scale_msa = c_scale_msa.squeeze(2)
c_gate_msa = c_gate_msa.squeeze(2)
else:
# temb: batch_size, 6, inner_dim (wan2.1/wan2.2 14B)
e = self.scale_shift_table + temb.float()
shift_msa, scale_msa, gate_msa, c_shift_msa, c_scale_msa, c_gate_msa = e.chunk(
6, dim=1)
assert shift_msa.dtype == torch.float32
# 1. Self-attention
@@ -649,9 +665,21 @@ class WanTransformer3DModel(CachableDiT):
hidden_states = self.patch_embedding(hidden_states)
hidden_states = hidden_states.flatten(2).transpose(1, 2)
# timestep shape: batch_size, or batch_size, seq_len (wan 2.2 ti2v)
if timestep.dim() == 2:
ts_seq_len = timestep.shape[1]
timestep = timestep.flatten() # batch_size * seq_len
else:
ts_seq_len = None
temb, timestep_proj, encoder_hidden_states, encoder_hidden_states_image = self.condition_embedder(
timestep, encoder_hidden_states, encoder_hidden_states_image)
timestep_proj = timestep_proj.unflatten(1, (6, -1))
timestep, encoder_hidden_states, encoder_hidden_states_image, timestep_seq_len=ts_seq_len)
if ts_seq_len is not None:
# batch_size, seq_len, 6, inner_dim
timestep_proj = timestep_proj.unflatten(2, (6, -1))
else:
# batch_size, 6, inner_dim
timestep_proj = timestep_proj.unflatten(1, (6, -1))
if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat(
@@ -688,8 +716,15 @@ class WanTransformer3DModel(CachableDiT):
if enable_teacache:
self.maybe_cache_states(hidden_states, original_hidden_states)
# 5. Output norm, projection & unpatchify
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2,
dim=1)
if temb.dim() == 3:
# batch_size, seq_len, inner_dim (wan 2.2 ti2v)
shift, scale = (self.scale_shift_table.unsqueeze(0) + temb.unsqueeze(2)).chunk(2, dim=2)
shift = shift.squeeze(2)
scale = scale.squeeze(2)
else:
# batch_size, inner_dim
shift, scale = (self.scale_shift_table + temb.unsqueeze(1)).chunk(2, dim=1)
hidden_states = self.norm_out(hidden_states, shift, scale)
hidden_states = self.proj_out(hidden_states)
@@ -793,3 +828,4 @@ class WanTransformer3DModel(CachableDiT):
return hidden_states + self.previous_residual_even
else:
return hidden_states + self.previous_residual_odd
@@ -636,8 +636,14 @@ class FlowMatchEulerDiscreteScheduler(SchedulerMixin, ConfigMixin,
timestep: torch.IntTensor,
) -> torch.Tensor:
self.sigmas = self.sigmas.to(noise.device)
timestep = timestep.expand(clean_latent.shape[0])
# 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)
+6
View File
@@ -6,6 +6,9 @@ from typing import Any
import torch
from fastvideo.logger import init_logger
logger = init_logger(__name__)
# TODO(PY): move it elsewhere
def auto_attributes(init_func):
"""
@@ -146,6 +149,9 @@ def pred_noise_to_pred_video(pred_noise: torch.Tensor,
"""
Convert predicted noise to clean latent.
"""
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])
dtype = pred_noise.dtype
device = pred_noise.device
+3
View File
@@ -27,6 +27,7 @@ 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
@@ -35,6 +36,7 @@ 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,
@@ -1129,6 +1131,7 @@ 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
@@ -228,7 +228,10 @@ class CausalDMDDenosingStage(DenoisingStage):
dim=2)
# Prepare inputs
t_expand = t_cur.repeat(latent_model_input.shape[0])
t_expand = t_cur.expand(latent_model_input.shape[0])
# t_expand = t_cur * torch.ones((latent_model_input.shape[0], 1), device=latent_model_input.device, dtype=torch.long)
# t_expand = t_expand.repeat(1, self.sliding_window_num_frames)
# Attention metadata if needed
if (vsa_available and self.attn_backend
@@ -262,10 +265,11 @@ class CausalDMDDenosingStage(DenoisingStage):
attn_metadata=attn_metadata,
forward_batch=batch):
# Run transformer; follow DMD stage pattern
t_expanded_noise= t_cur * torch.ones((latent_model_input.shape[0], 1), device=latent_model_input.device, dtype=torch.long)
pred_noise_btchw = self.transformer(
latent_model_input,
prompt_embeds,
t_expand,
t_expanded_noise,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
@@ -326,10 +330,11 @@ class CausalDMDDenosingStage(DenoisingStage):
set_forward_context(current_timestep=0,
attn_metadata=attn_metadata,
forward_batch=batch):
t_expanded_context = t_context * torch.ones((context_bcthw.shape[0], 1), device=context_bcthw.device, dtype=torch.long)
_ = self.transformer(
context_bcthw,
prompt_embeds,
t_context,
t_expanded_context,
kv_cache=self.kv_cache1,
crossattn_cache=self.crossattn_cache,
current_start=(pos_start_base + start_index) *
@@ -52,17 +52,14 @@ 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
@@ -303,6 +300,7 @@ 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():
@@ -310,6 +308,7 @@ 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'],
@@ -327,14 +326,17 @@ class SelfForcingDistillationPipeline(DistillationPipeline):
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)
@@ -349,6 +351,7 @@ 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)
@@ -374,11 +377,13 @@ 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)
).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(
@@ -428,7 +433,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).to(dtype)
image_latent = self.vae.encode(frame).mean.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)