| 21 | |
| 22 | |
| 23 | class FixPipeline(BasePipeline): |
| 24 | |
| 25 | def __init__(self, device="cuda", torch_dtype=torch.float16, tokenizer_path=None): |
| 26 | super().__init__(device=device, torch_dtype=torch_dtype) |
| 27 | self.scheduler = FlowMatchScheduler(shift=5, sigma_min=0.0, extra_one_step=True) |
| 28 | self.prompter = WanPrompter(tokenizer_path=tokenizer_path) |
| 29 | self.text_encoder: WanTextEncoder = None |
| 30 | self.image_encoder: WanImageEncoder = None |
| 31 | self.dit: WanModel = None |
| 32 | self.vae: WanVideoVAE = None |
| 33 | self.model_names = ['text_encoder', 'dit', 'vae'] |
| 34 | self.height_division_factor = 16 |
| 35 | self.width_division_factor = 16 |
| 36 | |
| 37 | |
| 38 | def enable_vram_management(self, num_persistent_param_in_dit=None): |
| 39 | dtype = next(iter(self.text_encoder.parameters())).dtype |
| 40 | enable_vram_management( |
| 41 | self.text_encoder, |
| 42 | module_map = { |
| 43 | torch.nn.Linear: AutoWrappedLinear, |
| 44 | torch.nn.Embedding: AutoWrappedModule, |
| 45 | T5RelativeEmbedding: AutoWrappedModule, |
| 46 | T5LayerNorm: AutoWrappedModule, |
| 47 | }, |
| 48 | module_config = dict( |
| 49 | offload_dtype=dtype, |
| 50 | offload_device="cpu", |
| 51 | onload_dtype=dtype, |
| 52 | onload_device="cpu", |
| 53 | computation_dtype=self.torch_dtype, |
| 54 | computation_device=self.device, |
| 55 | ), |
| 56 | ) |
| 57 | dtype = next(iter(self.dit.parameters())).dtype |
| 58 | enable_vram_management( |
| 59 | self.dit, |
| 60 | module_map = { |
| 61 | torch.nn.Linear: AutoWrappedLinear, |
| 62 | torch.nn.Conv3d: AutoWrappedModule, |
| 63 | torch.nn.LayerNorm: AutoWrappedModule, |
| 64 | RMSNorm: AutoWrappedModule, |
| 65 | }, |
| 66 | module_config = dict( |
| 67 | offload_dtype=dtype, |
| 68 | offload_device="cpu", |
| 69 | onload_dtype=dtype, |
| 70 | onload_device=self.device, |
| 71 | computation_dtype=self.torch_dtype, |
| 72 | computation_device=self.device, |
| 73 | ), |
| 74 | max_num_param=num_persistent_param_in_dit, |
| 75 | overflow_module_config = dict( |
| 76 | offload_dtype=dtype, |
| 77 | offload_device="cpu", |
| 78 | onload_dtype=dtype, |
| 79 | onload_device="cpu", |
| 80 | computation_dtype=self.torch_dtype, |