Spaces:
Running on Zero
Running on Zero
Download app.py from ML-Intern-lab/agate-preview-002-4step-live: direct link, hf CLI and curl.
- Browser
- Download file 15.2 kB
-
https://huggingface.co/spaces/ML-Intern-lab/agate-preview-002-4step-live/resolve/main/app.py
- Command line
-
hf download hf://spaces/ML-Intern-lab/agate-preview-002-4step-live/app.py
-
curl -L -o app.py https://huggingface.co/spaces/ML-Intern-lab/agate-preview-002-4step-live/resolve/main/app.py
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) | |
| 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) ---------------- | |
| 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) | |
| 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) | |