ofou's picture
download
raw
5.03 kB
#!/usr/bin/env python3
"""MLP Net for 1D RandOpt / Neural Thickets in tinygrad."""
from __future__ import annotations
import numpy as np
from tinygrad import Tensor, nn
class Net:
"""MLP: input -> [Linear+ReLU]*(depth-1) -> Linear -> scalar next-step pred."""
def __init__(self, width: int, depth: int, dim_in: int, dim_out: int = 1, init_type: str = "xavier"):
self.width = width
self.depth = depth
self.dim_in = dim_in
self.dim_out = dim_out
self.init_type = init_type
self.linears: list[nn.Linear] = []
# Match official: Linear(in,w) then (depth-2) of ReLU+Linear(w,w) then ReLU+Linear(w,out)
self.linears.append(nn.Linear(dim_in, width))
for _ in range(depth - 2):
self.linears.append(nn.Linear(width, width))
self.linears.append(nn.Linear(width, dim_out))
def parameters(self) -> list[Tensor]:
return nn.state.get_parameters(self)
def n_params(self) -> int:
return int(sum(int(np.prod(p.shape)) for p in self.parameters()))
def __call__(self, ctx: Tensor) -> Tensor:
was_1d = len(ctx.shape) == 1
if was_1d:
ctx = ctx.unsqueeze(0)
h = ctx
for i, layer in enumerate(self.linears):
h = layer(h)
if i < len(self.linears) - 1:
h = h.relu()
if was_1d:
h = h.squeeze(0)
# squeeze last dim if dim_out==1 -> [B] or scalar
if self.dim_out == 1 and len(h.shape) >= 1 and h.shape[-1] == 1:
h = h.squeeze(-1)
return h
def compute_loss(self, ctx: Tensor, y: Tensor) -> Tensor:
y_pred = self(ctx)
target = y.squeeze(-1) if len(y.shape) > 1 and y.shape[-1] == 1 else y
return ((y_pred - target) ** 2).mean()
def init_weights(self) -> None:
for layer in self.linears:
fan_in = int(layer.weight.shape[1])
fan_out = int(layer.weight.shape[0])
if self.init_type == "xavier":
a = float(np.sqrt(6.0 / (fan_in + fan_out)))
w = (np.random.uniform(-a, a, size=tuple(layer.weight.shape))).astype(np.float32)
layer.weight.assign(Tensor(w))
elif self.init_type == "kaiming":
layer.weight.assign(Tensor.kaiming_uniform(*layer.weight.shape))
else:
raise ValueError(f"Invalid init_type: {self.init_type}")
if layer.bias is not None:
layer.bias.assign(Tensor.zeros(*layer.bias.shape))
layer.weight.realize()
if layer.bias is not None:
layer.bias.realize()
def perturb_weights(self, seed: int, sigma: float) -> None:
"""In-place Gaussian perturbation with numpy RNG for reproducibility."""
rng = np.random.RandomState(seed)
for p in self.parameters():
noise = rng.randn(*p.shape).astype(np.float32) * sigma
p.assign(p + Tensor(noise))
p.realize()
def snapshot_weights(self) -> list[np.ndarray]:
return [p.numpy().copy() for p in self.parameters()]
def load_weights(self, arrays: list[np.ndarray]) -> None:
for p, arr in zip(self.parameters(), arrays):
p.assign(Tensor(arr.astype(np.float32)))
p.realize()
def perturb_from_snapshot(self, snapshot: list[np.ndarray], seed: int, sigma: float) -> None:
"""Reset to snapshot then apply Gaussian noise (avoids expensive Net reconstruction)."""
rng = np.random.RandomState(seed)
for p, base in zip(self.parameters(), snapshot):
noise = rng.randn(*base.shape).astype(np.float32) * sigma
p.assign(Tensor((base + noise).astype(np.float32)))
p.realize()
def clone(self) -> "Net":
"""Deep-copy parameters into a new Net."""
other = Net(self.width, self.depth, self.dim_in, self.dim_out, self.init_type)
other.load_weights(self.snapshot_weights())
return other
def AR_rollout(self, ctx: Tensor, T: int) -> Tensor:
"""Autoregressive rollout for T steps. ctx: [B, ctx_sz] -> [B, T]."""
cur = ctx.realize()
outs = []
for _ in range(T):
y_pred = self(cur).realize()
outs.append(y_pred.numpy())
cur = cur.cat(y_pred.unsqueeze(-1), dim=1)[:, 1:].realize()
return Tensor(np.stack(outs, axis=1).astype(np.float32))
def compute_mse(y_pred: Tensor, y_true: Tensor) -> float:
err = ((y_pred - y_true) ** 2).numpy().reshape(-1)
return float(err.mean())
def eval_model(model: Net, ctx_y: Tensor, fut_y: Tensor, fut_sz: int) -> float:
y_pred = model.AR_rollout(ctx_y, fut_sz)
return compute_mse(y_pred, fut_y)
def eval_next_step(model: Net, ctx_y: Tensor, fut_y: Tensor) -> float:
"""Fast single-step MSE (matches pretrain objective). Used for density sweeps."""
pred = model(ctx_y).realize()
target = fut_y[:, 0].realize()
return float(((pred - target) ** 2).mean().numpy())

Xet Storage Details

Size:
5.03 kB
·
Xet hash:
0770f916ea1eb0ebef2b88c5ebed2fb24ee63c7064209062cc5f1f120f3a939f

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.