MCPcopy Create free account
hub / github.com/NVlabs/LongLive / _initialize_models

Method _initialize_models

model/base.py:49–92  ·  view source on GitHub ↗
(self, args, device)

Source from the content-addressed store, hash-verified

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,

Callers 1

__init__Method · 0.95

Calls 6

section_getFunction · 0.90
WanDiffusionWrapperClass · 0.90
WanTextEncoderClass · 0.90
WanVAEWrapperClass · 0.90
toMethod · 0.80
get_schedulerMethod · 0.45

Tested by

no test coverage detected