LH-Tech-AI commited on
Commit
3039135
·
verified ·
1 Parent(s): 59aacb2

Create inference.py

Browse files
Files changed (1) hide show
  1. inference.py +411 -0
inference.py ADDED
@@ -0,0 +1,411 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ #!/usr/bin/env python3
2
+ """
3
+ Supra2-IMG inference — standalone text-to-image (DiT + Flan-T5-Base + SD VAE).
4
+
5
+ Works on Linux and Windows. No imports from other project files.
6
+
7
+ Usage:
8
+ python inference.py --prompt "a sea jellyfish floating in the pitch-black ocean depths" \\
9
+ --seed 0 --cfg 3.0 --steps 50 --n 1 --out jellyfish.png
10
+
11
+ If ./model_final_ema.pt is missing, it is downloaded from Hugging Face:
12
+ SupraLabs/Supra2-IMG
13
+ """
14
+
15
+ from __future__ import annotations
16
+
17
+ import argparse
18
+ import math
19
+ import os
20
+ import sys
21
+ import time
22
+
23
+ import torch
24
+ import torch.nn as nn
25
+ import torch.nn.functional as F
26
+
27
+ # ---------------------------------------------------------------------------
28
+ # Architecture constants (must match the trained checkpoint)
29
+ # ---------------------------------------------------------------------------
30
+ IMG_SIZE = 256
31
+ LATENT_SIZE = 32 # 256 / 8 (f8 VAE)
32
+ LATENT_CH = 4
33
+ PATCH = 2
34
+ NUM_TOKENS = (LATENT_SIZE // PATCH) ** 2
35
+
36
+ D_MODEL = 576
37
+ DEPTH = 14
38
+ N_HEADS = 9
39
+ MLP_RATIO = 4.0
40
+ D_CTX = 768 # Flan-T5-Base
41
+ MAX_CTX_LEN = 128
42
+ T5_NAME = "google/flan-t5-base"
43
+ VAE_NAME = "stabilityai/sd-vae-ft-mse"
44
+ VAE_SCALE = 0.18215
45
+
46
+ HF_REPO = "SupraLabs/Supra2-IMG"
47
+ DEFAULT_CKPT = os.path.join(".", "model_final_ema.pt")
48
+
49
+
50
+ def log(msg: str) -> None:
51
+ """Print immediately so the user sees live progress."""
52
+ print(msg, flush=True)
53
+
54
+
55
+ def pick_device() -> torch.device:
56
+ if torch.cuda.is_available():
57
+ torch.cuda.set_device(0)
58
+ name = torch.cuda.get_device_name(0)
59
+ mem = torch.cuda.get_device_properties(0).total_memory / 1e9
60
+ log(f"[device] CUDA: {name} ({mem:.1f} GB)")
61
+ return torch.device("cuda:0")
62
+ if hasattr(torch.backends, "mps") and torch.backends.mps.is_available():
63
+ log("[device] Apple MPS")
64
+ return torch.device("mps")
65
+ log("[device] CPU (this will be slow)")
66
+ return torch.device("cpu")
67
+
68
+
69
+ # ===========================================================================
70
+ # MODEL
71
+ # ===========================================================================
72
+ def modulate(x: torch.Tensor, shift: torch.Tensor, scale: torch.Tensor) -> torch.Tensor:
73
+ """AdaLN: x * (1 + scale) + shift."""
74
+ return x * (1 + scale.unsqueeze(1)) + shift.unsqueeze(1)
75
+
76
+
77
+ class TimestepEmbedder(nn.Module):
78
+ """Sinusoidal timestep embedding followed by an MLP."""
79
+
80
+ def __init__(self, hidden_size: int, freq_dim: int = 256) -> None:
81
+ super().__init__()
82
+ self.freq_dim = freq_dim
83
+ self.mlp = nn.Sequential(
84
+ nn.Linear(freq_dim, hidden_size),
85
+ nn.SiLU(),
86
+ nn.Linear(hidden_size, hidden_size),
87
+ )
88
+
89
+ def _sinusoidal(self, t: torch.Tensor) -> torch.Tensor:
90
+ half = self.freq_dim // 2
91
+ freqs = torch.exp(
92
+ -math.log(10000.0) * torch.arange(half, device=t.device) / half
93
+ )
94
+ args = t[:, None].float() * freqs[None] * 1000.0
95
+ emb = torch.cat([torch.cos(args), torch.sin(args)], dim=-1)
96
+ if self.freq_dim % 2:
97
+ emb = F.pad(emb, (0, 1))
98
+ return emb
99
+
100
+ def forward(self, t: torch.Tensor) -> torch.Tensor:
101
+ return self.mlp(self._sinusoidal(t))
102
+
103
+
104
+ class Attention(nn.Module):
105
+ """Multi-head self- or cross-attention."""
106
+
107
+ def __init__(self, dim: int, n_heads: int, ctx_dim: int | None = None) -> None:
108
+ super().__init__()
109
+ self.n_heads = n_heads
110
+ self.head_dim = dim // n_heads
111
+ self.is_self = ctx_dim is None
112
+ if self.is_self:
113
+ self.qkv = nn.Linear(dim, dim * 3, bias=True)
114
+ else:
115
+ self.q = nn.Linear(dim, dim, bias=True)
116
+ self.kv = nn.Linear(ctx_dim, dim * 2, bias=True)
117
+ self.proj = nn.Linear(dim, dim, bias=True)
118
+
119
+ def forward(
120
+ self,
121
+ x: torch.Tensor,
122
+ ctx: torch.Tensor | None = None,
123
+ ctx_mask: torch.Tensor | None = None,
124
+ ) -> torch.Tensor:
125
+ B, N, C = x.shape
126
+ if self.is_self:
127
+ qkv = self.qkv(x).view(B, N, 3, self.n_heads, self.head_dim)
128
+ q, k, v = (qkv[:, :, i].transpose(1, 2) for i in range(3))
129
+ else:
130
+ M = ctx.shape[1]
131
+ q = self.q(x).view(B, N, self.n_heads, self.head_dim).transpose(1, 2)
132
+ kv = self.kv(ctx).view(B, M, 2, self.n_heads, self.head_dim)
133
+ k, v = kv[:, :, 0].transpose(1, 2), kv[:, :, 1].transpose(1, 2)
134
+
135
+ attn_mask = None
136
+ if ctx_mask is not None:
137
+ attn_mask = ctx_mask.bool()[:, None, None, :]
138
+
139
+ out = F.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
140
+ out = out.transpose(1, 2).reshape(B, N, C)
141
+ return self.proj(out)
142
+
143
+
144
+ class DiTBlock(nn.Module):
145
+ """DiT block: AdaLN-Zero self-attn + cross-attn + MLP."""
146
+
147
+ def __init__(self, dim: int, n_heads: int, ctx_dim: int, mlp_ratio: float) -> None:
148
+ super().__init__()
149
+ self.norm1 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
150
+ self.self_attn = Attention(dim, n_heads)
151
+ self.norm_ca = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
152
+ self.cross_attn = Attention(dim, n_heads, ctx_dim=dim)
153
+ self.norm2 = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
154
+ hidden = int(dim * mlp_ratio)
155
+ self.mlp = nn.Sequential(
156
+ nn.Linear(dim, hidden),
157
+ nn.GELU(approximate="tanh"),
158
+ nn.Linear(hidden, dim),
159
+ )
160
+ self.adaln = nn.Sequential(nn.SiLU(), nn.Linear(dim, 6 * dim, bias=True))
161
+
162
+ def forward(
163
+ self,
164
+ x: torch.Tensor,
165
+ c: torch.Tensor,
166
+ ctx: torch.Tensor,
167
+ ctx_mask: torch.Tensor | None,
168
+ ) -> torch.Tensor:
169
+ shift_sa, scale_sa, gate_sa, shift_mlp, scale_mlp, gate_mlp = self.adaln(c).chunk(6, dim=1)
170
+ x = x + gate_sa.unsqueeze(1) * self.self_attn(
171
+ modulate(self.norm1(x), shift_sa, scale_sa)
172
+ )
173
+ x = x + self.cross_attn(self.norm_ca(x), ctx=ctx, ctx_mask=ctx_mask)
174
+ x = x + gate_mlp.unsqueeze(1) * self.mlp(modulate(self.norm2(x), shift_mlp, scale_mlp))
175
+ return x
176
+
177
+
178
+ class FinalLayer(nn.Module):
179
+ def __init__(self, dim: int, out_ch: int) -> None:
180
+ super().__init__()
181
+ self.norm = nn.LayerNorm(dim, elementwise_affine=False, eps=1e-6)
182
+ self.linear = nn.Linear(dim, out_ch, bias=True)
183
+ self.adaln = nn.Sequential(nn.SiLU(), nn.Linear(dim, 2 * dim, bias=True))
184
+
185
+ def forward(self, x: torch.Tensor, c: torch.Tensor) -> torch.Tensor:
186
+ shift, scale = self.adaln(c).chunk(2, dim=1)
187
+ return self.linear(modulate(self.norm(x), shift, scale))
188
+
189
+
190
+ class SupraDiT(nn.Module):
191
+ """~100M parameter DiT for rectified-flow text-to-image."""
192
+
193
+ def __init__(
194
+ self,
195
+ latent_ch: int = LATENT_CH,
196
+ d_model: int = D_MODEL,
197
+ depth: int = DEPTH,
198
+ n_heads: int = N_HEADS,
199
+ ctx_dim: int = D_CTX,
200
+ mlp_ratio: float = MLP_RATIO,
201
+ num_tokens: int = NUM_TOKENS,
202
+ ) -> None:
203
+ super().__init__()
204
+ self.num_tokens = num_tokens
205
+ self.patch = PATCH
206
+ self.x_embed = nn.Linear(latent_ch * PATCH * PATCH, d_model)
207
+ self.pos_embed = nn.Parameter(torch.zeros(1, num_tokens, d_model))
208
+ self.t_embed = TimestepEmbedder(d_model)
209
+ self.ctx_proj = nn.Linear(ctx_dim, d_model)
210
+ self.blocks = nn.ModuleList(
211
+ [DiTBlock(d_model, n_heads, d_model, mlp_ratio) for _ in range(depth)]
212
+ )
213
+ self.final = FinalLayer(d_model, latent_ch * PATCH * PATCH)
214
+
215
+ def forward(
216
+ self,
217
+ z: torch.Tensor,
218
+ t: torch.Tensor,
219
+ ctx: torch.Tensor,
220
+ ctx_mask: torch.Tensor | None = None,
221
+ ) -> torch.Tensor:
222
+ B, C, H, W = z.shape
223
+ P = self.patch
224
+ h, w = H // P, W // P
225
+ x = z.view(B, C, h, P, w, P).permute(0, 2, 4, 1, 3, 5).reshape(B, h * w, C * P * P)
226
+ x = self.x_embed(x) + self.pos_embed
227
+ c = self.t_embed(t)
228
+ ctx = self.ctx_proj(ctx)
229
+ for blk in self.blocks:
230
+ x = blk(x, c, ctx, ctx_mask)
231
+ x = self.final(x, c)
232
+ return x.view(B, h, w, C, P, P).permute(0, 3, 1, 4, 2, 5).reshape(B, C, H, W)
233
+
234
+
235
+ # ===========================================================================
236
+ # Checkpoint download
237
+ # ===========================================================================
238
+ def ensure_checkpoint(path: str) -> str:
239
+ """Use local checkpoint or download model_final_ema.pt from Hugging Face."""
240
+ if os.path.isfile(path):
241
+ log(f"[ckpt] found {path}")
242
+ return path
243
+
244
+ log(f"[ckpt] {path} not found — downloading from {HF_REPO} ...")
245
+ try:
246
+ from huggingface_hub import hf_hub_download
247
+ except ImportError:
248
+ log("[ckpt] installing huggingface_hub ...")
249
+ import subprocess
250
+
251
+ subprocess.check_call([sys.executable, "-m", "pip", "install", "-q", "huggingface_hub"])
252
+ from huggingface_hub import hf_hub_download
253
+
254
+ filename = os.path.basename(path) or "model_final_ema.pt"
255
+ downloaded = hf_hub_download(
256
+ repo_id=HF_REPO,
257
+ filename=filename,
258
+ local_dir=".",
259
+ local_dir_use_symlinks=False,
260
+ )
261
+ # Prefer the expected local path if the hub placed it elsewhere
262
+ if os.path.isfile(filename) and os.path.abspath(filename) != os.path.abspath(path):
263
+ import shutil
264
+
265
+ shutil.copy2(filename, path)
266
+ log(f"[ckpt] copied to {path}")
267
+ return path
268
+ log(f"[ckpt] downloaded: {downloaded}")
269
+ return downloaded if os.path.isfile(downloaded) else path
270
+
271
+
272
+ # ===========================================================================
273
+ # Sampling (Euler integration of the flow ODE + optional CFG)
274
+ # ===========================================================================
275
+ @torch.no_grad()
276
+ def generate(args: argparse.Namespace, device: torch.device) -> None:
277
+ from transformers import AutoTokenizer, T5EncoderModel
278
+ from diffusers import AutoencoderKL
279
+ import torchvision.utils as vutils
280
+
281
+ ckpt_path = ensure_checkpoint(DEFAULT_CKPT)
282
+
283
+ log("[model] building SupraDiT ...")
284
+ model = SupraDiT().to(device).eval()
285
+ n_params = sum(p.numel() for p in model.parameters())
286
+ log(f"[model] {n_params / 1e6:.1f}M parameters")
287
+
288
+ log(f"[model] loading weights from {ckpt_path} ...")
289
+ t0 = time.perf_counter()
290
+ state = torch.load(ckpt_path, map_location=device, weights_only=False)
291
+ cfg = state.get("config", {}) if isinstance(state, dict) else {}
292
+ if isinstance(state, dict):
293
+ if cfg.get("patch", PATCH) != PATCH:
294
+ raise SystemExit(f"Checkpoint PATCH={cfg['patch']} != script PATCH={PATCH}")
295
+ weights = state["ema"] if "ema" in state else state.get("model", state)
296
+ else:
297
+ weights = state
298
+ model.load_state_dict(weights, strict=True)
299
+ log(f"[model] weights loaded in {time.perf_counter() - t0:.1f}s")
300
+
301
+ ctx_len = int(cfg.get("ctx_len", MAX_CTX_LEN))
302
+ log(f"[text] ctx_len={ctx_len}")
303
+
304
+ log(f"[text] loading tokenizer + {T5_NAME} ...")
305
+ tokenizer = AutoTokenizer.from_pretrained(T5_NAME)
306
+ text_model = T5EncoderModel.from_pretrained(T5_NAME).to(device).eval()
307
+ for p in text_model.parameters():
308
+ p.requires_grad = False
309
+
310
+ log(f"[vae] loading {VAE_NAME} ...")
311
+ vae = AutoencoderKL.from_pretrained(VAE_NAME).to(device).eval()
312
+
313
+ torch.manual_seed(args.seed)
314
+ if device.type == "cuda":
315
+ torch.cuda.manual_seed_all(args.seed)
316
+
317
+ prompts = [args.prompt] * args.n
318
+ n_tok = len(tokenizer(args.prompt)["input_ids"])
319
+ log(f"[text] prompt tokens={n_tok} n={args.n} seed={args.seed} cfg={args.cfg} steps={args.steps}")
320
+ if n_tok > ctx_len:
321
+ log(f"[text] WARNING: prompt has {n_tok} tokens, truncated to ctx_len={ctx_len}")
322
+
323
+ tok = tokenizer(
324
+ prompts,
325
+ padding="max_length",
326
+ truncation=True,
327
+ max_length=ctx_len,
328
+ return_tensors="pt",
329
+ ).to(device)
330
+
331
+ use_amp = device.type == "cuda"
332
+ with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_amp):
333
+ ctx = text_model(**tok).last_hidden_state.float()
334
+ cmask = tok["attention_mask"].float()
335
+
336
+ use_cfg = args.cfg > 1.0
337
+ if use_cfg:
338
+ if isinstance(cfg, dict) and "uncond_text" in cfg:
339
+ uncond_ctx = cfg["uncond_text"].to(device).float().unsqueeze(0).expand(args.n, -1, -1)
340
+ uncond_mask = cfg["uncond_mask"].to(device).float().unsqueeze(0).expand(args.n, -1)
341
+ log("[cfg] using stored unconditional embeddings")
342
+ else:
343
+ u_tok = tokenizer(
344
+ [""] * args.n,
345
+ padding="max_length",
346
+ truncation=True,
347
+ max_length=ctx_len,
348
+ return_tensors="pt",
349
+ ).to(device)
350
+ with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_amp):
351
+ uncond_ctx = text_model(**u_tok).last_hidden_state.float()
352
+ uncond_mask = u_tok["attention_mask"].float()
353
+ log("[cfg] encoded empty string as unconditional")
354
+ ctx_all = torch.cat([ctx, uncond_ctx], 0)
355
+ mask_all = torch.cat([cmask, uncond_mask], 0)
356
+
357
+ z = torch.randn(args.n, LATENT_CH, LATENT_SIZE, LATENT_SIZE, device=device)
358
+ dt = 1.0 / args.steps
359
+ log(f"[sample] Euler flow, {args.steps} steps ...")
360
+ t_sample = time.perf_counter()
361
+
362
+ for i in range(args.steps):
363
+ t = torch.full((args.n,), i * dt, device=device)
364
+ with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_amp):
365
+ if use_cfg:
366
+ v_both = model(torch.cat([z, z], 0), torch.cat([t, t], 0), ctx_all, mask_all)
367
+ v_cond, v_uncond = v_both.float().chunk(2, 0)
368
+ v = v_uncond + args.cfg * (v_cond - v_uncond)
369
+ else:
370
+ v = model(z, t, ctx, cmask).float()
371
+ z = z + dt * v
372
+ if (i + 1) % max(1, args.steps // 10) == 0 or i == 0:
373
+ log(f" step {i + 1}/{args.steps}")
374
+
375
+ log(f"[sample] denoising done in {time.perf_counter() - t_sample:.1f}s")
376
+ log("[vae] decoding latents ...")
377
+ with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_amp):
378
+ imgs = vae.decode(z / VAE_SCALE).sample
379
+ imgs = (imgs.clamp(-1, 1) + 1) / 2
380
+
381
+ out_dir = os.path.dirname(os.path.abspath(args.out))
382
+ if out_dir:
383
+ os.makedirs(out_dir, exist_ok=True)
384
+ vutils.save_image(imgs, args.out, nrow=int(math.ceil(math.sqrt(args.n))))
385
+ log(f"[done] saved {args.n} image(s) -> {args.out}")
386
+
387
+
388
+ def parse_args() -> argparse.Namespace:
389
+ p = argparse.ArgumentParser(description="Supra2-IMG standalone inference")
390
+ p.add_argument(
391
+ "--prompt",
392
+ default="a sea jellyfish floating in the pitch-black ocean depths",
393
+ help="Text prompt",
394
+ )
395
+ p.add_argument("--seed", type=int, default=0, help="RNG seed")
396
+ p.add_argument("--cfg", type=float, default=3.0, help="Classifier-free guidance scale")
397
+ p.add_argument("--steps", type=int, default=50, help="Euler ODE steps")
398
+ p.add_argument("--n", type=int, default=1, help="Number of images")
399
+ p.add_argument("--out", default="jellyfish.png", help="Output image path")
400
+ return p.parse_args()
401
+
402
+
403
+ def main() -> None:
404
+ args = parse_args()
405
+ log("=== Supra2-IMG inference ===")
406
+ device = pick_device()
407
+ generate(args, device)
408
+
409
+
410
+ if __name__ == "__main__":
411
+ main()