SDXL-IP-Adapter / app.py
cereeenn120's picture
Initial commit
b5f9f05
Raw
History Blame Contribute Delete
2.38 kB
import spaces
import torch
from diffusers import (
StableDiffusionXLPipeline,
StableDiffusionXLImg2ImgPipeline
)
from PIL import Image
import gradio as gr
MODEL_ID = "stabilityai/stable-diffusion-xl-base-1.0"
IP_ADAPTER_REPO = "h94/IP-Adapter"
# Load pipelines
pipe_txt2img = StableDiffusionXLPipeline.from_pretrained(
MODEL_ID,
torch_dtype=torch.float32
).to("cuda" if torch.cuda.is_available() else "cpu")
pipe_img2img = StableDiffusionXLImg2ImgPipeline.from_pretrained(
MODEL_ID,
torch_dtype=torch.float32
).to("cuda" if torch.cuda.is_available() else "cpu")
# Load IP-Adapter
for p in [pipe_txt2img, pipe_img2img]:
p.load_ip_adapter(
IP_ADAPTER_REPO,
subfolder="sdxl_models",
weight_name="ip-adapter_sdxl.bin"
)
p.set_ip_adapter_scale(0.7)
@spaces.GPU
def generate(prompt, init_image, ref_images):
ref_list = []
if ref_images:
for item in ref_images:
img = item[0] if isinstance(item, tuple) else item
if img is not None:
if isinstance(img, str):
img = Image.open(img)
ref_list.append(img.convert("RGB"))
common_kwargs = {
"prompt": prompt,
"ip_adapter_image": [ref_list] if len(ref_list) > 0 else None,
"num_inference_steps": 30,
"guidance_scale": 8.0
}
# CASE 1: Img2Img
if init_image is not None:
init_image = init_image.convert("RGB")
result = pipe_img2img(
image=init_image,
strength=0.5,
**common_kwargs
).images[0]
# CASE 2: Text2Img
else:
result = pipe_txt2img(
**common_kwargs
).images[0]
return result
with gr.Blocks() as demo:
gr.Markdown("# SDXL Character Generator (IP-Adapter + Multi-Image + Img2Img)")
with gr.Row():
prompt = gr.Textbox(label="Prompt", placeholder="cinematic portrait, ultra detailed, same person")
with gr.Row():
init_image = gr.Image(label="Init Image (Img2Img)", type="pil")
ref_images = gr.Gallery(label="Reference Images (Multi-image)", columns=3, object_fit="contain")
generate_btn = gr.Button("Generate")
output = gr.Image(label="Output")
generate_btn.click(
fn=generate,
inputs=[prompt, init_image, ref_images],
outputs=output
)
demo.launch()