Download agate/pipeline.py from Logolabs/agate-preview-003: direct link, hf CLI and curl.
- Browser
- Download file 9.74 kB
-
https://huggingface.co/Logolabs/agate-preview-003/resolve/4571b3c532b0656069d1f6f9e8ab950ed6eed577/agate/pipeline.py
- Command line
-
hf download hf://Logolabs/agate-preview-003@4571b3c532b0656069d1f6f9e8ab950ed6eed577/agate/pipeline.py
-
curl -L -o pipeline.py https://huggingface.co/Logolabs/agate-preview-003/resolve/4571b3c532b0656069d1f6f9e8ab950ed6eed577/agate/pipeline.py
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() | |
| 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 | |
| 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 | |