Buckets:
| #!/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 | |
| 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.