(self, args, device)
| 120 | self.er_skip_block_0 = bool(cfg_dict.get("skip_block_0", False)) |
| 121 | |
| 122 | def _initialize_models(self, args, device): |
| 123 | model_name = getattr(args.model_kwargs, "model_name", "Wan2.2-TI2V-5B") |
| 124 | if "5B" not in model_name: |
| 125 | raise ValueError(f"Only Wan2.2-TI2V-5B is supported in this release, got {model_name}") |
| 126 | self.generator = WanDiffusionWrapper(**getattr(args, "model_kwargs", {}), is_causal=True) |
| 127 | self.generator.model.requires_grad_(True) |
| 128 | |
| 129 | self.text_encoder = WanTextEncoder() |
| 130 | self.text_encoder.requires_grad_(False) |
| 131 | |
| 132 | self.vae = WanVAEWrapper() |
| 133 | self.vae.requires_grad_(False) |
| 134 | |
| 135 | self.scheduler = self.generator.get_scheduler() |
| 136 | self.scheduler.timesteps = self.scheduler.timesteps.to(device) |
| 137 | |
| 138 | def generator_loss( |
| 139 | self, |
nothing calls this directly
no test coverage detected