Compare commits

...
Author SHA1 Message Date
SolitaryThinker 2dce62c034 update 2025-11-24 00:58:19 +00:00
SolitaryThinker 4888a31750 cleanup 2025-11-23 06:40:58 +00:00
SolitaryThinker fd899c16f9 update 2025-11-23 06:16:23 +00:00
SolitaryThinker a46f5ea380 working 2025-11-23 05:33:37 +00:00
SolitaryThinker 1f1974bfa2 wip 2025-11-22 02:57:23 +00:00
SolitaryThinker c147f97d3c lint 2025-11-21 00:18:15 +00:00
SolitaryThinker 10384b1ad3 wip 2025-11-21 00:17:24 +00:00
7 changed files with 208 additions and 72 deletions
@@ -12,20 +12,34 @@ def main():
generator = VideoGenerator.from_pretrained(
model_name,
# FastVideo will automatically handle distributed setup
num_gpus=1,
num_gpus=4,
use_fsdp_inference=True,
text_encoder_cpu_offload=False,
dit_cpu_offload=False,
num_frame_per_block=4,
)
sampling_param = SamplingParam.from_pretrained(model_name)
# sampling_param.num_frames = 13
sampling_param.num_frames = 77
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
# return
import time
start_time = time.perf_counter()
for i in range(10):
prompt = (
"A curious raccoon peers through a vibrant field of yellow sunflowers, its eyes "
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
video = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, sampling_param=sampling_param)
end_time = time.perf_counter()
print(f"Time taken to generate 10 videos: {end_time - start_time} seconds, average time per video: {(end_time - start_time) / 10} seconds")
if __name__ == "__main__":
main()
@@ -13,20 +13,21 @@ def main():
generator = VideoGenerator.from_pretrained(
"FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers",
# FastVideo will automatically handle distributed setup
num_gpus=1,
num_gpus=4,
use_fsdp_inference=True,
dit_cpu_offload=True, # DiT need to be offloaded for MoE
dit_cpu_offload=False, # DiT need to be offloaded for MoE
dit_precision="fp32",
vae_cpu_offload=False,
text_encoder_cpu_offload=True,
dmd_denoising_steps=[1000, 850, 700, 550, 350, 275, 200, 125],
# Set pin_cpu_memory to false if CPU RAM is limited and there're no frequent CPU-GPU transfer
pin_cpu_memory=True,
num_frame_per_block=4,
# image_encoder_cpu_offload=False,
)
sampling_param = SamplingParam.from_pretrained("FastVideo/SFWan2.2-I2V-A14B-Preview-Diffusers")
sampling_param.num_frames = 81
sampling_param.num_frames = 77
sampling_param.width = 832
sampling_param.height = 480
sampling_param.seed = 1000
@@ -1,6 +1,7 @@
# NOTE: This is still a work in progress, and the checkpoints are not released yet.
from fastvideo import VideoGenerator
import time
# from fastvideo.configs.sample import SamplingParam
@@ -10,6 +11,7 @@ def main():
# model.
# If a local path is provided, FastVideo will make a best effort
# attempt to identify the optimal arguments.
load_start_time = time.perf_counter()
generator = VideoGenerator.from_pretrained(
"rand0nmr/SFWan2.2-T2V-A14B-Diffusers",
# FastVideo will automatically handle distributed setup
@@ -23,9 +25,10 @@ def main():
pin_cpu_memory=True,
init_weights_from_safetensors="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_inference_transformer/",
init_weights_from_safetensors_2="/mnt/sharefs/users/hao.zhang/wei/SFwan2.2_distill_self_forcing_release_cfg2/checkpoint-246_weight_only/generator_2_inference_transformer/",
num_frame_per_block=7,
# image_encoder_cpu_offload=False,
num_frame_per_block=4,
)
load_end_time = time.perf_counter()
load_time = load_end_time - load_start_time
# sampling_param = SamplingParam.from_pretrained("Wan-AI/Wan2.1-T2V-1.3B-Diffusers")
# sampling_param.num_frames = 45
@@ -36,8 +39,16 @@ def main():
"wide with interest. The playful yet serene atmosphere is complemented by soft "
"natural light filtering through the petals. Mid-shot, warm and cheerful tones."
)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=True, num_frames=81)
_ = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=False, num_frames=77)
# return
start_time = time.perf_counter()
for _ in range(10):
video2 = generator.generate_video(prompt, output_path=OUTPUT_PATH, save_video=False, num_frames=77)
end_time = time.perf_counter()
gen_time = (end_time - start_time) / 10
print(f"Time taken to load model: {load_time} seconds")
print(f"Time taken to generate 10 videos: {gen_time} seconds")
if __name__ == "__main__":
main()
+41 -7
View File
@@ -11,6 +11,9 @@ from fastvideo.distributed.parallel_state import (get_sp_parallel_rank,
from fastvideo.forward_context import ForwardContext, get_forward_context
from fastvideo.platforms import AttentionBackendEnum
from fastvideo.utils import get_compute_dtype
from fastvideo.logger import init_logger
logger = init_logger(__name__)
class DistributedAttention(nn.Module):
@@ -64,7 +67,11 @@ class DistributedAttention(nn.Module):
replicated_q: torch.Tensor | None = None,
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
k_cache: torch.Tensor | None = None,
v_cache: torch.Tensor | None = None,
return_current_kv: bool = False,
) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None,
torch.Tensor | None]:
"""Forward pass for distributed attention.
Args:
@@ -97,6 +104,7 @@ class DistributedAttention(nn.Module):
qkv = sequence_model_parallel_all_to_all_4D(qkv,
scatter_dim=2,
gather_dim=1)
# Apply backend-specific preprocess_qkv
qkv = self.attn_impl.preprocess_qkv(qkv, ctx_attn_metadata)
@@ -112,9 +120,22 @@ class DistributedAttention(nn.Module):
heads_per_rank]
qkv = torch.cat([qkv, replicated_qkv], dim=1)
q, k, v = qkv.chunk(3, dim=0)
q, k_new, v_new = qkv.chunk(3, dim=0)
k_total = k_new
v_total = v_new
output = self.attn_impl.forward(q, k, v, ctx_attn_metadata)
if k_cache is not None or v_cache is not None:
# assert False, "KV cache is not supported"
assert k_cache is not None and v_cache is not None
assert k_cache.shape == v_cache.shape
assert k_cache.shape[2] == v_cache.shape[
2], "Number of heads must be the same"
if k_cache.shape[1] > 0:
k_total = torch.cat([k_cache, k_new], dim=1)
v_total = torch.cat([v_cache, v_new], dim=1)
assert k_total.shape == v_total.shape, "Key and value shapes must be the same"
output = self.attn_impl.forward(q, k_total, v_total, ctx_attn_metadata)
# Redistribute back if using sequence parallelism
replicated_output = None
@@ -124,13 +145,17 @@ class DistributedAttention(nn.Module):
# TODO: make this asynchronous
replicated_output = sequence_model_parallel_all_gather(
replicated_output.contiguous(), dim=2)
# Apply backend-specific postprocess_output
output = self.attn_impl.postprocess_output(output, ctx_attn_metadata)
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
gather_dim=2)
return output, replicated_output
if return_current_kv:
return output, replicated_output, k_new, v_new # [batch_size, seq_len, num_heads/sp_world_size, head_dim]
else:
return output, replicated_output, None, None
class DistributedAttention_VSA(DistributedAttention):
@@ -147,7 +172,9 @@ class DistributedAttention_VSA(DistributedAttention):
replicated_k: torch.Tensor | None = None,
replicated_v: torch.Tensor | None = None,
gate_compress: torch.Tensor | None = None,
) -> tuple[torch.Tensor, torch.Tensor | None]:
**kwargs,
) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None,
torch.Tensor | None]:
"""Forward pass for distributed attention.
Args:
@@ -164,8 +191,15 @@ class DistributedAttention_VSA(DistributedAttention):
- o (torch.Tensor): Output tensor after attention for the main sequence
- replicated_o (Optional[torch.Tensor]): Output tensor for replicated tokens, if provided
"""
# Check text tokens are not supported for VSA now
k_cache = kwargs.pop("k_cache", None)
v_cache = kwargs.pop("v_cache", None)
return_current_kv = kwargs.pop("return_current_kv", False)
assert not kwargs, f"Unexpected kwargs for DistributedAttention_VSA: {kwargs.keys()}"
# Check unsupported features
assert replicated_q is None and replicated_k is None and replicated_v is None, "Replicated QKV is not supported for VSA now"
assert not return_current_kv, "return_current_kv is not supported for VSA attention"
assert k_cache is None and v_cache is None, "KV cache is not supported for VSA attention"
# Check input shapes
assert q.dim() == 4 and k.dim() == 4 and v.dim(
) == 4, "Expected 4D tensors"
@@ -197,7 +231,7 @@ class DistributedAttention_VSA(DistributedAttention):
output = sequence_model_parallel_all_to_all_4D(output,
scatter_dim=1,
gather_dim=2)
return output, replicated_output
return output, replicated_output, None, None
class LocalAttention(nn.Module):
+10 -6
View File
@@ -174,6 +174,9 @@ class FastVideoArgs:
init_weights_from_safetensors: str = "" # path to safetensors file for initial weight loading
init_weights_from_safetensors_2: str = "" # path to safetensors file for initial weight loading for transformer_2
# Self-forcing specific arguments
num_frame_per_block: int = 3
# # DMD parameters
# dmd_denoising_steps: List[int] | None = field(default=None)
@@ -433,6 +436,13 @@ class FastVideoArgs:
type=str,
help="Path to safetensors file for initial weight loading")
# Self-forcing specific arguments
parser.add_argument(
"--num-frame-per-block",
type=int,
default=FastVideoArgs.num_frame_per_block,
help="Number of frames per block for causal generation")
# Add pipeline configuration arguments
PipelineConfig.add_cli_args(parser)
@@ -751,7 +761,6 @@ class TrainingArgs(FastVideoArgs):
warp_denoising_step: bool = False
# Self-forcing specific arguments
num_frame_per_block: int = 3
independent_first_frame: bool = False
enable_gradient_masking: bool = True
gradient_mask_last_n_frames: int = 21
@@ -1146,11 +1155,6 @@ class TrainingArgs(FastVideoArgs):
)
# Self-forcing specific arguments
parser.add_argument(
"--num-frame-per-block",
type=int,
default=TrainingArgs.num_frame_per_block,
help="Number of frames per block for causal generation")
parser.add_argument(
"--independent-first-frame",
action=StoreBoolean,
+66 -35
View File
@@ -20,6 +20,7 @@ import fastvideo.envs as envs
from fastvideo.attention import (DistributedAttention,
LocalAttention)
from fastvideo.configs.models.dits import WanVideoConfig
from fastvideo.distributed import sequence_model_parallel_all_gather
from fastvideo.distributed.parallel_state import get_sp_world_size
from fastvideo.forward_context import get_forward_context
from fastvideo.layers.layernorm import (FP32LayerNorm, LayerNormScaleShift,
@@ -59,7 +60,7 @@ class CausalWanSelfAttention(nn.Module):
self.max_attention_size = 32760 if local_attn_size == -1 else local_attn_size * 1560
# Scaled dot product attention
self.attn = LocalAttention(
self.attn = DistributedAttention(
num_heads=num_heads,
head_size=self.head_dim,
dropout_rate=0,
@@ -121,42 +122,65 @@ class CausalWanSelfAttention(nn.Module):
)[:, :, :-padded_length].transpose(2, 1)
else:
frame_seqlen = q.shape[1]
current_end = current_start + roped_query.shape[1]
sink_tokens = self.sink_size * frame_seqlen
# If we are using local attention and the current KV cache size is larger than the local attention size, we need to truncate the KV cache
kv_cache_size = kv_cache["k"].shape[1]
num_new_tokens = roped_query.shape[1]
if self.local_attn_size != -1 and (current_end > kv_cache["global_end_index"].item()) and (
num_new_tokens + kv_cache["local_end_index"].item() > kv_cache_size):
# Calculate the number of new tokens added in this step
# Shift existing cache content left to discard oldest tokens
# Clone the source slice to avoid overlapping memory error
num_evicted_tokens = num_new_tokens + kv_cache["local_end_index"].item() - kv_cache_size
num_rolled_tokens = kv_cache["local_end_index"].item() - num_evicted_tokens - sink_tokens
kv_cache["k"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["k"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
kv_cache["v"][:, sink_tokens:sink_tokens + num_rolled_tokens] = \
kv_cache["v"][:, sink_tokens + num_evicted_tokens:sink_tokens + num_evicted_tokens + num_rolled_tokens].clone()
# Insert the new keys/values at the end
local_end_index = kv_cache["local_end_index"].item() + current_end - \
kv_cache["global_end_index"].item() - num_evicted_tokens
local_start_index = local_end_index - num_new_tokens
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
sp_world_size = max(get_sp_world_size(), 1)
num_new_tokens = roped_query.shape[1] * sp_world_size
prev_local_end = current_start
current_end = current_start + num_new_tokens
assert self.local_attn_size == -1, "Local attention is not supported"
# if self.local_attn_size != -1 and (current_end > prev_global_end) and (
# num_new_tokens + prev_local_end > kv_cache_size):
# num_evicted_tokens = num_new_tokens + prev_local_end - kv_cache_size
# num_rolled_tokens = prev_local_end - num_evicted_tokens - sink_tokens
# if num_rolled_tokens > 0:
# src_slice = slice(sink_tokens + num_evicted_tokens,
# sink_tokens + num_evicted_tokens + num_rolled_tokens)
# dst_slice = slice(sink_tokens, sink_tokens + num_rolled_tokens)
# kv_cache["k"][:, dst_slice] = kv_cache["k"][:, src_slice].clone()
# kv_cache["v"][:, dst_slice] = kv_cache["v"][:, src_slice].clone()
# prev_local_end = max(sink_tokens, prev_local_end - num_evicted_tokens)
# current_end = current_start + roped_query.shape[1]
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
max_history_tokens = max(self.max_attention_size - num_new_tokens, 0)
history_len = min(prev_local_end, max_history_tokens)
history_start = prev_local_end - history_len
if history_len > 0:
k_history = kv_cache["k"][:, history_start:history_start + history_len]
v_history = kv_cache["v"][:, history_start:history_start + history_len]
else:
# 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"].detach()
kv_cache["v"] = kv_cache["v"].detach()
# logger.info("kv_cache['k'] is in comp graph: %s", kv_cache["k"].requires_grad or kv_cache["k"].grad_fn is not None)
kv_cache["k"][:, local_start_index:local_end_index] = roped_key
kv_cache["v"][:, local_start_index:local_end_index] = v
x = self.attn(
k_history = None
v_history = None
attn_output, _, k_new, v_new = self.attn(
roped_query,
kv_cache["k"][:, max(0, local_end_index - self.max_attention_size):local_end_index],
kv_cache["v"][:, max(0, local_end_index - self.max_attention_size):local_end_index]
roped_key,
v,
k_cache=k_history,
v_cache=v_history,
return_current_kv=True,
)
assert k_new.shape[1] == num_new_tokens == v_new.shape[1], "New keys and values must have the same number of tokens"
# assert actual_new_tokens == num_new_tokens, "Actual new tokens must match expected new tokens"
# num_new_tokens = actual_new_tokens
# current_end = current_start + num_new_tokens
new_k = k_new.detach().to(kv_cache["k"].dtype)
new_v = v_new.detach().to(kv_cache["v"].dtype)
kv_cache["k"] = kv_cache["k"].detach()
kv_cache["v"] = kv_cache["v"].detach()
kv_cache["k"][:, local_start_index:local_end_index] = new_k
kv_cache["v"][:, local_start_index:local_end_index] = new_v
x = attn_output
kv_cache["global_end_index"].fill_(current_end)
kv_cache["local_end_index"].fill_(local_end_index)
@@ -371,7 +395,8 @@ class CausalWanTransformer3DModel(BaseDiT):
# Causal-specific
self.block_mask = None
self.num_frame_per_block = config.arch_config.num_frames_per_block
assert self.num_frame_per_block <= 3
self.num_frame_per_block = 4
# assert self.num_frame_per_block <= 3
self.independent_first_frame = False
self.__post_init__()
@@ -472,6 +497,8 @@ class CausalWanTransformer3DModel(BaseDiT):
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
# For Ulysses SP, we need full rotary embeddings (no temporal sharding)
# because after all-to-all, each rank will have the full sequence
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
@@ -480,7 +507,8 @@ class CausalWanTransformer3DModel(BaseDiT):
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000,
start_frame=start_frame # Assume that start_frame is 0 when kv_cache is None
start_frame=start_frame, # Assume that start_frame is 0 when kv_cache is None
# use_sp_shard=False # Don't shard rotary embeddings for Ulysses SP
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
@@ -571,6 +599,8 @@ class CausalWanTransformer3DModel(BaseDiT):
# Get rotary embeddings
d = self.hidden_size // self.num_attention_heads
rope_dim_list = [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)]
# For Ulysses SP, we need full rotary embeddings (no temporal sharding)
# because after all-to-all, each rank will have the full sequence
freqs_cos, freqs_sin = get_rotary_pos_embed(
(post_patch_num_frames * get_sp_world_size(), post_patch_height,
post_patch_width),
@@ -579,7 +609,8 @@ class CausalWanTransformer3DModel(BaseDiT):
rope_dim_list,
dtype=torch.float32 if current_platform.is_mps() else torch.float64,
rope_theta=10000,
start_frame=start_frame
start_frame=start_frame,
use_sp_shard=False # Don't shard rotary embeddings for Ulysses SP
)
freqs_cos = freqs_cos.to(hidden_states.device)
freqs_sin = freqs_sin.to(hidden_states.device)
+57 -16
View File
@@ -1,6 +1,9 @@
import torch # type: ignore
from fastvideo.distributed import get_local_torch_device
from fastvideo.distributed import (get_local_torch_device,
sequence_model_parallel_all_gather,
get_sp_parallel_rank, get_sp_world_size)
from einops import rearrange
from fastvideo.fastvideo_args import FastVideoArgs
from fastvideo.forward_context import set_forward_context
from fastvideo.logger import init_logger
@@ -47,6 +50,7 @@ class CausalDMDDenosingStage(DenoisingStage):
# Model-dependent constants (aligned with causal_inference.py assumptions)
self.num_transformer_blocks = len(self.transformer.blocks)
self.num_frames_per_block = self.transformer.config.arch_config.num_frames_per_block
self.num_frames_per_block = 4
self.sliding_window_num_frames = self.transformer.config.arch_config.sliding_window_num_frames
try:
@@ -218,16 +222,44 @@ class CausalDMDDenosingStage(DenoisingStage):
block_sizes.pop(0)
latents[:, :, :1, :, :] = first_frame_latent
# Setup noise generator for SP
sp_world_size = get_sp_world_size()
sp_rank = get_sp_parallel_rank()
sp_group = sp_world_size > 1
noise_generator = batch.generator[0] if isinstance(
batch.generator, list) else batch.generator
if sp_group:
try:
# Use initial_seed + rank to ensure diversity across ranks
seed = noise_generator.initial_seed()
noise_generator = torch.Generator(
device=noise_generator.device).manual_seed(seed + sp_rank)
except Exception as e:
logger.warning(f"Could not create rank-specific generator: {e}")
# DMD loop in causal blocks
with self.progress_bar(total=len(block_sizes) *
len(timesteps)) as progress_bar:
for current_num_frames in block_sizes:
for idx, current_num_frames in enumerate(block_sizes):
current_latents = latents[:, :, start_index:start_index +
current_num_frames, :, :]
sp_world_size, rank_in_sp_group = get_sp_world_size(
), get_sp_parallel_rank()
sp_group = sp_world_size > 1
if sp_group:
current_latents = rearrange(current_latents,
"b c (n t) h w -> b c n t h w",
n=sp_world_size).contiguous()
current_latents = current_latents[:, :,
rank_in_sp_group, :, :, :]
# use BTCHW for DMD conversion routines
noise_latents_btchw = current_latents.permute(0, 2, 1, 3, 4)
video_raw_latent_shape = noise_latents_btchw.shape
logger.info(
f"before DMD: Noise latents shape: {noise_latents_btchw.shape}"
)
for i, t_cur in enumerate(timesteps):
if boundary_timestep is not None and t_cur < boundary_timestep:
current_model = self.transformer_2
@@ -283,6 +315,7 @@ class CausalDMDDenosingStage(DenoisingStage):
(latent_model_input.shape[0], 1),
device=latent_model_input.device,
dtype=torch.long)
assert current_model is not None, "current_model is not initialized"
pred_noise_btchw = current_model(
latent_model_input,
prompt_embeds,
@@ -319,12 +352,10 @@ class CausalDMDDenosingStage(DenoisingStage):
[1],
dtype=torch.long,
device=pred_video_btchw.device)
noise = torch.randn(
video_raw_latent_shape,
dtype=pred_video_btchw.dtype,
generator=(batch.generator[0] if isinstance(
batch.generator, list) else
batch.generator)).to(self.device)
noise = torch.randn(video_raw_latent_shape,
dtype=pred_video_btchw.dtype,
generator=noise_generator).to(
self.device)
noise_btchw = noise
if boundary_timestep is not None and i < len(
high_noise_timesteps) - 1:
@@ -352,10 +383,6 @@ class CausalDMDDenosingStage(DenoisingStage):
if progress_bar is not None:
progress_bar.update()
# Write back and advance
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
# Re-run with context timestep to update KV cache using clean context
context_noise = getattr(fastvideo_args.pipeline_config,
"context_noise", 0)
@@ -363,6 +390,7 @@ class CausalDMDDenosingStage(DenoisingStage):
device=latents.device,
dtype=torch.long) * int(context_noise)
context_bcthw = current_latents.to(target_dtype)
with torch.autocast(device_type="cuda",
dtype=target_dtype,
enabled=autocast_enabled), \
@@ -372,6 +400,7 @@ class CausalDMDDenosingStage(DenoisingStage):
t_expanded_context = t_context.unsqueeze(1)
if boundary_timestep is not None:
assert self.transformer_2 is not None, "transformer_2 is not initialized"
self.transformer_2(
context_bcthw,
prompt_embeds,
@@ -398,6 +427,13 @@ class CausalDMDDenosingStage(DenoisingStage):
**pos_cond_kwargs,
)
if sp_group:
current_latents = sequence_model_parallel_all_gather(
current_latents, dim=2)
# Write back and advance
latents[:, :, start_index:start_index +
current_num_frames, :, :] = current_latents
start_index += current_num_frames
if boundary_timestep is not None:
@@ -413,6 +449,11 @@ class CausalDMDDenosingStage(DenoisingStage):
"""
kv_cache1 = []
num_attention_heads = self.transformer.num_attention_heads
sp_world_size = get_sp_world_size()
assert num_attention_heads % max(
sp_world_size,
1) == 0, "num_attention_heads must be divisible by sp_world_size"
heads_per_rank = num_attention_heads // max(sp_world_size, 1)
attention_head_dim = self.transformer.attention_head_dim
if self.local_attn_size != -1:
kv_cache_size = self.local_attn_size * self.frame_seq_length
@@ -423,14 +464,14 @@ class CausalDMDDenosingStage(DenoisingStage):
kv_cache1.append({
"k":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
batch_size, kv_cache_size, heads_per_rank,
attention_head_dim
],
dtype=dtype,
device=device),
"v":
torch.zeros([
batch_size, kv_cache_size, num_attention_heads,
batch_size, kv_cache_size, heads_per_rank,
attention_head_dim
],
dtype=dtype,
@@ -494,4 +535,4 @@ class CausalDMDDenosingStage(DenoisingStage):
result.add_check(
"negative_prompt_embeds", batch.negative_prompt_embeds, lambda x:
not batch.do_classifier_free_guidance or V.list_not_empty(x))
return result
return result