agate-preview-003 / agate /pipeline.py
Incorporo-user's picture
Agate Preview 003: weights, pipeline, model card, figures
4770b6a verified
Raw History Blame
9.74 kB
"""AgatePipeline (Preview 003): prompt -> 512x512 image (or 256x256).
from agate import AgatePipeline
pipe = AgatePipeline.from_pretrained("Logolabs/agate-preview-003", device="cuda")
images = pipe("a red cube on top of a blue sphere", seed=0) # 512 x 512
images = pipe('a shop sign that says "OPEN"', seed=0, resolution=256) # 256 x 256
Preview 003 is the multi-resolution model (arch fcdm_t2mr). It was trained with a prompt pipeline, and
sampling uses the same one (prompt_norm.py):
* normaliser: whitespace, SHOUTED or Title Cased prompts lower-cased outside quotes, number words and
small digits made canonical ("THREE", "3" -> "three");
* negatives: "without X" / "no X" are cut from the prompt and become the negative prompt (CFG pushes
away from it; the model itself barely reacts to negation);
* spelling: text inside double quotes is also spelled out letter by letter after " || spell: ", so the
text encoder sees one token per letter;
* count code: the value of every object count ("three cats") is passed to the model at those tokens.
normalize=False / spell=False switch the steps off (not how the model was trained).
Sampler: Euler from t=0 (noise) to t=1 (data), 50 steps, classifier-free guidance 3 against the empty
prompt (or the negatives), velocity prediction. 512 px uses the SD3 resolution shift 2 on the time grid,
as trained; 256 px uses no shift. On CUDA each denoising step is recorded once as a CUDA graph and
replayed; on CPU it runs eagerly.
"""
from __future__ import annotations
import json
from pathlib import Path
import torch
import torch.nn.functional as F
from . import prompt_norm as PN
from .fcdm_thinker2_mr import FCDMThinker2MR
from .text_encoder import EttinTextEncoder
from .marking import add_watermark, provenance
BUCKETS = (64, 128, 256, 512)
def _resolve(repo_or_dir: str) -> Path:
p = Path(repo_or_dir)
if p.is_dir():
return p
from huggingface_hub import snapshot_download
return Path(snapshot_download(repo_or_dir))
def _pad_to(ctx: torch.Tensor, mask: torch.Tensor, L: int):
return F.pad(ctx, (0, 0, 0, L - ctx.shape[1])), F.pad(mask, (0, L - mask.shape[1]))
def _bucket(*lengths: int) -> int:
need = max(lengths)
return next((b for b in BUCKETS if b >= need), BUCKETS[-1])
def shift_t(t: float, shift: float) -> float:
"""SD3's resolution shift in Agate's convention (t = 0 noise): with the noise level s = 1 - t,
s' = shift * s / (1 + (shift - 1) * s). shift = 1 is the identity."""
if shift == 1.0:
return t
s = 1 - t
return 1 - shift * s / (1 + (shift - 1) * s)
class _Graphed:
"""model(z, t, ctx, mask, counts) recorded as a CUDA graph per input shape and replayed."""
def __init__(self, model):
self.model, self.cache = model, {}
def _run(self, z, t, ctx, mask, counts):
with torch.autocast("cuda", dtype=torch.bfloat16):
return self.model(z, t, ctx, mask, counts=counts)
def __call__(self, z, t, ctx, mask, counts):
key = (tuple(z.shape), tuple(ctx.shape))
if key not in self.cache:
static = [z.clone(), t.clone(), ctx.clone(), mask.clone(), counts.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, counts)):
dst.copy_(src)
graph.replay()
return out.float()
class _Eager:
def __init__(self, model):
self.model = model
def __call__(self, z, t, ctx, mask, counts):
with torch.autocast(z.device.type, dtype=torch.bfloat16, enabled=z.device.type == "cuda"):
return self.model(z, t, ctx, mask, counts=counts).float()
class AgatePipeline:
def __init__(self, root: Path, device: str = "cuda", fast_vae: bool = False, cuda_graphs: bool = True):
from diffusers import AutoencoderKL, AutoencoderTiny
self.root, self.device = Path(root), torch.device(device)
self.cfg = json.loads((self.root / "config.json").read_text())
on_cuda = self.device.type == "cuda"
if on_cuda:
torch.backends.cudnn.benchmark = True
torch.backends.cuda.enable_cudnn_sdp(False) # the attention kernels Agate was trained with
self._graphs = on_cuda and cuda_graphs
self._dtype = torch.bfloat16 if on_cuda else torch.float32
self.model = self._load_generator(self.root / "generator.safetensors")
self.text = self._load_text(self.root / self.cfg["text_encoder"], self.cfg["text_max_len"])
vdtype = torch.float16 if on_cuda else torch.float32
if fast_vae:
self.vae = AutoencoderTiny.from_pretrained(self.cfg["fast_vae"], torch_dtype=vdtype)
self.vae_div = 1.0 # TAESD decodes the scaled latents directly
else:
self.vae = AutoencoderKL.from_pretrained(self.cfg["vae"], torch_dtype=vdtype)
self.vae_div = self.cfg["vae_scale"]
self.vae = self.vae.to(self.device).eval()
@classmethod
def from_pretrained(cls, repo_or_dir: str = "Logolabs/agate-preview-003", device: str = "cuda", **kw):
return cls(_resolve(repo_or_dir), device=device, **kw)
def _load_generator(self, path: Path):
from safetensors.torch import load_file
m = FCDMThinker2MR(**self.cfg["model_kw"])
m.load_state_dict(load_file(str(path)))
m = m.to(self.device, dtype=self._dtype).eval()
if self.device.type == "cuda":
m = m.to(memory_format=torch.channels_last)
return _Graphed(m) if self._graphs else _Eager(m)
def _load_text(self, path: Path, max_len: int) -> EttinTextEncoder:
text = EttinTextEncoder(str(path), self.device, max_len=max_len)
text.model.to(dtype=self._dtype)
return text
def prepare(self, prompt: str, negative_prompt: str = "", normalize: bool = True, spell: bool = True):
"""The prompt pipeline alone: -> (prompt as the text encoder sees it, negative prompt)."""
negs = []
if normalize:
prompt, negs = PN.split_negatives(PN.normalize(prompt))
if spell:
prompt = PN.add_spelling(prompt)
negative = ", ".join(n for n in [negative_prompt.strip(), *negs] if n)
return prompt, negative
@torch.no_grad()
def __call__(self, prompt: str, negative_prompt: str = "", seed: int = 0, steps: int = 50, cfg: float = 3.0,
num_images: int = 1, resolution: int | None = None, normalize: bool = True, spell: bool = True,
watermark: bool = True, metadata: bool = True, record_prompt: bool = False):
"""resolution: 512 (default) or 256. negative_prompt is combined with the negatives the normaliser
cut from the prompt ("without X", "no X"); when both are empty the unconditional prompt is "".
watermark (default True): an invisible dwtDct watermark (marking.PAYLOAD) in every image; metadata (default
True): provenance entries in img.info ("ai_generated", "generator", "model"); record_prompt (default False)
also stores the prompt and seed there. Save with agate.save(img, path) to keep the entries in a PNG."""
n, dev = int(num_images), self.device
res = self.cfg["resolutions"][str(int(resolution or self.cfg["resolution"]))]
hw, shift = res["latent_hw"], float(res["shift"])
text, negative = self.prepare(prompt, negative_prompt, normalize, spell)
ids, am = self.text.tokenize([text] * n)
ctx, mask = self.text.encode(ids, am)
counts = PN.count_tensor(self.text, [text] * n, ids.shape[1]) if normalize else torch.zeros(n, ids.shape[1])
u_ctx, u_mask = self.text([negative] * n)
L = _bucket(ctx.shape[1], u_ctx.shape[1])
ctx, mask = _pad_to(ctx, mask, L)
u_ctx, u_mask = _pad_to(u_ctx, u_mask, L)
counts = F.pad(counts.to(dev).float(), (0, L - counts.shape[1]))
both_ctx, both_mask = torch.cat([ctx, u_ctx]), torch.cat([mask, u_mask])
both_counts = torch.cat([counts, torch.zeros_like(counts)]) # the unconditional half gets none
gen = torch.Generator(device=dev).manual_seed(int(seed))
z = torch.randn(n, 4, hw, hw, device=dev, generator=gen)
grid = [shift_t(i / steps, shift) for i in range(steps + 1)]
for i in range(steps):
t = torch.full((n,), grid[i], device=dev)
vc, vu = self.model(torch.cat([z, z]), torch.cat([t, t]), both_ctx, both_mask, both_counts).chunk(2)
z = z + (grid[i + 1] - grid[i]) * (vu + cfg * (vc - vu))
x = self.vae.decode((z / self.vae_div).to(self.vae.dtype)).sample
x = ((x.float().clamp(-1, 1) + 1) * 127.5).round().byte().permute(0, 2, 3, 1).cpu().numpy()
from PIL import Image
imgs = [Image.fromarray(a) for a in x]
if watermark:
imgs = [add_watermark(im) for im in imgs]
if metadata:
for k, im in enumerate(imgs):
im.info.update(provenance(prompt if record_prompt else None, int(seed) if record_prompt else None,
{"image_index": str(k)} if n > 1 else None))
return imgs