Views
No views yet
1
2import torch
3
4from diffusers import StableDiffusionPipeline, DPMSolverMultistepScheduler
5
6from animatediff.models.unet import UNet3DConditionModel
7from animatediff.models.sparse_controlnet import SparseControlNetModel
8from animatediff.pipelines.pipeline_animation import AnimationPipeline
9from animatediff.utils.util import load_weights
10
11sdpipe = StableDiffusionPipeline.from_single_file(pretrained_model_path, use_safetensors=True, add_watermarker=False).to(dtype=torch.float16)
12sdpipe.load_lora_weights(lora_model_path)
13sdpipe.fuse_lora(lora_scale=0.3)
14
15text_encoder = sdpipe.text_encoder.cuda()
16vae = sdpipe.vae.cuda()
17tokenizer = sdpipe.tokenizer
18
19unet_additional_kwargs = params["unet_additional_kwargs"]
20controlnet_additional_kwargs = params["controlnet_additional_kwargs"]
21
22unet = UNet3DConditionModel.from_pretrained_2d(pretrained_model_path, subfolder="unet", unet_config=sdpipe.unet.config, unet_additional_kwargs=unet_additional_kwargs).cuda()
23unet.config.num_attention_heads = 8
24unet.config.projection_class_embeddings_input_dim = None
25unet.to(dtype=torch.float16)
26
27controlnet = SparseControlNetModel.from_unet(unet, controlnet_additional_kwargs=controlnet_additional_kwargs)
28controlnet_path = "models/motion_module/v3_sd15_sparsectrl_rgb.ckpt"
29
30print(f"loading controlnet checkpoint from {controlnet_path} ...")
31controlnet_state_dict = torch.load(controlnet_path, map_location="cpu")
32controlnet_state_dict = controlnet_state_dict["controlnet"] if "controlnet" in controlnet_state_dict else controlnet_state_dict
33controlnet_state_dict = {name: param for name, param in controlnet_state_dict.items() if "pos_encoder.pe" not in name}
34controlnet_state_dict.pop("animatediff_config", "")
35controlnet.load_state_dict(controlnet_state_dict)
36controlnet.to(dtype=torch.float16)
37controlnet.cuda()
38
39pipe = load_weights(
40 pipeline,
41 # motion module
42 motion_module_path = "models/Motion_Module/v3_sd15_mm.ckpt",
43 motion_module_lora_configs = [],
44 # domain adapter
45 adapter_lora_path = "models/Motion_Module/v3_sd15_adapter.ckpt",
46 adapter_lora_scale = 1.0,
47 # image layers
48 dreambooth_model_path = pretrained_model_path,
49 lora_model_path = "",
50 lora_alpha = 0.8,
51).to("cuda")
52pipe.to(dtype=torch.float16)
53pipe.enable_vae_slicing()
54pipe.enable_model_cpu_offload()
55
56pipe.scheduler = DPMSolverMultistepScheduler(
57 beta_start = 0.00075,
58 beta_end = 0.0145,
59 beta_schedule = "linear",
60 use_karras_sigmas = True,
61)
621controlnet_additional_kwargs = params["controlnet_additional_kwargs"]
2
3pipe = AnimationPipeline.from_pretrained("models/animatediff_model")
4unet = pipe.unet
5vae = pipe.vae
6
7unet.config.num_attention_heads = 8
8unet.config.projection_class_embeddings_input_dim = None
9unet.to(dtype=torch.float16)
10
11controlnet = SparseControlNetModel.from_unet(unet, controlnet_additional_kwargs=controlnet_additional_kwargs)
12controlnet_path = "./models/motion_module/v3_sd15_sparsectrl_rgb.ckpt"
13
14print(f"loading controlnet checkpoint from {controlnet_path} ...")
15controlnet_state_dict = torch.load(controlnet_path, map_location="cpu")
16controlnet_state_dict = controlnet_state_dict["controlnet"] if "controlnet" in controlnet_state_dict else controlnet_state_dict
17controlnet_state_dict = {name: param for name, param in controlnet_state_dict.items() if "pos_encoder.pe" not in name}
18controlnet_state_dict.pop("animatediff_config", "")
19controlnet.load_state_dict(controlnet_state_dict)
20controlnet.to(dtype=torch.float16)
21controlnet.cuda()
22
23pipe.controlnet = controlnet
24
25without_xformers = False
26if is_xformers_available() and (not without_xformers):
27 unet.enable_xformers_memory_efficient_attention()
28 if controlnet is not None:
29 print("\nenable_xformers_memory_efficient_attention\n")
30 controlnet.enable_xformers_memory_efficient_attention()
31
32pipe.to(dtype=torch.float16)
33pipe.enable_vae_slicing()
34pipe.enable_model_cpu_offload()
35pipe.to("cuda")