Spaces:
Running on Zero
Running on Zero
| import torch.nn.functional as F | |
| from typing import Optional, Tuple | |
| import torch | |
| import time | |
| from model.base import SelfForcingModel | |
| import torch.distributed as dist | |
| from model.dmd import DMD | |
| from pipeline.streaming_switch_training import StreamingSwitchTrainingPipeline | |
| from einops import rearrange | |
| class DMDSwitch(DMD): | |
| def _initialize_inference_pipeline(self): | |
| self.inference_pipeline = StreamingSwitchTrainingPipeline(denoising_step_list=self.denoising_step_list, scheduler=self.scheduler, generator=self.generator, num_frame_per_block=self.num_frame_per_block, same_step_across_blocks=self.args.same_step_across_blocks, last_step_only=self.args.last_step_only, context_noise=self.args.context_noise, local_attn_size=getattr(self.args, 'model_kwargs', {}).get('local_attn_size', -1), slice_last_frames=getattr(self.args, 'slice_last_frames', 21), global_sink=getattr(self.args, 'global_sink', False), apr_enabled=getattr(self.args, 'apr_enabled', False), apr_alpha_max=getattr(self.args, 'apr_alpha_max', 0.8), apr_d_window=getattr(self.args, 'apr_d_window', None), apr_blend_sink=getattr(self.args, 'apr_blend_sink', False)) | |