File size: 9,741 Bytes
4770b6a | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 | """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
|