SemaSci/DiffModels
0
1import gradio as gr2import numpy as np3import random4import os5import torch6from diffusers import StableDiffusionPipeline, ControlNetModel, StableDiffusionControlNetPipeline7from diffusers.utils import load_image8from peft import PeftModel, LoraConfig9 10from rembg import remove 11 12 13device = "cuda" if torch.cuda.is_available() else "cpu"14model_id_default = "stable-diffusion-v1-5/stable-diffusion-v1-5"15 16if torch.cuda.is_available():17 torch_dtype = torch.float1618else:19 torch_dtype = torch.float3220 21MAX_SEED = np.iinfo(np.int32).max22MAX_IMAGE_SIZE = 102423 24 25# @spaces.GPU #[uncomment to use ZeroGPU]26def infer(27 prompt,28 negative_prompt,29 width=512,30 height=512,31 model_id=model_id_default,32 seed=42,33 guidance_scale=7.0,34 lora_scale=1.0,35 num_inference_steps=20,36 controlnet_checkbox=False, 37 controlnet_strength=0.0, 38 controlnet_mode="edge_detection", 39 controlnet_image=None, 40 ip_adapter_checkbox=False, 41 ip_adapter_scale=0.0, 42 ip_adapter_image=None, 43 remove_bg=None, 44 progress=gr.Progress(track_tqdm=True), 45): 46 ckpt_dir='./lora_pussinboots_logos'47 unet_sub_dir = os.path.join(ckpt_dir, "unet")48 #text_encoder_sub_dir = os.path.join(ckpt_dir, "text_encoder")49 50 if model_id is None:51 raise ValueError("Please specify the base model name or path")52 53 generator = torch.Generator(device).manual_seed(seed)54 params = {'prompt': prompt,55 'negative_prompt': negative_prompt,56 'guidance_scale': guidance_scale,57 'num_inference_steps': num_inference_steps,58 'width': width,59 'height': height,60 'generator': generator61 }62 63 if controlnet_checkbox:64 if controlnet_mode == "depth_map":65 controlnet = ControlNetModel.from_pretrained(66 "lllyasviel/sd-controlnet-depth",67 cache_dir="./models_cache",68 torch_dtype=torch_dtype69 )70 elif controlnet_mode == "pose_estimation":71 controlnet = ControlNetModel.from_pretrained(72 "lllyasviel/sd-controlnet-openpose",73 cache_dir="./models_cache",74 torch_dtype=torch_dtype75 )76 elif controlnet_mode == "normal_map":77 controlnet = ControlNetModel.from_pretrained(78 "lllyasviel/sd-controlnet-normal",79 cache_dir="./models_cache",80 torch_dtype=torch_dtype81 )82 elif controlnet_mode == "scribbles":83 controlnet = ControlNetModel.from_pretrained(84 "lllyasviel/sd-controlnet-scribble",85 cache_dir="./models_cache",86 torch_dtype=torch_dtype87 )88 else:89 controlnet = ControlNetModel.from_pretrained(90 "lllyasviel/sd-controlnet-canny",91 cache_dir="./models_cache",92 torch_dtype=torch_dtype93 )94 pipe = StableDiffusionControlNetPipeline.from_pretrained(model_id, 95 controlnet=controlnet,96 torch_dtype=torch_dtype, 97 safety_checker=None).to(device)98 params['image'] = controlnet_image99 params['controlnet_conditioning_scale'] = float(controlnet_strength)100 else:101 pipe = StableDiffusionPipeline.from_pretrained(model_id, 102 torch_dtype=torch_dtype, 103 safety_checker=None).to(device)104 105 pipe.unet = PeftModel.from_pretrained(pipe.unet, unet_sub_dir)106 #pipe.text_encoder = PeftModel.from_pretrained(pipe.text_encoder, text_encoder_sub_dir)107 108 # исправляем ошибку устанорвки lora_scale - меняем на параметр "cross_attention_kwargs"109 # pipe.unet.load_state_dict({k: lora_scale*v for k, v in pipe.unet.state_dict().items()})110 params['cross_attention_kwargs'] = {"scale": float(lora_scale)}111 #pipe.text_encoder.load_state_dict({k: lora_scale*v for k, v in pipe.text_encoder.state_dict().items()})112 113 if torch_dtype in (torch.float16, torch.bfloat16):114 pipe.unet.half()115 #pipe.text_encoder.half()116 117 if ip_adapter_checkbox:118 pipe.load_ip_adapter("h94/IP-Adapter", subfolder="models", weight_name="ip-adapter-plus_sd15.bin")119 pipe.set_ip_adapter_scale(ip_adapter_scale)120 params['ip_adapter_image'] = ip_adapter_image121 122 pipe.to(device)123 124 image = pipe(**params).images[0]125 126 # Если выбрано удаление фона127 if remove_bg:128 image = remove(image) 129 130 return image131 132examples = [133 "Puss in Boots wearing a sombrero crosses the Grand Canyon on a tightrope with a guitar.",134 "Cat wearing a sombrero crosses the Grand Canyon on a tightrope with a guitar.",135 "A cat is playing a song called ""About the Cat"" on an accordion by the sea at sunset. The sun is quickly setting behind the horizon, and the light is fading.",136 "A cat walks through the grass on the streets of an abandoned city. The camera view is always focused on the cat's face.",137 "A young lady in a Russian embroidered kaftan is sitting on a beautiful carved veranda, holding a cup to her mouth and drinking tea from the cup. With her other hand, the girl holds a saucer. The cup and saucer are painted with gzhel. Next to the girl on the table stands a samovar, and steam can be seen above it.",138 "Astronaut in a jungle, cold color palette, muted colors, detailed, 8k",139 "An astronaut riding a green horse",140 "A delicious ceviche cheesecake slice",141]142 143css = """144#col-container {145 margin: 0 auto;146 max-width: 640px;147}148"""149 150def controlnet_params(show_extra):151 return gr.update(visible=show_extra)152 153with gr.Blocks(css=css, fill_height=True) as demo:154 with gr.Column(elem_id="col-container"):155 gr.Markdown(" # Text-to-Image demo")156 157 with gr.Row():158 model_id = gr.Textbox(159 label="Model ID",160 max_lines=1,161 placeholder="Enter model id",162 value=model_id_default,163 )164 165 prompt = gr.Textbox(166 label="Prompt",167 max_lines=1,168 placeholder="Enter your prompt",169 )170 171 negative_prompt = gr.Textbox(172 label="Negative prompt",173 max_lines=1,174 placeholder="Enter your negative prompt",175 )176 177 with gr.Row():178 seed = gr.Number(179 label="Seed",180 minimum=0,181 maximum=MAX_SEED,182 step=1,183 value=42,184 )185 186 guidance_scale = gr.Slider(187 label="Guidance scale",188 minimum=0.0,189 maximum=30.0,190 step=0.1,191 value=7.0, # Replace with defaults that work for your model192 )193 with gr.Row():194 lora_scale = gr.Slider(195 label="LoRA scale",196 minimum=0.0,197 maximum=1.0,198 step=0.01,199 value=1.0,200 )201 202 num_inference_steps = gr.Slider(203 label="Number of inference steps",204 minimum=1,205 maximum=100,206 step=1,207 value=20, # Replace with defaults that work for your model208 )209 with gr.Row():210 controlnet_checkbox = gr.Checkbox(211 label="ControlNet",212 value=False213 )214 with gr.Column(visible=False) as controlnet_params:215 controlnet_strength = gr.Slider(216 label="ControlNet conditioning scale",217 minimum=0.0,218 maximum=1.0,219 step=0.01,220 value=1.0, 221 )222 controlnet_mode = gr.Dropdown(223 label="ControlNet mode",224 choices=["edge_detection", 225 "depth_map",226 "pose_estimation", 227 "normal_map",228 "scribbles"],229 value="edge_detection",230 max_choices=1231 )232 controlnet_image = gr.Image(233 label="ControlNet condition image",234 type="pil",235 format="png"236 )237 controlnet_checkbox.change(238 fn=lambda x: gr.Row.update(visible=x),239 inputs=controlnet_checkbox,240 outputs=controlnet_params241 )242 243 with gr.Row():244 ip_adapter_checkbox = gr.Checkbox(245 label="IPAdapter",246 value=False247 )248 with gr.Column(visible=False) as ip_adapter_params:249 ip_adapter_scale = gr.Slider(250 label="IPAdapter scale",251 minimum=0.0,252 maximum=1.0,253 step=0.01,254 value=1.0, 255 )256 ip_adapter_image = gr.Image(257 label="IPAdapter condition image",258 type="pil"259 )260 ip_adapter_checkbox.change(261 fn=lambda x: gr.Row.update(visible=x),262 inputs=ip_adapter_checkbox,263 outputs=ip_adapter_params264 )265 266 with gr.Accordion("Optional Settings", open=False):267 268 with gr.Row():269 width = gr.Slider(270 label="Width",271 minimum=256,272 maximum=MAX_IMAGE_SIZE,273 step=32,274 value=512, # Replace with defaults that work for your model275 )276 277 height = gr.Slider(278 label="Height",279 minimum=256,280 maximum=MAX_IMAGE_SIZE,281 step=32,282 value=512, # Replace with defaults that work for your model283 )284 285 # Удаление фона------------------------------------------------------------------------------------------------286 # Checkbox для удаления фона287 remove_bg = gr.Checkbox(288 label="Remove Background",289 value=False,290 interactive=True291 )292 # -------------------------------------------------------------------------------------------------------------293 294 295 gr.Examples(examples=examples, inputs=[prompt])296 297 298 run_button = gr.Button("Run", scale=0, variant="primary")299 result = gr.Image(label="Result", show_label=False)300 301 gr.on(302 triggers=[run_button.click],303 fn=infer,304 inputs=[305 prompt,306 negative_prompt,307 width,308 height,309 model_id,310 seed,311 guidance_scale, 312 lora_scale,313 num_inference_steps,314 controlnet_checkbox,315 controlnet_strength,316 controlnet_mode,317 controlnet_image,318 ip_adapter_checkbox,319 ip_adapter_scale,320 ip_adapter_image, 321 remove_bg, 322 ],323 outputs=[result],324 )325 326if __name__ == "__main__":327 demo.launch()