echo-infinity / pipeline /streaming_training.py
multimodalart's picture
multimodalart HF Staff
Upload folder using huggingface_hub
3e936b2 verified
Raw
History Blame Contribute Delete
8.98 kB
from utils.wan_wrapper import WanDiffusionWrapper
from utils.scheduler import SchedulerInterface
from typing import List, Optional, Tuple
import torch
import torch.distributed as dist
class StreamingTrainingPipeline:
def __init__(self, denoising_step_list: List[int], scheduler: SchedulerInterface, generator: WanDiffusionWrapper, num_frame_per_block=3, same_step_across_blocks: bool=False, last_step_only: bool=False, context_noise: int=0, **kwargs):
super().__init__()
self.scheduler = scheduler
self.generator = generator
self.denoising_step_list = denoising_step_list
if self.denoising_step_list[-1] == 0:
self.denoising_step_list = self.denoising_step_list[:-1]
self.num_transformer_blocks = 30
self.frame_seq_length = 1560
self.num_frame_per_block = num_frame_per_block
self.context_noise = context_noise
self.kv_cache1 = None
self.crossattn_cache = None
self.same_step_across_blocks = same_step_across_blocks
self.last_step_only = last_step_only
self.local_attn_size = kwargs.get('local_attn_size', -1)
slice_last_frames: int = int(kwargs.get('slice_last_frames', 21))
self.kv_cache_size = (self.local_attn_size + slice_last_frames) * self.frame_seq_length
def generate_and_sync_list(self, num_blocks, num_denoising_steps, device):
rank = dist.get_rank() if dist.is_initialized() else 0
if rank == 0:
indices = torch.randint(low=0, high=num_denoising_steps, size=(num_blocks,), device=device)
if self.last_step_only:
indices = torch.ones_like(indices) * (num_denoising_steps - 1)
else:
indices = torch.empty(num_blocks, dtype=torch.long, device=device)
if dist.is_initialized():
dist.broadcast(indices, src=0)
return indices.tolist()
def generate_chunk_with_cache(self, noise: torch.Tensor, conditional_dict: dict, *, current_start_frame: int=0, requires_grad: bool=True, return_sim_step: bool=False) -> Tuple[torch.Tensor, Optional[int], Optional[int]]:
batch_size, chunk_frames, num_channels, height, width = noise.shape
assert chunk_frames % self.num_frame_per_block == 0
num_blocks = chunk_frames // self.num_frame_per_block
all_num_frames = [self.num_frame_per_block] * num_blocks
output = torch.zeros_like(noise)
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)
if not requires_grad:
start_gradient_frame_index = chunk_frames
else:
start_gradient_frame_index = 0
local_start_frame = 0
self.generator.model.local_attn_size = int(self.local_attn_size)
self._set_all_modules_max_attention_size(int(self.local_attn_size))
for block_index, current_num_frames in enumerate(all_num_frames):
noisy_input = noise[:, local_start_frame:local_start_frame + current_num_frames]
for step_idx, current_timestep in enumerate(self.denoising_step_list):
exit_flag = step_idx == exit_flags[0] if self.same_step_across_blocks else step_idx == exit_flags[block_index]
timestep = torch.ones([batch_size, current_num_frames], device=noise.device, dtype=torch.int64) * current_timestep
if not exit_flag:
with torch.no_grad():
_, denoised_pred = self.generator(noisy_image_or_video=noisy_input, conditional_dict=conditional_dict, timestep=timestep, kv_cache=self.kv_cache1, crossattn_cache=self.crossattn_cache, current_start=(current_start_frame + local_start_frame) * self.frame_seq_length)
if step_idx < len(self.denoising_step_list) - 1:
next_timestep = self.denoising_step_list[step_idx + 1]
noisy_input = self.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])
else:
enable_grad = local_start_frame >= start_gradient_frame_index
context_manager = torch.enable_grad() if enable_grad else torch.no_grad()
with context_manager:
_, denoised_pred = self.generator(noisy_image_or_video=noisy_input, conditional_dict=conditional_dict, timestep=timestep, kv_cache=self.kv_cache1, crossattn_cache=self.crossattn_cache, current_start=(current_start_frame + local_start_frame) * self.frame_seq_length)
break
output[:, local_start_frame:local_start_frame + current_num_frames] = denoised_pred
context_timestep = torch.ones_like(timestep) * self.context_noise
context_noisy = self.scheduler.add_noise(denoised_pred.flatten(0, 1), torch.randn_like(denoised_pred.flatten(0, 1)), context_timestep.flatten(0, 1)).unflatten(0, denoised_pred.shape[:2])
with torch.no_grad():
self.generator(noisy_image_or_video=context_noisy, conditional_dict=conditional_dict, timestep=context_timestep, kv_cache=self.kv_cache1, crossattn_cache=self.crossattn_cache, current_start=(current_start_frame + local_start_frame) * self.frame_seq_length)
local_start_frame += current_num_frames
if not self.same_step_across_blocks:
denoised_timestep_from, denoised_timestep_to = (None, None)
elif exit_flags[0] == len(self.denoising_step_list) - 1:
denoised_timestep_to = 0
denoised_timestep_from = 1000 - torch.argmin((self.scheduler.timesteps.cuda() - self.denoising_step_list[exit_flags[0]].cuda()).abs(), dim=0).item()
else:
denoised_timestep_to = 1000 - torch.argmin((self.scheduler.timesteps.cuda() - self.denoising_step_list[exit_flags[0] + 1].cuda()).abs(), dim=0).item()
denoised_timestep_from = 1000 - torch.argmin((self.scheduler.timesteps.cuda() - self.denoising_step_list[exit_flags[0]].cuda()).abs(), dim=0).item()
if return_sim_step:
return (output, denoised_timestep_from, denoised_timestep_to, exit_flags[0] + 1)
return (output, denoised_timestep_from, denoised_timestep_to)
def _initialize_kv_cache(self, batch_size, dtype, device):
kv_cache1 = []
for _ in range(self.num_transformer_blocks):
kv_cache1.append({'k': torch.zeros([batch_size, self.kv_cache_size, 12, 128], dtype=dtype, device=device), 'v': torch.zeros([batch_size, self.kv_cache_size, 12, 128], dtype=dtype, device=device), 'global_end_index': torch.tensor([0], dtype=torch.long, device=device), 'local_end_index': torch.tensor([0], dtype=torch.long, device=device)})
self.kv_cache1 = kv_cache1
def _initialize_crossattn_cache(self, batch_size, dtype, device):
crossattn_cache = []
for _ in range(self.num_transformer_blocks):
crossattn_cache.append({'k': torch.zeros([batch_size, 512, 12, 128], dtype=dtype, device=device), 'v': torch.zeros([batch_size, 512, 12, 128], dtype=dtype, device=device), 'is_init': False})
self.crossattn_cache = crossattn_cache
def clear_kv_cache(self):
if getattr(self, 'kv_cache1', None) is not None:
for blk in self.kv_cache1:
blk['k'].zero_()
blk['v'].zero_()
if 'global_end_index' in blk:
blk['global_end_index'].zero_()
if 'local_end_index' in blk:
blk['local_end_index'].zero_()
if getattr(self, 'crossattn_cache', None) is not None:
for blk in self.crossattn_cache:
blk['k'].zero_()
blk['v'].zero_()
blk['is_init'] = False
def _set_all_modules_max_attention_size(self, local_attn_size_value: int):
if isinstance(local_attn_size_value, (list, tuple)):
raise ValueError('_set_all_modules_max_attention_size expects an int, got list/tuple.')
if int(local_attn_size_value) == -1:
target_size = 32760
policy = 'global'
else:
target_size = int(local_attn_size_value) * self.frame_seq_length
policy = 'local'
if hasattr(self.generator.model, 'max_attention_size'):
try:
_ = getattr(self.generator.model, 'max_attention_size')
except Exception:
pass
setattr(self.generator.model, 'max_attention_size', target_size)
for name, module in self.generator.model.named_modules():
if hasattr(module, 'max_attention_size'):
try:
setattr(module, 'max_attention_size', target_size)
except Exception:
pass