Spaces:
Paused
Paused
File size: 6,143 Bytes
573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c 3c08420 573e34c | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 | # server.py (robust: try ONNX if available, fallback to diffusers)
import os
import time
import base64
from io import BytesIO
from pathlib import Path
from fastapi import FastAPI
from pydantic import BaseModel
import numpy as np
from PIL import Image
import torch
# optional imports; may fail if packages not installed but we'll handle that
try:
import onnxruntime as ort
ONNX_AVAILABLE = True
except Exception:
ONNX_AVAILABLE = False
from transformers import CLIPTokenizer
from diffusers import LMSDiscreteScheduler, StableDiffusionPipeline
from huggingface_hub import snapshot_download
app = FastAPI()
MODEL_DIR = Path(os.environ.get("MODEL_DIR", "onnx_models_quant"))
HF_ONNX_REPO = os.environ.get("HF_ONNX_REPO") # optional HF repo id to download ONNX artifacts
HF_TOKEN = os.environ.get("HF_TOKEN") # optional token for private repo
INTRA_THREADS = int(os.environ.get("INTRA_THREADS", "4"))
# globals
use_onnx = False
onnx_sessions = {}
tokenizer = None
diffusers_pipe = None
scheduler = LMSDiscreteScheduler(beta_start=0.00085, beta_end=0.012, beta_schedule="scaled_linear", num_train_timesteps=1000)
def download_onnx_if_needed():
if MODEL_DIR.exists() and any((MODEL_DIR / f).exists() for f in ("text_encoder.onnx", "unet.onnx", "vae_decoder.onnx")):
return True
if not HF_ONNX_REPO:
return False
# attempt to download snapshot into MODEL_DIR
try:
snapshot_download(repo_id=HF_ONNX_REPO, cache_dir=str(MODEL_DIR), token=HF_TOKEN)
return MODEL_DIR.exists()
except Exception as e:
print("Failed to download ONNX artifacts from HF Hub:", e)
return False
def init_onnx():
global onnx_sessions, tokenizer, use_onnx
if not ONNX_AVAILABLE:
print("onnxruntime not available; skipping ONNX init")
use_onnx = False
return
# Ensure files exist or try download
ok = download_onnx_if_needed()
if not ok:
print(f"ONNX models not found in {MODEL_DIR} and no HF_ONNX_REPO configured; skipping ONNX")
use_onnx = False
return
# load tokenizer from model dir if available else fallback
try:
tokenizer = CLIPTokenizer.from_pretrained(str(MODEL_DIR))
except Exception:
tokenizer = CLIPTokenizer.from_pretrained("openai/clip-vit-large-patch14")
sess_opts = ort.SessionOptions()
sess_opts.intra_op_num_threads = INTRA_THREADS
sess_opts.inter_op_num_threads = 1
sess_opts.graph_optimization_level = ort.GraphOptimizationLevel.ORT_ENABLE_ALL
onnx_sessions["text"] = ort.InferenceSession(str(MODEL_DIR / "text_encoder.onnx"), sess_opts, providers=["CPUExecutionProvider"])
onnx_sessions["unet"] = ort.InferenceSession(str(MODEL_DIR / "unet.onnx"), sess_opts, providers=["CPUExecutionProvider"])
onnx_sessions["vae"] = ort.InferenceSession(str(MODEL_DIR / "vae_decoder.onnx"), sess_opts, providers=["CPUExecutionProvider"])
use_onnx = True
print("Loaded ONNX sessions from", MODEL_DIR)
def init_diffusers():
global diffusers_pipe, tokenizer
if diffusers_pipe is None:
model_id = os.environ.get("DIFFUSERS_MODEL", "runwayml/stable-diffusion-v1-5")
print("Loading diffusers pipeline:", model_id)
diffusers_pipe = StableDiffusionPipeline.from_pretrained(model_id)
device = "cuda" if torch.cuda.is_available() else "cpu"
diffusers_pipe.to(device)
tokenizer = diffusers_pipe.tokenizer
# Initialize best available backend
init_onnx()
if not use_onnx:
init_diffusers()
class Prompt(BaseModel):
prompt: str
steps: int = 20
scale: float = 7.5
# ONNX helpers
def run_text_encoder_onnx(prompt: str):
toks = tokenizer(prompt, return_tensors="pt", padding="max_length", max_length=77, truncation=True)
input_ids = toks["input_ids"].cpu().numpy()
attention_mask = toks["attention_mask"].cpu().numpy()
ort_inputs = {"input_ids": input_ids, "attention_mask": attention_mask}
out = onnx_sessions["text"].run(None, ort_inputs)
return out[0].astype(np.float32)
def unet_predict_onnx(latent: np.ndarray, t: int, encoder_hidden_states: np.ndarray):
ort_inputs = {
"latent": latent.astype(np.float32),
"timestep": np.array([int(t)], dtype=np.int64),
"encoder_hidden_states": encoder_hidden_states.astype(np.float32)
}
out = onnx_sessions["unet"].run(None, ort_inputs)
return out[0].astype(np.float32)
def decode_vae_onnx(latents: np.ndarray):
out = onnx_sessions["vae"].run(None, {"latent": latents.astype(np.float32)})
return out[0].astype(np.float32)
# Generation endpoints use either ONNX or diffusers
@app.post("/generate")
def generate(p: Prompt):
start = time.time()
steps = int(p.steps)
if use_onnx:
scheduler.set_timesteps(steps)
encoder_hidden_states = run_text_encoder_onnx(p.prompt)
latents = np.random.randn(1, 4, 64, 64).astype(np.float32)
for t in scheduler.timesteps:
noise_pred = unet_predict_onnx(latents, int(t), encoder_hidden_states)
noise_pred_t = torch.from_numpy(noise_pred)
latents_t = torch.from_numpy(latents)
step_output = scheduler.step(noise_pred_t, t, latents_t, return_dict=False)
latents = step_output[0].numpy()
images = decode_vae_onnx(latents)
image = images[0]
image = (np.clip(image, -1, 1) + 1.0) / 2.0
image = (image * 255).round().astype(np.uint8)
image = np.transpose(image, (1, 2, 0))
pil = Image.fromarray(image)
elapsed = time.time() - start
return {"image_base64": base64.b64encode(pil_to_bytes(pil)).decode("utf-8"), "elapsed_s": elapsed, "mode": "onnx"}
else:
# diffusers fallback
out = diffusers_pipe(p.prompt, num_inference_steps=steps, guidance_scale=float(p.scale))
pil = out.images[0]
elapsed = time.time() - start
return {"image_base64": base64.b64encode(pil_to_bytes(pil)).decode("utf-8"), "elapsed_s": elapsed, "mode": "diffusers"}
def pil_to_bytes(pil):
buf = BytesIO()
pil.save(buf, format="PNG")
return buf.getvalue() |