File size: 8,037 Bytes
e465a2f | 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 198 199 200 201 202 203 204 205 206 207 208 209 210 211 | """
PixelModel v2 - the weights ARE the image.
model.png stores every parameter of the network as pixel values.
v2 changes vs v1:
- Wider trunk and decoder: 200,259 parameters rather than 23,747, while
keeping the same prompt encoder, coordinate features, 64x64 native
resolution, and 16-bit PNG codec.
- The larger capacity is intended for the same ~20K COCO training rows.
v1 retained these changes from v0:
- Coordinate-conditioned decoder (CPPN-style): parameters no longer scale
with output resolution. v0 spent 196,608 of its 202,752 params on the
output head; v1 has 23,747 params total and renders at ANY resolution
(native 64x64 for eval).
- Hashed character-trigram + word prompt embedding (deterministic, 0 params)
instead of v0's char-sum, so different prompts actually get different,
partially compositional embeddings.
- Biases everywhere (v0 had none).
- 16-bit weight codec: R = high byte, G = low byte, B = reserved.
v0 wasted G and stored weights at 8 bits; v1's quantization error is
~6e-5 per weight. model.png is 160x149 px - the model is a thumbnail.
Architecture:
prompt -> hashed trigram/word embedding (64)
-> T1 (256x64)+b tanh -> T2 (128x256)+b tanh -> latent z (128)
per pixel: concat(z, fourier(x,y) (18))
-> D1 (320x146)+b tanh -> D2 (320x320)+b tanh -> D3 (3x320)+b sigmoid -> RGB
"""
import numpy as np
import torch
from PIL import Image
# -- config -------------------------------------------------------------------
EMB_DIM = 64 # prompt embedding size (hashed, no learned params)
TRUNK_HIDDEN = 256
LATENT = 128
COORD_FEATS = 18 # x, y + sin/cos at freqs (1,2,4,8) for each axis
DEC_HIDDEN = 320
NATIVE_RES = 64 # native eval resolution (decoder works at any res)
FREQS = (1.0, 2.0, 4.0, 8.0)
# every parameter tensor, in canonical flattening order
PARAM_SPECS = [
("T1", (TRUNK_HIDDEN, EMB_DIM)),
("b1", (TRUNK_HIDDEN,)),
("T2", (LATENT, TRUNK_HIDDEN)),
("b2", (LATENT,)),
("D1", (DEC_HIDDEN, LATENT + COORD_FEATS)),
("bd1", (DEC_HIDDEN,)),
("D2", (DEC_HIDDEN, DEC_HIDDEN)),
("bd2", (DEC_HIDDEN,)),
("D3", (3, DEC_HIDDEN)),
("bd3", (3,)),
]
N_PARAMS = sum(int(np.prod(s)) for _, s in PARAM_SPECS) # 200,259
# model.png geometry: one weight per pixel, 16-bit (R=high, G=low, B=reserved)
MODEL_W = 256
MODEL_H = -(-N_PARAMS // MODEL_W) # 783
WMAX = 2.0 # weights live in [-2, 2]
# -- prompt embedding (deterministic, zero parameters) ------------------------
def _fnv1a(data: bytes) -> int:
h = 0x811C9DC5
for byte in data:
h ^= byte
h = (h * 0x01000193) & 0xFFFFFFFF
return h
def prompt_to_embedding(prompt: str) -> torch.Tensor:
"""Hashed bag of character trigrams + words -> EMB_DIM vector, L2-normed."""
text = "".join(c if c.isalnum() or c == " " else " " for c in prompt.lower())
text = " ".join(text.split())
vec = np.zeros(EMB_DIM, dtype=np.float32)
padded = f" {text} "
for i in range(len(padded) - 2):
h = _fnv1a(padded[i:i + 3].encode("utf-8"))
sign = 1.0 if (h >> 16) & 1 else -1.0
vec[h % EMB_DIM] += sign
for word in text.split():
h = _fnv1a(b"w:" + word.encode("utf-8"))
sign = 1.0 if (h >> 16) & 1 else -1.0
vec[h % EMB_DIM] += 2.0 * sign # words weigh more than trigrams
norm = np.linalg.norm(vec)
if norm > 0:
vec /= norm
return torch.from_numpy(vec)
def prompts_to_embeddings(prompts) -> torch.Tensor:
return torch.stack([prompt_to_embedding(p) for p in prompts])
# -- coordinate features ------------------------------------------------------
_coord_cache = {}
def coord_features(res: int) -> torch.Tensor:
"""(res*res, COORD_FEATS) fourier features of the pixel grid, cached."""
if res not in _coord_cache:
axis = torch.linspace(-1.0, 1.0, res)
yy, xx = torch.meshgrid(axis, axis, indexing="ij")
x = xx.reshape(-1)
y = yy.reshape(-1)
feats = [x, y]
for f in FREQS:
feats += [torch.sin(f * torch.pi * x), torch.cos(f * torch.pi * x),
torch.sin(f * torch.pi * y), torch.cos(f * torch.pi * y)]
_coord_cache[res] = torch.stack(feats, dim=1) # (res*res, 18)
return _coord_cache[res]
# -- forward ------------------------------------------------------------------
def encode_prompt(weights: dict, emb: torch.Tensor) -> torch.Tensor:
"""emb (B, EMB_DIM) -> latent z (B, LATENT)."""
z = torch.tanh(emb @ weights["T1"].T + weights["b1"])
z = torch.tanh(z @ weights["T2"].T + weights["b2"])
return z
def decode_pixels(weights: dict, z: torch.Tensor, feats: torch.Tensor) -> torch.Tensor:
"""z (B, LATENT), feats (P, COORD_FEATS) -> RGB (B, P, 3) in [0, 1]."""
B, P = z.shape[0], feats.shape[0]
inp = torch.cat([z.unsqueeze(1).expand(B, P, LATENT),
feats.unsqueeze(0).expand(B, P, COORD_FEATS)], dim=2)
h = torch.tanh(inp @ weights["D1"].T + weights["bd1"])
h = torch.tanh(h @ weights["D2"].T + weights["bd2"])
return torch.sigmoid(h @ weights["D3"].T + weights["bd3"])
def forward(weights: dict, prompts, res: int = NATIVE_RES) -> torch.Tensor:
"""prompts: str or list of str -> (B, res, res, 3) float in [0, 1]."""
if isinstance(prompts, str):
prompts = [prompts]
emb = prompts_to_embeddings(prompts)
z = encode_prompt(weights, emb)
rgb = decode_pixels(weights, z, coord_features(res))
return rgb.reshape(len(prompts), res, res, 3)
# -- 16-bit PNG weight codec --------------------------------------------------
def weights_to_pixels(weights: dict) -> np.ndarray:
"""weight dict -> (MODEL_H, MODEL_W, 3) uint8. R=high byte, G=low, B=0."""
flat = torch.cat([weights[n].detach().reshape(-1) for n, _ in PARAM_SPECS])
q = ((flat.clamp(-WMAX, WMAX) / WMAX + 1.0) / 2.0 * 65535.0).round()
q = q.to(torch.int64).cpu().numpy()
q = np.pad(q, (0, MODEL_W * MODEL_H - N_PARAMS))
img = np.zeros((MODEL_H * MODEL_W, 3), dtype=np.uint8)
img[:, 0] = q >> 8
img[:, 1] = q & 0xFF
return img.reshape(MODEL_H, MODEL_W, 3)
def pixels_to_weights(arr: np.ndarray) -> dict:
"""(MODEL_H, MODEL_W, 3) uint8 -> weight dict."""
flat = arr.reshape(-1, 3).astype(np.int64)
q = (flat[:, 0] << 8) | flat[:, 1]
vals = (q.astype(np.float32) / 65535.0 * 2.0 - 1.0) * WMAX
vals = torch.from_numpy(vals[:N_PARAMS])
weights, i = {}, 0
for name, shape in PARAM_SPECS:
n = int(np.prod(shape))
weights[name] = vals[i:i + n].reshape(shape).clone()
i += n
return weights
def save_model(weights: dict, path: str):
Image.fromarray(weights_to_pixels(weights), mode="RGB").save(path)
def load_model(path: str) -> dict:
img = Image.open(path).convert("RGB")
arr = np.array(img, dtype=np.uint8)
assert arr.shape == (MODEL_H, MODEL_W, 3), \
f"expected {MODEL_H}x{MODEL_W} model.png, got {arr.shape[0]}x{arr.shape[1]}"
return pixels_to_weights(arr)
def init_weights(seed: int = 0) -> dict:
g = torch.Generator().manual_seed(seed)
weights = {}
for name, shape in PARAM_SPECS:
if len(shape) == 1:
weights[name] = torch.zeros(shape)
else:
fan_in = shape[1]
scale = 1.0 / fan_in ** 0.5
if name == "D3":
scale *= 0.1 # start near grey, not saturated
weights[name] = torch.randn(shape, generator=g) * scale
return weights
if __name__ == "__main__":
print(f"params: {N_PARAMS} model.png: {MODEL_W}x{MODEL_H} px")
for name, shape in PARAM_SPECS:
print(f" {name:>4} {str(tuple(shape)):>12} {int(np.prod(shape)):>6}")
|