BuildingVirtuallyAI/RenderingModel
0
1from __future__ import annotations
2import gc
3import numpy as np
4import PIL.Image
5import torch
6from diffusers import (
7 ControlNetModel,
8 DiffusionPipeline,
9 StableDiffusionControlNetPipeline,
10 UniPCMultistepScheduler,
11)
12
13from preprocessor import Preprocessor
14from settings import *
15
16
17class Model:
18 def __init__(self, base_model_id: str = "runwayml/stable-diffusion-v1-5", task_name: str = "lineart"):
19 self.device = torch.device("cuda:0" if torch.cuda.is_available() else "cpu")
20 self.base_model_id = ""
21 self.task_name = ""
22 self.pipe = self.load_pipe(base_model_id, task_name)
23 self.preprocessor = Preprocessor()
24
25 def load_pipe(self, base_model_id: str, task_name) -> DiffusionPipeline:
26 if (
27 base_model_id == self.base_model_id
28 and task_name == self.task_name
29 and hasattr(self, "pipe")
30 and self.pipe is not None
31 ):
32 return self.pipe
33 controlnet = ControlNetModel.from_pretrained(model_id, torch_dtype=torch.float16)
34 pipe = StableDiffusionControlNetPipeline.from_pretrained(
35 base_model_id, safety_checker=None, controlnet=controlnet, torch_dtype=torch.float16
36 )
37 pipe.scheduler = UniPCMultistepScheduler.from_config(pipe.scheduler.config)
38 if self.device.type == "cuda":
39 pipe.enable_xformers_memory_efficient_attention()
40 pipe.to(self.device)
41 torch.cuda.empty_cache()
42 gc.collect()
43 self.base_model_id = base_model_id
44 self.task_name = task_name
45 return pipe
46
47 def set_base_model(self, base_model_id: str) -> str:
48 if not base_model_id or base_model_id == self.base_model_id:
49 return self.base_model_id
50 del self.pipe
51 torch.cuda.empty_cache()
52 gc.collect()
53 try:
54 self.pipe = self.load_pipe(base_model_id, self.task_name)
55 except Exception:
56 self.pipe = self.load_pipe(self.base_model_id, self.task_name)
57 return self.base_model_id
58
59 def load_controlnet_weight(self, task_name: str) -> None:
60 if task_name == self.task_name:
61 return
62 if self.pipe is not None and hasattr(self.pipe, "controlnet"):
63 del self.pipe.controlnet
64 torch.cuda.empty_cache()
65 gc.collect()
66 controlnet = ControlNetModel.from_pretrained(model_id, torch_dtype=torch.float16)
67 controlnet.to(self.device)
68 torch.cuda.empty_cache()
69 gc.collect()
70 self.pipe.controlnet = controlnet
71 self.task_name = task_name
72
73 def get_prompt(self, prompt: str, additional_prompt: str) -> str:
74 if not prompt:
75 prompt = additional_prompt
76 else:
77 prompt = f"{prompt}, {additional_prompt}"
78 return prompt
79
80 @torch.autocast("cuda")
81 def run_pipe(
82 self,
83 control_image: PIL.Image.Image,
84 ) -> list[PIL.Image.Image]:
85 generator = torch.Generator().manual_seed(randomize_seed)
86 return self.pipe(
87 prompt=prompt + ' ' + a_prompt,
88 negative_prompt=n_prompt,
89 guidance_scale=guidance_scale,
90 num_images_per_prompt=DEFAULT_NUM_IMAGES,
91 num_inference_steps=num_steps,
92 generator=generator,
93 image=control_image,
94 ).images
95
96 def process_lineart(
97 self,
98 image: np.ndarray,
99 ) -> list[PIL.Image.Image]:
100 if image is None:
101 raise ValueError
102
103 else:
104
105 self.preprocessor.load("Lineart")
106 control_image = self.preprocessor(
107 image=image,
108 image_resolution=DEFAULT_IMAGE_RESOLUTION,
109 detect_resolution=preprocess_resolution,
110 )
111 self.load_controlnet_weight("lineart")
112 results = self.run_pipe(
113 control_image=control_image
114 )
115 return [control_image] + results