multimodalart's picture
multimodalart HF Staff
Update handler.py
4363538 verified
Raw
History Blame Contribute Delete
3.79 kB
from typing import Dict, List, Any
import torch
from torch import autocast
from huggingface_hub import hf_hub_download
from diffusers import DiffusionPipeline
import base64
from io import BytesIO
from safetensors.torch import load_file
from cog_sdxl.dataset_and_utils import TokenEmbeddingsHandler
from cog_sdxl.no_init import no_init_or_tensor
from diffusers.models.attention_processor import LoRAAttnProcessor2_0
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("device ~>", device)
class EndpointHandler:
def __init__(self, path=""):
print("path ~>", path)
self.pipe = DiffusionPipeline.from_pretrained(
"stabilityai/stable-diffusion-xl-base-1.0",
torch_dtype=torch.float16 if device.type == "cuda" else None,
variant="fp16",
).to(device)
lora_path = hf_hub_download(repo_id="SvenN/sdxl-emoji", filename="lora.safetensors", repo_type="model")
embeddings_path = hf_hub_download(repo_id="SvenN/sdxl-emoji", filename="embeddings.pti", repo_type="model")
#Load the LoRA into the UNet
unet = self.pipe.unet
tensors = load_file(lora_path)
unet_lora_attn_procs = {}
name_rank_map = {}
for tk, tv in tensors.items():
# up is N, d
tensors[tk] = tv.half()
if tk.endswith("up.weight"):
proc_name = ".".join(tk.split(".")[:-3])
r = tv.shape[1]
name_rank_map[proc_name] = r
for name, attn_processor in unet.attn_processors.items():
cross_attention_dim = (
None
if name.endswith("attn1.processor")
else unet.config.cross_attention_dim
)
if name.startswith("mid_block"):
hidden_size = unet.config.block_out_channels[-1]
elif name.startswith("up_blocks"):
block_id = int(name[len("up_blocks.")])
hidden_size = list(reversed(unet.config.block_out_channels))[
block_id
]
elif name.startswith("down_blocks"):
block_id = int(name[len("down_blocks.")])
hidden_size = unet.config.block_out_channels[block_id]
with no_init_or_tensor():
module = LoRAAttnProcessor2_0(
hidden_size=hidden_size,
cross_attention_dim=cross_attention_dim,
rank=name_rank_map[name],
).half()
unet_lora_attn_procs[name] = module.to("cuda", non_blocking=True)
unet.set_attn_processor(unet_lora_attn_procs)
unet.load_state_dict(tensors, strict=False)
#Load the text embeddings into the text encoder/tokenizer
handler = TokenEmbeddingsHandler(
[self.pipe.text_encoder, self.pipe.text_encoder_2], [self.pipe.tokenizer, self.pipe.tokenizer_2]
)
handler.load_embeddings(embeddings_path)
def __call__(self, data: Any) -> List[List[Dict[str, float]]]:
"""
Args:
data (:obj:):
includes the input data and the parameters for the inference.
Return:
A :obj:`dict`:. base64 encoded image
"""
inputs = data.pop("inputs", data)
# Automatically add trigger tokens to the beginning of the prompt
images = self.pipe(
inputs,
cross_attention_kwargs={"scale": 0.6},
**data['parameters']
).images
image = images[0]
return image
if __name__ == "__main__":
handler = EndpointHandler()
print(handler)
output = handler({"inputs": "emoji of a tiger face, white background"})
print(output)