Team Ai
Modelpublic

BuildingVirtuallyAI/RenderingModel

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
model.py115 linesDownload Raw Back to root
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