(self, args, device)
| 47 | ) |
| 48 | |
| 49 | def _initialize_models(self, args, device): |
| 50 | self.real_model_name = getattr(args, "real_name", "Wan2.2-TI2V-5B") |
| 51 | self.fake_model_name = getattr(args, "fake_name", "Wan2.2-TI2V-5B") |
| 52 | self.local_attn_size = section_get( |
| 53 | args, |
| 54 | "inference", |
| 55 | "local_attn_size", |
| 56 | getattr(args, "model_kwargs", {}).get("local_attn_size", -1), |
| 57 | aliases=("inference_local_attn_size",), |
| 58 | ) |
| 59 | all_causal = getattr(args, "all_causal", False) |
| 60 | score_is_causal = all_causal |
| 61 | |
| 62 | model_name = args.model_kwargs.get("model_name", "Wan2.2-TI2V-5B") |
| 63 | if "5B" not in model_name: |
| 64 | raise ValueError(f"Only Wan2.2-TI2V-5B is supported in this release, got {model_name}") |
| 65 | if not dist.is_initialized() or dist.get_rank() == 0: |
| 66 | tag = "all-causal 5B mode" if all_causal else "Wan2.2-TI2V-5B" |
| 67 | print(f"Using {tag}") |
| 68 | |
| 69 | # Generator |
| 70 | generator_is_causal = getattr(args, "generator_is_causal", True) |
| 71 | self.generator = WanDiffusionWrapper(**getattr(args, "model_kwargs", {}), is_causal=generator_is_causal) |
| 72 | self.generator.model.requires_grad_(True) |
| 73 | |
| 74 | # Real Score |
| 75 | real_kwargs = getattr(args, "real_model_kwargs", {"model_name": self.real_model_name}) |
| 76 | self.real_score = WanDiffusionWrapper(**real_kwargs, is_causal=score_is_causal) |
| 77 | self.real_score.model.requires_grad_(False) |
| 78 | |
| 79 | # Fake Score |
| 80 | fake_kwargs = getattr(args, "fake_model_kwargs", {"model_name": self.fake_model_name}) |
| 81 | self.fake_score = WanDiffusionWrapper(**fake_kwargs, is_causal=score_is_causal) |
| 82 | self.fake_score.model.requires_grad_(True) |
| 83 | |
| 84 | # Text Encoder & VAE |
| 85 | self.text_encoder = WanTextEncoder() |
| 86 | self.text_encoder.requires_grad_(False) |
| 87 | |
| 88 | self.vae = WanVAEWrapper() |
| 89 | self.vae.requires_grad_(False) |
| 90 | |
| 91 | self.scheduler = self.generator.get_scheduler() |
| 92 | self.scheduler.timesteps = self.scheduler.timesteps.to(device) |
| 93 | |
| 94 | def _get_timestep( |
| 95 | self, |
no test coverage detected