ysharma's picture
ysharma HF Staff
Move to ML-Intern-lab: new repo names
7b14807 verified
Raw History Blame Contribute Delete
15.2 kB
"""Agate 4-step, live: an image on every keystroke, on ZeroGPU or any dedicated GPU (or CPU).
ML-Intern-lab/agate-preview-002-4step is an unofficial 4-step, guidance-free distillation of LogoLabs' Agate Preview 002
(256 px, 0.19B generator + 68M text encoder). One render is 4 network passes plus a VAE decode.
Two serving modes, picked from the hardware at startup:
* ZeroGPU: every @spaces.GPU call lands in a GPU worker that spends ~1.2 s warming CUDA up before its first render
(measured), which would dominate a 4-step model. So the first keystroke opens ONE GPU call that stays open while
you type: each keystroke writes the newest prompt to a small per-browser-session file (a plain, GPU-free handler),
and the session generator polls that file and streams a new image whenever the prompt or seed changes. It closes
after IDLE_S seconds without changes, returning the GPU; the next keystroke opens a new one.
* Dedicated GPU (T4, L4, A10G, L40S, A100, H100, ...) or CPU: the model stays resident and warm, so every keystroke is
simply one render. trigger_mode="always_last" drops keystrokes that arrive during a render except the newest, and
on GPUs each denoising step is replayed from a CUDA graph (recorded per text-length bucket at startup/first use).
Precision follows the GPU: bf16 on Ampere and newer (as Agate was trained), fp16 on Turing (T4 has no bf16), fp32 on
CPU. The AgatePipeline hard-codes bf16 autocast, so this app drives its modules with its own copy of the 4-step
sampler (identical maths: uniform Euler, t = 0 noise, v = model(z, t, ctx, mask), no CFG).
"""
import base64
import io
import json
import os
import sys
import tempfile
import threading
import time
from pathlib import Path
import gradio as gr
import spaces
import torch
import torch.nn.functional as F
from huggingface_hub import snapshot_download
from PIL import Image, ImageFilter
from transformers import pipeline
REPO = "ML-Intern-lab/agate-preview-002-4step"
STEPS = 4
IDLE_S = 6.0 # ZeroGPU: a live session ends after this long without a prompt/seed change
SESSION_S = 50.0 # ZeroGPU: hard cap per GPU call (duration below leaves headroom)
POLL_S = 0.01
BUCKETS = (64, 128, 256, 512)
ON_ZEROGPU = bool(os.environ.get("SPACES_ZERO_GPU"))
HAS_CUDA = torch.cuda.is_available() # also True on ZeroGPU (spaces patches torch)
DEVICE = torch.device("cuda" if HAS_CUDA else "cpu")
if ON_ZEROGPU:
DTYPE = torch.bfloat16 # ZeroGPU runs on recent datacentre GPUs
elif HAS_CUDA:
DTYPE = torch.bfloat16 if torch.cuda.get_device_capability(0)[0] >= 8 else torch.float16
else:
DTYPE = torch.float32
USE_GRAPHS = HAS_CUDA and not ON_ZEROGPU # ZeroGPU workers are short-lived; graphs would be re-recorded each call
AUTOCAST = DTYPE != torch.float32
LIVE_DIR = Path(tempfile.gettempdir()) / "agate_live"
LIVE_DIR.mkdir(exist_ok=True)
path = snapshot_download(REPO, allow_patterns=["agate/*", "text_encoder/*", "config.json", "generator.safetensors"])
sys.path.insert(0, path)
from agate import AgatePipeline # noqa: E402
from agate.marking import add_watermark, provenance # noqa: E402
# The pipeline loads the generator, Ettin text encoder and SD-VAE; we use its modules, in our dtype.
pipe = AgatePipeline.from_pretrained(path, device=str(DEVICE), cuda_graphs=False)
NET = pipe.model.model.to(dtype=DTYPE) # FCDMThinker2 (pipe.model is the pipeline's eager wrapper)
TEXT = pipe.text # EttinTextEncoder: .tokenize() and .model (ModernBERT)
TEXT.model.to(dtype=DTYPE)
VAE, VAE_DIV = pipe.vae, pipe.vae_div
if DEVICE.type == "cuda":
torch.backends.cuda.enable_cudnn_sdp(False) # the attention kernels Agate was trained with
safety = pipeline("image-classification", model="Falconsai/nsfw_image_detection", device=DEVICE)
GPU_LOCK = threading.Lock() # dedicated GPU: one render at a time (CUDA graphs share buffers)
def _autocast():
return torch.autocast(DEVICE.type, dtype=DTYPE, enabled=AUTOCAST)
class StepRunner:
"""net(z, t, ctx, mask) -> velocity (fp32); on dedicated GPUs recorded as a CUDA graph per input shape."""
def __init__(self, net, graphs: bool):
self.net, self.graphs, self.cache = net, graphs, {}
def _run(self, z, t, ctx, mask):
with _autocast():
return self.net(z, t, ctx, mask)
def __call__(self, z, t, ctx, mask):
if not self.graphs:
return self._run(z, t, ctx, mask).float()
key = (tuple(z.shape), tuple(ctx.shape))
if key not in self.cache:
static = [z.clone(), t.clone(), ctx.clone(), mask.clone()]
side = torch.cuda.Stream()
side.wait_stream(torch.cuda.current_stream())
with torch.cuda.stream(side):
for _ in range(2): # warm-up: cuDNN autotune, allocator
self._run(*static)
torch.cuda.current_stream().wait_stream(side)
graph = torch.cuda.CUDAGraph()
with torch.cuda.graph(graph):
out = self._run(*static)
self.cache[key] = (graph, static, out)
graph, static, out = self.cache[key]
for dst, src in zip(static, (z, t, ctx, mask)):
dst.copy_(src)
graph.replay()
return out.float()
STEP = StepRunner(NET, USE_GRAPHS)
@torch.no_grad()
def _sample(prompt: str, seed: int) -> Image.Image:
ids, mask = TEXT.tokenize([prompt])
ids, mask = ids.to(DEVICE), mask.to(DEVICE)
with _autocast():
ctx = TEXT.model(input_ids=ids, attention_mask=mask).last_hidden_state.float()
L = next((b for b in BUCKETS if b >= ctx.shape[1]), BUCKETS[-1])
ctx, mask = F.pad(ctx, (0, 0, 0, L - ctx.shape[1])), F.pad(mask, (0, L - mask.shape[1]))
gen = torch.Generator(device=DEVICE).manual_seed(int(seed))
hw = pipe.cfg["latent_hw"]
z = torch.randn(1, 4, hw, hw, device=DEVICE, generator=gen)
dt = 1.0 / STEPS
for i in range(STEPS):
t = torch.full((1,), i * dt, device=DEVICE)
z = z + dt * STEP(z, t, ctx, mask)
x = VAE.decode((z / VAE_DIV).to(VAE.dtype)).sample
x = ((x.float().clamp(-1, 1) + 1) * 127.5).round().byte().permute(0, 2, 3, 1).cpu().numpy()
return Image.fromarray(x[0])
def _gpu_name() -> str:
return torch.cuda.get_device_name(0) if DEVICE.type == "cuda" else "CPU"
def _render(prompt: str, seed: int):
"""-> (PIL image, ms, flagged). Must run where the GPU is attached (inside @spaces.GPU on ZeroGPU)."""
t0 = time.perf_counter()
img = _sample(prompt, seed)
img = add_watermark(img) # AGATE002, as the pipeline does by default
img.info.update(provenance())
ms = (time.perf_counter() - t0) * 1000
flagged = {r["label"]: r["score"] for r in safety(img)}.get("nsfw", 0.0) > 0.5
if flagged:
img = img.filter(ImageFilter.GaussianBlur(24))
return img, ms, flagged
def _html(img: Image.Image) -> str:
"""The image inline as a WebP data URL: it arrives inside the event message, so the browser needs no second
request (one network round trip per keystroke saved). The invisible watermark is in the pixels."""
buf = io.BytesIO()
img.save(buf, format="WEBP", quality=90, method=4)
b64 = base64.b64encode(buf.getvalue()).decode()
return (f'<img src="data:image/webp;base64,{b64}" alt="Generated image (AI-generated)" '
'style="width:100%;max-width:512px;aspect-ratio:1;display:block;border-radius:8px">')
PLACEHOLDER = ('<div style="width:100%;max-width:512px;aspect-ratio:1;border-radius:8px;display:grid;place-items:center;'
'background:var(--block-background-fill);color:var(--body-text-color-subdued)">Start typing</div>')
def _stats(ms: float, flagged: bool, n: int | None = None) -> str:
s = f"**{ms:.0f} ms** per image on {_gpu_name()} ({str(DTYPE).removeprefix('torch.')})"
if n is not None:
s += f" · live session: {n} image{'s' if n != 1 else ''}, closes after {IDLE_S:.0f} s without typing"
if flagged:
s += " · blurred by the safety filter"
return s
# ---------------- one render (all hardware) ----------------
@spaces.GPU(duration=10)
def _single(prompt: str, seed: int):
if ON_ZEROGPU:
torch.backends.cudnn.benchmark = False # autotuning in a fresh ZeroGPU worker costs seconds
with GPU_LOCK:
img, ms, flagged = _render(prompt, seed)
return _html(img), _stats(ms, flagged)
def generate(prompt: str, seed: float):
prompt = (prompt or "").strip()
if not prompt:
return gr.skip(), gr.skip()
return _single(prompt, int(seed or 0))
# ---------------- ZeroGPU live sessions ----------------
def _state_file(sid: str) -> Path:
return LIVE_DIR / f"{sid}.json"
def _read_state(sid: str):
try:
d = json.loads(_state_file(sid).read_text())
return d["prompt"], int(d["seed"])
except Exception:
return None
def push_prompt(prompt: str, seed: float, request: gr.Request):
"""GPU-free: record the newest prompt for this browser session (atomic replace)."""
p = _state_file(request.session_hash)
tmp = p.with_suffix(".tmp")
tmp.write_text(json.dumps({"prompt": prompt or "", "seed": int(seed or 0)}))
os.replace(tmp, p)
@spaces.GPU(duration=60)
def _live_session(sid: str, prompt: str, seed: int):
torch.backends.cudnn.benchmark = False
last, n, idle_since, t_end = None, 0, time.monotonic(), time.monotonic() + SESSION_S
while time.monotonic() < t_end:
cur = _read_state(sid) or (prompt, seed)
if cur != last and cur[0].strip():
last = cur
img, ms, flagged = _render(cur[0].strip(), cur[1])
n += 1
idle_since = time.monotonic()
yield _html(img), _stats(ms, flagged, n)
elif time.monotonic() - idle_since > IDLE_S:
break
else:
time.sleep(POLL_S)
def live(prompt: str, seed: float, request: gr.Request):
yield from _live_session(request.session_hash, prompt or "", int(seed or 0))
def new_seed():
return int.from_bytes(os.urandom(4), "little") % 1_000_000
# ---------------- warm-up on dedicated hardware ----------------
if HAS_CUDA and not ON_ZEROGPU:
t0 = time.perf_counter()
for words in (1, 40): # record the graphs for the two common text buckets (64, 128)
_render("warm up " * words, 0)
print(f"[agate] warm on {_gpu_name()} ({DTYPE}) in {time.perf_counter() - t0:.1f} s; CUDA graphs: {USE_GRAPHS}")
EXAMPLES = [
"a red cube on top of a blue sphere",
"a lighthouse on a rocky coast under a stormy sky",
"a minimalist logo of a fox head, orange, flat design, white background",
"a cabin in a snowy forest at night, warm light in the windows",
"an astronaut riding a horse on the moon",
"a watercolor painting of a harbor with red sailboats",
]
CSS = ".gradio-container { max-width: 1000px !important; margin: 0 auto; }"
if ON_ZEROGPU:
HOW = ("The first keystroke opens a GPU session (about a second to warm up); while you keep typing, each change is "
f"drawn as soon as the GPU is free. The session closes after {IDLE_S:.0f} s without typing.")
else:
HOW = f"Running on {_gpu_name()}: the model stays warm, and every keystroke is one render."
with gr.Blocks(title="Agate 4-step live") as demo:
gr.Markdown(
"# Agate 4-step, live\n"
"Start typing: the image redraws as you type. "
"[ML-Intern-lab/agate-preview-002-4step](https://huggingface.co/ML-Intern-lab/agate-preview-002-4step) is an **unofficial** 4-step distillation of "
"[Agate Preview 002](https://huggingface.co/Logolabs/agate-preview-002) (LogoLabs, MIT): 4 network passes per image "
"instead of 100, GenEval 0.536 against 0.563 for the 50-step teacher. The seed stays fixed while you type, so you can "
"watch the picture follow your words."
)
with gr.Row():
with gr.Column(scale=1):
prompt = gr.Textbox(label="Prompt", value=EXAMPLES[1], lines=3, autofocus=True,
placeholder="Type anything...")
with gr.Row():
seed = gr.Number(label="Seed", value=0, precision=0, minimum=0, maximum=999_999)
dice = gr.Button("New seed", variant="secondary")
stats = gr.Markdown()
gr.Examples(EXAMPLES, inputs=prompt, cache_examples=False)
with gr.Column(scale=1):
out = gr.HTML(PLACEHOLDER, label="256 × 256, shown at 2×", show_label=True)
gr.Markdown(
f"{HOW}\n\n"
"256 × 256 only. No negative prompts (guidance is baked into the weights). Agate's weaknesses remain: exact text, "
"counts above three, negation. Images carry Agate's invisible watermark and an AI-generated note in the PNG "
"metadata; an NSFW classifier blurs flagged outputs. Not affiliated with LogoLabs. "
"In-browser version: [ML-Intern-lab/agate-preview-002-4step-webgpu](https://huggingface.co/spaces/ML-Intern-lab/agate-preview-002-4step-webgpu)."
)
if ON_ZEROGPU:
# 1) every keystroke: record the newest prompt (no GPU, no queue)
for ev, vis in ((prompt.input, "public"), (seed.input, "private")):
ev(push_prompt, inputs=[prompt, seed], outputs=None, queue=False, show_progress="hidden",
api_name="push", api_visibility=vis)
# 2) every keystroke: make sure a live session is running; trigger_mode="once" ignores triggers while one is
# open, and the open session picks up the newest prompt by itself
for ev, vis in ((prompt.input, "public"), (seed.input, "private")):
ev(live, inputs=[prompt, seed], outputs=[out, stats], trigger_mode="once", show_progress="hidden",
api_name="live", api_visibility=vis)
else:
for ev, vis in ((prompt.input, "public"), (seed.input, "private")):
# queue=False: one HTTP request per keystroke instead of queue join + event stream (two network round
# trips); GPU_LOCK serializes renders, and always_last coalescing happens in the browser
ev(generate, inputs=[prompt, seed], outputs=[out, stats], trigger_mode="always_last", queue=False,
show_progress="hidden", api_name="keystroke", api_visibility=vis)
dice.click(new_seed, outputs=seed).then(generate, inputs=[prompt, seed], outputs=[out, stats],
show_progress="hidden", api_visibility="private")
demo.load(generate, inputs=[prompt, seed], outputs=[out, stats], show_progress="hidden", api_name="generate")
if __name__ == "__main__":
demo.queue(default_concurrency_limit=None if ON_ZEROGPU else 1).launch(css=CSS, ssr_mode=False)