Text-to-Image
English
numpy
machine-learning
deep-learning
generative-ai
text-2-image
image-generation
open-weights
model-weights
ai-art
pixel-art
game-development
gamedev
game-assets
asset-generator
sprite-generator
offline
tiny-model
numpy-runtime
int8-quantization
self-supervised
procedural-data
gpt
english-prompts
awesome-ai
| """PXG-Tiny offline inference runtime — pure NumPy, no torch. | |
| Loads an INT8 weight bundle (gen_int8.npz + vq_int8.npz + runtime.json) and | |
| generates 16x16 sprite grids from English captions with a KV-cached causal | |
| decoder. Weights are stored as per-output-channel INT8 (weight-only | |
| quantization, `key.q`/`key.s` int8/float32 pairs) and dequantized once at | |
| load; all compute is float32. The intent encoder mirrors the torch | |
| TransformerEncoderLayer (norm_first, ReLU FFN); decoder blocks use | |
| pre-LN + tanh-GELU exactly as trained. | |
| """ | |
| import json | |
| from pathlib import Path | |
| import numpy as np | |
| SQ2PI = float(np.sqrt(2.0 / np.pi)) | |
| def layer_norm(x, w, b, eps=1e-5): | |
| m = x.mean(axis=-1, keepdims=True) | |
| v = x.var(axis=-1, keepdims=True) | |
| return (x - m) / np.sqrt(v + eps) * w + b | |
| def tanh_gelu(x): | |
| return 0.5 * x * (1.0 + np.tanh(SQ2PI * (x + 0.044715 * x ** 3))) | |
| def softmax(x, axis=-1): | |
| x = x - x.max(axis=axis, keepdims=True) | |
| e = np.exp(x) | |
| return e / e.sum(axis=axis, keepdims=True) | |
| def dequant(z, name): | |
| """name.q int8 (out,in) + name.s scales -> dequantized (W, b); | |
| falls back to plain fp32 `name.w` when the layer was exported unquantized | |
| (intent encoder stays fp32 for prompt-understanding fidelity).""" | |
| if name + ".q" in z: | |
| w = z[name + ".q"].astype(np.float32) * z[name + ".s"].astype( | |
| np.float32)[:, None] | |
| else: | |
| w = z[name + ".w"].astype(np.float32) | |
| return w, z[name + ".b"].astype(np.float32) | |
| class OfflinePipeline: | |
| def __init__(self, bundle_dir): | |
| bdir = Path(bundle_dir) | |
| cfg = json.loads((bdir / "runtime.json").read_text()) | |
| self.d = cfg["d_model"] | |
| self.L = cfg["n_layers"] | |
| self.H = cfg["n_heads"] | |
| self.dh = self.d // self.H | |
| self.prefix_len = cfg["prefix_len"] | |
| self.seq_vis = cfg["seq_vis"] | |
| self.caps_len = cfg["caps_len"] | |
| self.temperature = cfg["sampling"]["temperature"] | |
| self.top_k = cfg["sampling"]["top_k"] | |
| z = np.load(bdir / "gen_int8.npz") | |
| self.char_emb = z["char_emb"].astype(np.float32) | |
| self.vis_emb = z["vis_emb"].astype(np.float32) | |
| self.pos_emb = z["pos_emb"].astype(np.float32) | |
| self.head_w = z["head.w"].astype(np.float32) | |
| self.head_b = z["head.b"].astype(np.float32) | |
| self.lnf = (z["lnf.w"].astype(np.float32), z["lnf.b"].astype(np.float32)) | |
| # intent encoder (bidirectional, ReLU FFN, norm-first) | |
| self.enc_attn = dequant(z, "enc.attn") | |
| self.enc_attn_out = dequant(z, "enc.attn_out") | |
| self.enc_ff1 = dequant(z, "enc.ff1") | |
| self.enc_ff2 = dequant(z, "enc.ff2") | |
| self.enc_n1 = (z["enc.n1.w"].astype(np.float32), z["enc.n1.b"].astype(np.float32)) | |
| self.enc_n2 = (z["enc.n2.w"].astype(np.float32), z["enc.n2.b"].astype(np.float32)) | |
| self.enc_proj = dequant(z, "enc.proj") | |
| self.enc_ln = (z["enc.ln.w"].astype(np.float32), z["enc.ln.b"].astype(np.float32)) | |
| self.blocks = [] | |
| for i in range(self.L): | |
| self.blocks.append({ | |
| "ln1": (z[f"b{i}.ln1.w"].astype(np.float32), | |
| z[f"b{i}.ln1.b"].astype(np.float32)), | |
| "qkv": dequant(z, f"b{i}.qkv"), | |
| "proj": dequant(z, f"b{i}.proj"), | |
| "ln2": (z[f"b{i}.ln2.w"].astype(np.float32), | |
| z[f"b{i}.ln2.b"].astype(np.float32)), | |
| "fc1": dequant(z, f"b{i}.fc1"), | |
| "fc2": dequant(z, f"b{i}.fc2"), | |
| }) | |
| vq = np.load(bdir / "vq_int8.npz") | |
| self.palette = vq["palette"].astype(np.uint8) # (32,4) RGBA | |
| # ------------------------------------------------------------ encoder -- | |
| def encode_text(self, ids): | |
| """ids (32,) int -> prefix (8, d). Mirrors the torch intent encoder: | |
| char emb -> TransformerEncoderLayer(norm_first, relu) -> chunk-mean -> | |
| Linear -> tanh-GELU -> LN.""" | |
| T = self.caps_len | |
| x = self.char_emb[ids] | |
| qkv = self._lin(layer_norm(x, *self.enc_n1), self.enc_attn) | |
| q, k, v = np.split(qkv, 3, axis=-1) | |
| q = q.reshape(T, self.H, self.dh).transpose(1, 0, 2) | |
| k = k.reshape(T, self.H, self.dh).transpose(1, 0, 2) | |
| v = v.reshape(T, self.H, self.dh).transpose(1, 0, 2) | |
| att = softmax(q @ k.transpose(0, 2, 1) / np.sqrt(self.dh), axis=-1) | |
| y = (att @ v).transpose(1, 0, 2).reshape(T, -1) | |
| x = x + self._lin(y, self.enc_attn_out) | |
| h = layer_norm(x, *self.enc_n2) | |
| x = x + self._lin(np.maximum(self._lin(h, self.enc_ff1), 0.0), self.enc_ff2) | |
| chunks = x.reshape(self.prefix_len, T // self.prefix_len, self.d).mean(axis=1) | |
| return layer_norm(tanh_gelu(self._lin(chunks, self.enc_proj)), *self.enc_ln) | |
| def _lin(x, wb): | |
| w, b = wb | |
| return x @ w.T + b | |
| # ----------------------------------------------------------- decoder -- | |
| def _block_step(self, i, x, Kc, Vc): | |
| ln1_w, ln1_b = self.blocks[i]["ln1"] | |
| h = layer_norm(x, ln1_w, ln1_b) | |
| qkv = self._lin(h, self.blocks[i]["qkv"]) | |
| q, k, v = np.split(qkv, 3, axis=-1) | |
| q = q.reshape(self.H, self.dh) | |
| Kc[i].append(k.reshape(self.H, self.dh)) | |
| Vc[i].append(v.reshape(self.H, self.dh)) | |
| K = np.stack(Kc[i], axis=1) # (H, t, dh) | |
| V = np.stack(Vc[i], axis=1) | |
| att = softmax(np.einsum("hd,htd->ht", q, K) / np.sqrt(self.dh), axis=-1) | |
| y = np.einsum("ht,htd->hd", att, V).reshape(-1) | |
| x = x + self._lin(y, self.blocks[i]["proj"]) | |
| h2 = tanh_gelu(self._lin(layer_norm(x, *self.blocks[i]["ln2"]), | |
| self.blocks[i]["fc1"])) | |
| return x + self._lin(h2, self.blocks[i]["fc2"]) | |
| # ---------------------------------------------------------- sampling -- | |
| def _sample(self, logits, rng, temperature, top_k, logit_bias=None): | |
| z = logits.astype(np.float64) | |
| if logit_bias is not None: | |
| z = z + np.asarray(logit_bias).reshape(-1).astype(np.float64) | |
| z = z / max(temperature, 1e-4) | |
| if top_k and top_k < z.shape[-1]: | |
| thr = np.partition(z, -top_k)[-top_k] | |
| z[z < thr] = -np.inf | |
| p = softmax(z) | |
| return int(rng.choice(len(p), p=p)) | |
| def _bias_at(self, logit_bias, j): | |
| """Flat (32,) bias or per-position (256, 32) matrix row j -> (32,).""" | |
| if logit_bias is None: | |
| return None | |
| arr = np.asarray(logit_bias) | |
| if arr.ndim == 1: | |
| return arr | |
| return arr[min(j, arr.shape[0] - 1)] | |
| # ------------------------------------------------------------ public -- | |
| def generate_grid(self, text, seed=0, temperature=None, top_k=None, | |
| return_logits=False, ids_override=None, logit_bias=None, | |
| prefix_tokens=None): | |
| from pxg_tiny.config import encode_caption | |
| rng = np.random.default_rng(seed) | |
| temperature = self.temperature if temperature is None else temperature | |
| top_k = self.top_k if top_k is None else top_k | |
| if ids_override is not None: | |
| ids = np.asarray(ids_override, dtype=np.int64) | |
| else: | |
| ids = np.array(encode_caption(text), dtype=np.int64) | |
| pref = self.encode_text(ids) | |
| Kc = [[] for _ in range(self.L)] | |
| Vc = [[] for _ in range(self.L)] | |
| x = None | |
| for pi in range(self.prefix_len): | |
| x = pref[pi] + self.pos_emb[pi] | |
| for i in range(self.L): | |
| x = self._block_step(i, x, Kc, Vc) | |
| forced = (list(prefix_tokens) if prefix_tokens is not None else []) | |
| tokens = [] | |
| teacher_logits = [] | |
| for j in range(self.seq_vis): | |
| if j < len(forced): | |
| t = int(forced[j]) # structural prior token | |
| teacher_logits.append(None) | |
| else: | |
| logits = self._lin(layer_norm(x, *self.lnf), | |
| (self.head_w, self.head_b)) | |
| t = self._sample(logits, rng, temperature, top_k, | |
| logit_bias=self._bias_at(logit_bias, j)) | |
| teacher_logits.append(logits) | |
| tokens.append(t) | |
| if j < self.seq_vis - 1: | |
| x = self.vis_emb[t] + self.pos_emb[self.prefix_len + j + 1] | |
| for i in range(self.L): | |
| x = self._block_step(i, x, Kc, Vc) | |
| grid = np.array(tokens, dtype=np.uint8).reshape(16, 16) | |
| if return_logits: | |
| return grid, np.stack(teacher_logits) | |
| return grid | |
| def teacher_forced_logits(self, caps_ids, vis_tokens): | |
| """Parity hook: feed the exact training layout and return the 256x32 | |
| logits the model assigns for each visual position (prefix only fed | |
| once, then ground-truth tokens streamed).""" | |
| Kc = [[] for _ in range(self.L)] | |
| Vc = [[] for _ in range(self.L)] | |
| pref = self.encode_text(np.array(caps_ids, dtype=np.int64)) | |
| outs = [] | |
| x = None | |
| for pi in range(self.prefix_len): | |
| x = pref[pi] + self.pos_emb[pi] | |
| for i in range(self.L): | |
| x = self._block_step(i, x, Kc, Vc) | |
| outs.append(self._lin(layer_norm(x, *self.lnf), (self.head_w, self.head_b))) | |
| for j in range(self.seq_vis - 1): | |
| x = self.vis_emb[int(vis_tokens[j])] + self.pos_emb[self.prefix_len + j] | |
| for i in range(self.L): | |
| x = self._block_step(i, x, Kc, Vc) | |
| outs.append(self._lin(layer_norm(x, *self.lnf), (self.head_w, self.head_b))) | |
| return np.stack(outs) | |
| def generate_rgba(self, text, seed=0, **kw): | |
| grid = self.generate_grid(text, seed=seed, **kw) | |
| rgba = np.zeros((16, 16, 4), dtype=np.uint8) | |
| for idx in range(1, len(self.palette)): | |
| m = grid == idx | |
| rgba[m] = self.palette[idx] | |
| return grid, rgba | |