pxg-tiny / pxg_tiny /runtime_pipeline.py
Tarul's picture
Upload pxg_tiny/runtime_pipeline.py with huggingface_hub
f21a310 verified
Raw
History Blame Contribute Delete
9.96 kB
"""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)
@staticmethod
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