MCPcopy Create free account
hub / github.com/YesianRohn/TextSSR / main

Function main

diffusers/scripts/convert_dance_diffusion_to_diffusers.py:258–332  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

256
257
258def main(args):
259 device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
260
261 model_name = args.model_path.split("/")[-1].split(".")[0]
262 if not os.path.isfile(args.model_path):
263 assert (
264 model_name == args.model_path
265 ), f"Make sure to provide one of the official model names {MODELS_MAP.keys()}"
266 args.model_path = download(model_name)
267
268 sample_rate = MODELS_MAP[model_name]["sample_rate"]
269 sample_size = MODELS_MAP[model_name]["sample_size"]
270
271 config = Object()
272 config.sample_size = sample_size
273 config.sample_rate = sample_rate
274 config.latent_dim = 0
275
276 diffusers_model = UNet1DModel(sample_size=sample_size, sample_rate=sample_rate)
277 diffusers_state_dict = diffusers_model.state_dict()
278
279 orig_model = DiffusionUncond(config)
280 orig_model.load_state_dict(torch.load(args.model_path, map_location=device)["state_dict"])
281 orig_model = orig_model.diffusion_ema.eval()
282 orig_model_state_dict = orig_model.state_dict()
283 renamed_state_dict = rename_orig_weights(orig_model_state_dict)
284
285 renamed_minus_diffusers = set(renamed_state_dict.keys()) - set(diffusers_state_dict.keys())
286 diffusers_minus_renamed = set(diffusers_state_dict.keys()) - set(renamed_state_dict.keys())
287
288 assert len(renamed_minus_diffusers) == 0, f"Problem with {renamed_minus_diffusers}"
289 assert all(k.endswith("kernel") for k in list(diffusers_minus_renamed)), f"Problem with {diffusers_minus_renamed}"
290
291 for key, value in renamed_state_dict.items():
292 assert (
293 diffusers_state_dict[key].squeeze().shape == value.squeeze().shape
294 ), f"Shape for {key} doesn't match. Diffusers: {diffusers_state_dict[key].shape} vs. {value.shape}"
295 if key == "time_proj.weight":
296 value = value.squeeze()
297
298 diffusers_state_dict[key] = value
299
300 diffusers_model.load_state_dict(diffusers_state_dict)
301
302 steps = 100
303 seed = 33
304
305 diffusers_scheduler = IPNDMScheduler(num_train_timesteps=steps)
306
307 generator = torch.manual_seed(seed)
308 noise = torch.randn([1, 2, config.sample_size], generator=generator).to(device)
309
310 t = torch.linspace(1, 0, steps + 1, device=device)[:-1]
311 step_list = get_crash_schedule(t)
312
313 pipe = DanceDiffusionPipeline(unet=diffusers_model, scheduler=diffusers_scheduler)
314
315 generator = torch.manual_seed(33)

Calls 14

UNet1DModelClass · 0.90
IPNDMSchedulerClass · 0.90
downloadFunction · 0.85
ObjectClass · 0.85
DiffusionUncondClass · 0.85
rename_orig_weightsFunction · 0.85
get_crash_scheduleFunction · 0.85
load_state_dictMethod · 0.80
deviceMethod · 0.45
state_dictMethod · 0.45
loadMethod · 0.45

Tested by

no test coverage detected