Compare commits
7
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
2dce62c034 | ||
|
|
4888a31750 | ||
|
|
fd899c16f9 | ||
|
|
a46f5ea380 | ||
|
|
1f1974bfa2 | ||
|
|
c147f97d3c | ||
|
|
10384b1ad3 |
@@ -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()
|
||||
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user