"""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'Generated image (AI-generated)') PLACEHOLDER = ('
Start typing
') 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)