whosouravsharma's picture
Move helpers into src/; propagate LangSmith trace to the backend
9b715b4 verified
Raw History Blame Contribute Delete
9.73 kB
"""Gradio UI for the DiffusionDB SD 1.5 LoRA.
Holds no model code. Generation happens in the inference Space
(whosouravsharma/diffusiondb-sd15-lora-inference), which is called here over
gradio_client. That keeps this Space light: no torch, no CUDA image, no
weights, so it builds in a minute and runs fine on free CPU while only the
backend needs a GPU.
"""
import json
import os
import time
import gradio as gr
from gradio_client import Client, handle_file # noqa: F401
from huggingface_hub import HfApi, hf_hub_download
from src.tracing import configure_logging, logger, traced_predict
configure_logging()
# ZeroGPU refuses to start a Space in which no @spaces.GPU function is
# defined. This Space runs no model — every generation happens in the
# inference Space over HTTP — so there is no real GPU work to decorate and
# this probe exists solely to satisfy that startup check. It is never called.
# Moving this Space to CPU hardware would let it go.
try:
import spaces
@spaces.GPU(duration=1)
def _zerogpu_probe():
return "ok"
except ImportError: # not installed off-Spaces; local runs do not need it
pass
BACKEND = os.environ.get(
"BACKEND_SPACE", "whosouravsharma/diffusiondb-sd15-lora-inference"
)
DATASET_REPO = "whosouravsharma/text-to-image-diffusiondb-2M"
MODEL_REPO = "whosouravsharma/diffusiondb-sd15-lora"
_client = None
# The backend sleeps after five idle minutes. Waking it means scheduling a
# T4, starting the container, importing torch and loading ~2 GB of weights.
# That is a wait, not a failure, so the UI waits it out instead of erroring.
WAKE_TIMEOUT = 600 # give up after ten minutes
POLL_INTERVAL = 5 # seconds between connection attempts
# Stages a Space cannot leave on its own. Waiting on these would just burn
# the full timeout and then report something unhelpful.
STUCK_STAGES = {"PAUSED", "RUNTIME_ERROR", "BUILD_ERROR", "CONFIG_ERROR"}
def _require_token() -> str:
token = os.environ.get("HF_TOKEN")
if not token:
raise gr.Error(
"HF_TOKEN is not set on this Space. The inference backend is "
"private and cannot be reached without it."
)
return token
def _backend_stage(token: str):
"""Current runtime stage of the backend Space, or None if unknown."""
try:
return HfApi(token=token).get_space_runtime(BACKEND).stage
except Exception:
return None
def _waking_message(stage, elapsed: int) -> str:
detail = {
"SLEEPING": "waking it up",
"BUILDING": "it is rebuilding",
"APP_STARTING": "loading the model",
"RUNNING": "loading the model",
}.get(stage, "waking it up")
return (
f"⏳ **Starting the server…** the GPU backend sleeps when idle, so "
f"{detail}. The first image takes about 40 seconds; later ones take "
f"a few.\n\n`{elapsed}s elapsed`"
)
def connect_backend():
"""Generator: yields status tuples while waiting, returns a live Client.
A sleeping Space wakes on its first HTTP request, so simply attempting
the connection starts it; the loop then keeps retrying until the app
inside is actually serving. Consumed with `yield from`, so every status
update reaches the browser as it happens.
"""
global _client
if _client is not None:
return _client
token = _require_token()
started = time.monotonic()
while True:
elapsed = int(time.monotonic() - started)
try:
# Passed positionally on purpose: this argument is named hf_token
# in gradio_client 1.x and token in 2.x, and the Space installs
# whichever version sdk_version pulls in. It is the second
# parameter in both. This call also wakes a sleeping Space.
_client = Client(BACKEND, token)
return _client
except Exception as error:
last_error = error
stage = _backend_stage(token)
if stage in STUCK_STAGES:
raise gr.Error(
f"The inference Space is {stage.lower().replace('_', ' ')} and "
f"cannot start itself. It needs to be restarted from its "
f"Settings page."
)
if time.monotonic() - started > WAKE_TIMEOUT:
raise gr.Error(
f"The inference Space did not come up within "
f"{WAKE_TIMEOUT // 60} minutes ({last_error})."
)
yield _waking_message(stage, elapsed), None, None
time.sleep(POLL_INTERVAL)
def load_examples() -> list[str]:
"""Real held-out prompts rather than invented ones.
These are the prompts the model was evaluated on and never trained on.
"""
try:
path = hf_hub_download(
DATASET_REPO, "eval_prompts.json",
repo_type="dataset", revision="v2-clean",
)
return json.loads(open(path).read())[:6]
except Exception:
return [
"a steampunk owl inside a glass jar, intricate detail",
"a cyberpunk silkscreen pop art portrait",
]
# The backend sleeps after five idle minutes. Waking it means scheduling a
# T4, starting the container, importing torch, and loading ~2 GB of weights
# before the first step runs -- around 40 seconds during which a bare spinner
# is indistinguishable from a hang. These say what is actually happening.
STATUS_IDLE = ""
STATUS_STARTING = "⏳ **Starting the server…** contacting the GPU backend."
STATUS_GENERATING = "🎨 **Generating…**"
def generate(prompt, negative, steps, guidance, lora_scale, seed,
request: gr.Request = None):
"""Streams status while the backend wakes, then returns the image.
A generator rather than a plain function so the first yield paints
immediately: the user sees that the server is starting instead of
watching a spinner for forty seconds. A sleeping backend is a wait, not
an error -- the only failures raised here are ones waiting cannot fix.
`request` is filled in by Gradio from the type annotation; it is not a UI
input and never appears in `inputs`.
"""
if not prompt or not prompt.strip():
raise gr.Error("Enter a prompt.")
yield STATUS_STARTING, None, None
# Waits out a cold backend, yielding progress the whole time.
client = yield from connect_backend()
yield STATUS_GENERATING, None, None
try:
image, used = traced_predict(
client,
prompt=prompt,
negative=negative,
steps=int(steps),
guidance=float(guidance),
lora_scale=float(lora_scale),
seed=int(seed),
request=request,
)
except Exception as error:
yield STATUS_IDLE, None, None
raise gr.Error(
f"Generation failed ({error}). The inference Space may have gone "
f"back to sleep mid-request — try again."
)
yield STATUS_IDLE, image, used
with gr.Blocks(title="DiffusionDB SD 1.5 LoRA") as demo:
gr.Markdown(
f"# DiffusionDB SD 1.5 LoRA\n"
f"Stable Diffusion 1.5 with a rank-32 LoRA trained on 13,598 "
f"prompt-image pairs from DiffusionDB "
f"([model]({'https://huggingface.co/' + MODEL_REPO}) · "
f"[dataset]({'https://huggingface.co/datasets/' + DATASET_REPO})).\n\n"
f"The adapter shifts style toward the DiffusionDB aesthetic — the "
f"keyword-heavy *artstation / intricate / octane render* idiom its "
f"users prompted with. Set **LoRA strength to 0** to render vanilla "
f"SD 1.5 at the same seed for a direct comparison."
)
with gr.Row():
with gr.Column(scale=3):
prompt = gr.Textbox(
label="Prompt", lines=3,
placeholder="a steampunk owl inside a glass jar, intricate detail",
)
gr.Examples(
examples=[[p] for p in load_examples()],
inputs=[prompt],
label="Example prompts",
examples_per_page=6,
)
negative = gr.Textbox(
label="Negative prompt", lines=1,
placeholder="blurry, watermark, text",
)
run = gr.Button("Generate", variant="primary")
with gr.Accordion("Settings", open=False):
lora_scale = gr.Slider(
0.0, 1.5, value=1.0, step=0.05, label="LoRA strength",
info="0 = base SD 1.5, 1 = as trained, >1 exaggerates",
)
steps = gr.Slider(10, 50, value=25, step=1, label="Steps")
guidance = gr.Slider(
1.0, 15.0, value=7.5, step=0.5, label="Guidance scale",
)
seed = gr.Number(
value=-1, precision=0, label="Seed", info="-1 for random",
)
with gr.Column(scale=4):
status = gr.Markdown(STATUS_IDLE)
output = gr.Image(label="Output", height=512)
used_seed = gr.Number(label="Seed used", interactive=False)
gr.Markdown(
"Trained on Stable Diffusion 1.x outputs, so it reproduces that "
"model's artifacts along with its style. Images are 512×512, the "
"resolution it was fine-tuned at."
)
for trigger in (run.click, prompt.submit):
trigger(
generate,
inputs=[prompt, negative, steps, guidance, lora_scale, seed],
outputs=[status, output, used_seed],
)
if __name__ == "__main__":
demo.queue(max_size=12).launch(
server_name="0.0.0.0",
server_port=int(os.environ.get("PORT", 7860)),
)