From 6a2168e6985fd609aa533573f871b9bcf291eade Mon Sep 17 00:00:00 2001 From: chenyu Date: Mon, 15 Apr 2024 14:57:21 -0400 Subject: [PATCH] TRAIN_BEAM and EVAL_BEAM for resnet (#4177) working on measuring compile time --- examples/mlperf/model_train.py | 36 +++++++++++++++++++--------------- 1 file changed, 20 insertions(+), 16 deletions(-) diff --git a/examples/mlperf/model_train.py b/examples/mlperf/model_train.py index e05d5d4547..df6d80deb1 100644 --- a/examples/mlperf/model_train.py +++ b/examples/mlperf/model_train.py @@ -5,7 +5,7 @@ from tqdm import tqdm import multiprocessing from tinygrad import Device, GlobalCounters, Tensor, TinyJit, dtypes -from tinygrad.helpers import getenv, BEAM, WINO +from tinygrad.helpers import getenv, BEAM, WINO, Context from tinygrad.nn.state import get_parameters, get_state_dict, safe_load, safe_save from tinygrad.nn.optim import LARS, SGD, OptimizerGroup @@ -105,23 +105,27 @@ def train_resnet(): def normalize(x): return (x.permute([0, 3, 1, 2]) - input_mean).cast(dtypes.default_float) @TinyJit def train_step(X, Y): - optimizer_group.zero_grad() - X = normalize(X) - out = model.forward(X) - loss = out.cast(dtypes.float32).sparse_categorical_crossentropy(Y, label_smoothing=0.1) - top_1 = (out.argmax(-1) == Y).sum() - (loss * loss_scaler).backward() - for t in optimizer_group.params: t.grad = t.grad.contiguous() / loss_scaler - optimizer_group.step() - scheduler_group.step() - return loss.realize(), top_1.realize() + with Context(BEAM=getenv("TRAIN_BEAM", BEAM.value)): + optimizer_group.zero_grad() + X = normalize(X) + out = model.forward(X) + loss = out.cast(dtypes.float32).sparse_categorical_crossentropy(Y, label_smoothing=0.1) + top_1 = (out.argmax(-1) == Y).sum() + (loss * loss_scaler).backward() + for t in optimizer_group.params: t.grad = t.grad.contiguous() / loss_scaler + optimizer_group.step() + scheduler_group.step() + return loss.realize(), top_1.realize() + @TinyJit def eval_step(X, Y): - X = normalize(X) - out = model.forward(X) - loss = out.cast(dtypes.float32).sparse_categorical_crossentropy(Y, label_smoothing=0.1) - top_1 = (out.argmax(-1) == Y).sum() - return loss.realize(), top_1.realize() + with Context(BEAM=getenv("EVAL_BEAM", BEAM.value)): + X = normalize(X) + out = model.forward(X) + loss = out.cast(dtypes.float32).sparse_categorical_crossentropy(Y, label_smoothing=0.1) + top_1 = (out.argmax(-1) == Y).sum() + return loss.realize(), top_1.realize() + def data_get(it): x, y, cookie = next(it) return x.shard(GPUS, axis=0).realize(), Tensor(y, requires_grad=False).shard(GPUS, axis=0), cookie