Spaces:
Sleeping
Sleeping
| 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 | |