"""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'')
PLACEHOLDER = ('