ofou's picture
download
raw
5.28 kB
#!/usr/bin/env python3
"""Pretrain / RandOpt / ES / SGD helpers in tinygrad."""
from __future__ import annotations
import time
from dataclasses import dataclass
from typing import Any
import numpy as np
from tinygrad import Tensor, nn
from datasets import load_data
from models import Net, eval_model, eval_next_step
@dataclass
class Args:
ctx_sz: int = 10
fut_sz: int = 60
res_x: float = 0.1
width: int = 128
depth: int = 5
base_init: str = "xavier"
pretrain_bsz: int = 256
pretrain_iters: int = 800
pretraining_lr: float = 0.001
posttrain_dataset_sz: int = 512
sigma: float = 0.002
N: int = 400
K: int = 20
global_seed: int = 0
def set_seed(seed: int) -> None:
import random
import numpy as np
random.seed(seed)
np.random.seed(seed)
Tensor.manual_seed(seed)
def create_model(args: Args) -> Net:
model = Net(
width=args.width,
depth=args.depth,
dim_in=args.ctx_sz,
dim_out=1,
init_type=args.base_init,
)
model.init_weights()
return model
def pretrain_base_model(model: Net, pretrain_dataset: str, args: Args) -> Net:
print(f"\n{'=' * 60}\nPRETRAINING (tinygrad)\n{'=' * 60}")
print(f"Batch size: {args.pretrain_bsz}, Iterations: {args.pretrain_iters}")
t0 = time.time()
optim = nn.optim.Adam(model.parameters(), lr=args.pretraining_lr)
log_interval = max(1, args.pretrain_iters // 10)
for i in range(args.pretrain_iters):
_, ctx_y, _, fut_y = load_data(args.pretrain_bsz, pretrain_dataset, args)
with Tensor.train():
optim.zero_grad()
loss = model.compute_loss(ctx_y, fut_y[:, 0:1])
loss.backward()
optim.step()
if i % log_interval == 0:
print(f"Iter {i + 1}/{args.pretrain_iters} - Loss: {float(loss.numpy()):.4f}")
print(f"Completed in {time.time() - t0:.2f}s")
return model
def randopt_scores(
base: Net,
ctx_y: Tensor,
fut_y: Tensor,
args: Args,
N: int | None = None,
sigma: float | None = None,
) -> list[tuple[int, float]]:
N = args.N if N is None else N
sigma = args.sigma if sigma is None else sigma
scores: list[tuple[int, float]] = []
for seed in range(N):
pert = base.clone()
pert.perturb_weights(seed, sigma)
scores.append((seed, eval_next_step(pert, ctx_y, fut_y)))
scores.sort(key=lambda x: x[1])
return scores
def ensemble_param_average(base: Net, seeds: list[int], sigma: float) -> Net:
snap = base.snapshot_weights()
acc = [np.zeros_like(s) for s in snap]
work = base.clone()
for seed in seeds:
work.perturb_from_snapshot(snap, seed, sigma)
for a, w in zip(acc, work.snapshot_weights()):
a += w
mean = [a / len(seeds) for a in acc]
ens = base.clone()
ens.load_weights(mean)
return ens
def evolutionary_strategies(
base: Net, ctx_y: Tensor, fut_y: Tensor, args: Args, N: int, sigma: float, lr: float = 0.05
) -> dict[str, Any]:
half = max(N // 2, 1)
snap = base.snapshot_weights()
work = base.clone()
pos_scores, neg_scores = [], []
# Store plus noise directions as weight deltas for the ES update
plus_deltas: list[list] = []
for seed in range(half):
work.perturb_from_snapshot(snap, seed, sigma)
pos_scores.append(eval_next_step(work, ctx_y, fut_y))
plus_arr = work.snapshot_weights()
# antithetic: 2*base - plus
minus_arr = [2 * b - p for b, p in zip(snap, plus_arr)]
work.load_weights(minus_arr)
neg_scores.append(eval_next_step(work, ctx_y, fut_y))
plus_deltas.append([p - b for p, b in zip(plus_arr, snap)])
updated = base.clone()
accum = [Tensor.zeros(*p.shape) for p in updated.parameters()]
for seed in range(half):
w = (neg_scores[seed] - pos_scores[seed]) / (2 * sigma + 1e-12)
for a, d in zip(accum, plus_deltas[seed]):
a.assign(a + Tensor(d.astype("float32")) * w)
a.realize()
for p, a in zip(updated.parameters(), accum):
p.assign(p + lr * a / half)
p.realize()
mse = eval_model(updated, ctx_y, fut_y, args.fut_sz)
return {
"method": "ES",
"N": N,
"K": None,
"sigma": sigma,
"best_mse": mse,
"ensemble_mse": mse,
"evals": 2 * half,
"rank_metric": "next_step_mse",
"report_metric": "ar_mse",
}
def sgd_posttrain(
base: Net, ctx_y: Tensor, fut_y: Tensor, args: Args, steps: int, lr: float = 1e-3
) -> dict[str, Any]:
"""Next-token Adam fine-tune for `steps` (matched budget proxy)."""
m = base.clone()
optim = nn.optim.Adam(m.parameters(), lr=lr)
for _ in range(steps):
with Tensor.train():
optim.zero_grad()
# next-step MSE (same as pretrain objective) — AR full-horizon backprop is too heavy
loss = m.compute_loss(ctx_y, fut_y[:, 0:1])
loss.backward()
optim.step()
mse = eval_model(m, ctx_y, fut_y, args.fut_sz)
return {
"method": "SGD",
"N": steps,
"K": None,
"sigma": None,
"best_mse": mse,
"ensemble_mse": mse,
"evals": steps,
}

Xet Storage Details

Size:
5.28 kB
·
Xet hash:
c46a6c87d0830fdcce46ab63b7a67755301c6bf2a41946736f1d6ba14e9af395

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