MCPcopy Create free account
hub / github.com/Vchitect/LaVie / setup

Method setup

predict.py:45–157  ·  view source on GitHub ↗

Load the model into memory to make running multiple predictions efficient

(self)

Source from the content-addressed store, hash-verified

43
44class Predictor(BasePredictor):
45 def setup(self) -> None:
46 """Load the model into memory to make running multiple predictions efficient"""
47 pretrained_path = "pretrained_models"
48 torch.set_grad_enabled(False)
49 self.device = "cuda:0"
50
51 # ============ base model ============
52 args_base = OmegaConf.load("base/configs/sample.yaml")
53 sd_path = pretrained_path + "/stable-diffusion-v1-4"
54 self.unet = get_models_base(args_base, sd_path).to(
55 self.device, dtype=torch.float16
56 )
57 state_dict = find_model(pretrained_path + "/lavie_base.pt")
58 self.unet.load_state_dict(state_dict)
59
60 self.vae = AutoencoderKL.from_pretrained(
61 sd_path, subfolder="vae", torch_dtype=torch.float16
62 ).to(self.device)
63 self.tokenizer_one = CLIPTokenizer.from_pretrained(
64 sd_path, subfolder="tokenizer"
65 )
66 self.text_encoder_one = CLIPTextModel.from_pretrained(
67 sd_path, subfolder="text_encoder", torch_dtype=torch.float16
68 ).to(self.device)
69
70 self.unet.eval()
71 self.vae.eval()
72 self.text_encoder_one.eval()
73
74 self.schedulers = {
75 "ddim": DDIMScheduler.from_pretrained(
76 sd_path,
77 subfolder="scheduler",
78 beta_start=args_base.beta_start,
79 beta_end=args_base.beta_end,
80 beta_schedule=args_base.beta_schedule,
81 ),
82 "eulerdiscrete": EulerDiscreteScheduler.from_pretrained(
83 sd_path,
84 subfolder="scheduler",
85 beta_start=args_base.beta_start,
86 beta_end=args_base.beta_end,
87 beta_schedule=args_base.beta_schedule,
88 ),
89 "ddpm": DDPMScheduler.from_pretrained(
90 sd_path,
91 subfolder="scheduler",
92 beta_start=args_base.beta_start,
93 beta_end=args_base.beta_end,
94 beta_schedule=args_base.beta_schedule,
95 ),
96 }
97
98 # ============ interpolation model ============
99 interpolation_ckpt_path = pretrained_path + "/lavie_interpolation.pt"
100 self.args_interpolation = OmegaConf.load(
101 "interpolation/configs/sample.yaml"
102 ).args

Callers

nothing calls this directly

Calls 3

find_modelFunction · 0.90
TextEmbedderClass · 0.90
create_diffusionFunction · 0.90

Tested by

no test coverage detected