From 993091adfa543e38fa38cabc8a182f7a0072f04f Mon Sep 17 00:00:00 2001 From: Elias Wahl <82230675+Eliulm@users.noreply.github.com> Date: Mon, 20 May 2024 23:08:35 +0200 Subject: [PATCH] loss scaler + nan fixes (#4661) --- examples/mlperf/model_train.py | 23 ++++++++++++----------- tinygrad/tensor.py | 2 +- 2 files changed, 13 insertions(+), 12 deletions(-) diff --git a/examples/mlperf/model_train.py b/examples/mlperf/model_train.py index a89ef4ea32..af5011cd37 100644 --- a/examples/mlperf/model_train.py +++ b/examples/mlperf/model_train.py @@ -354,7 +354,7 @@ def train_rnnt(): pass @TinyJit -def train_step_bert(model, optimizer, scheduler, input_ids:Tensor, segment_ids:Tensor, attention_mask:Tensor, masked_positions:Tensor, masked_lm_ids:Tensor, masked_lm_weights:Tensor, next_sentence_labels:Tensor): +def train_step_bert(model, optimizer, scheduler, loss_scaler:float, input_ids:Tensor, segment_ids:Tensor, attention_mask:Tensor, masked_positions:Tensor, masked_lm_ids:Tensor, masked_lm_weights:Tensor, next_sentence_labels:Tensor): lm_logits, clsf_logits = model(input_ids, segment_ids, attention_mask, masked_positions) lm_loss = lm_logits.sparse_categorical_crossentropy(masked_lm_ids, ignore_index=masked_lm_weights) clsf_loss = clsf_logits.binary_crossentropy_logits(next_sentence_labels) @@ -362,8 +362,9 @@ def train_step_bert(model, optimizer, scheduler, input_ids:Tensor, segment_ids:T if not getenv('DISABLE_BACKWARD', 0): optimizer.zero_grad() - loss.backward() + (loss * loss_scaler).backward() + for p in optimizer.params: p.grad /= loss_scaler optimizer.step() scheduler.step() return loss.realize() @@ -373,11 +374,11 @@ def eval_step_bert(model, input_ids:Tensor, segment_ids:Tensor, attention_mask:T lm_logits, clsf_logits = model(input_ids, segment_ids, attention_mask, masked_positions) clsf_predictions = clsf_logits.log_softmax().argmax(-1) - clsf_accuracy = (clsf_predictions == next_sentence_labels).float().mean() + clsf_accuracy = (clsf_predictions == next_sentence_labels).mean() mlm_predictions = lm_logits.log_softmax().argmax(-1) mask = (masked_lm_weights == 1.0) - mlm_accuracy = (mlm_predictions == masked_lm_ids).where(mask, 0).sum() / mask.float().sum() + mlm_accuracy = (mlm_predictions == masked_lm_ids).where(mask, 0).sum() / mask.sum() lm_loss = lm_logits.sparse_categorical_crossentropy(masked_lm_ids, ignore_index=masked_lm_weights) clsf_loss = clsf_logits.binary_crossentropy_logits(next_sentence_labels) @@ -386,7 +387,7 @@ def eval_step_bert(model, input_ids:Tensor, segment_ids:Tensor, attention_mask:T "masked_lm_loss": lm_loss.realize(), "next_sentence_accuracy": clsf_accuracy.realize(), "next_sentence_loss": clsf_loss.realize() - } + } def train_bert(): # NOTE: pip install tensorflow, wandb required @@ -408,13 +409,14 @@ def train_bert(): max_lr = config["OPT_BASE_LEARNING_RATE"] = getenv("OPT_BASE_LEARNING_RATE", 0.000004166 * BS) train_steps = config["TRAIN_STEPS"] = getenv("TRAIN_STEPS", 4800000 // BS) - warmup_steps = config["NUM_WARMUP_STEPS"] = getenv("NUM_WARMUP_STEPS", train_steps // 10) + warmup_steps = config["NUM_WARMUP_STEPS"] = getenv("NUM_WARMUP_STEPS", 1) max_eval_steps = config["MAX_EVAL_STEPS"] = getenv("MAX_EVAL_STEPS", (10000 + EVAL_BS - 1) // EVAL_BS) # EVAL_BS * MAX_EVAL_STEPS >= 10000 eval_step_freq = config["EVAL_STEP_FREQ"] = getenv("EVAL_STEP_FREQ", int((math.floor(0.05 * (230.23 * BS + 3000000) / 25000) * 25000) / BS)) # Round down save_ckpt_freq = config["SAVE_CKPT_FREQ"] = getenv("SAVE_CKPT_FREQ", 1000) keep_ckpt_amount = config["KEEP_CKPT_AMOUNT"] = getenv("KEEP_CKPT_AMOUNT", 5) init_ckpt = config["INIT_CKPT_DIR"] = getenv("INIT_CKPT_DIR", BASEDIR) + loss_scaler = config["loss_scaler"] = getenv("LOSS_SCALER", 2**9 if dtypes.default_float == dtypes.float16 else 1.0) decay = config["decay"] = getenv("DECAY", 0.01) poly_power = config["poly_power"] = getenv("POLY_POWER", 1.0) @@ -436,14 +438,13 @@ def train_bert(): assert 10000 <= (EVAL_BS * max_eval_steps), "Evaluation batchsize * max_eval_steps must greater or equal 10000 to iterate over full eval dataset" # ** Log hparams ** - for key, value in config.items(): - print(f'HParam: "{key}": {value}') + for key, value in config.items(): print(f'HParam: "{key}": {value}') # ** Optimizer ** skip_list = [v for k, v in get_state_dict(model).items() if "bias" in k or "LayerNorm" in k] parameters = [x for x in parameters if x not in set(skip_list)] - optimizer = LAMB(parameters, 1 / warmup_steps, eps=1e-6, wd=decay, adam=False) - optimizer_skip = LAMB(skip_list, 1 / warmup_steps, eps=1e-6, wd=0.0, adam=False) + optimizer = LAMB(parameters, lr=max_lr, eps=1e-6, wd=decay, adam=False) + optimizer_skip = LAMB(skip_list, lr=max_lr, eps=1e-6, wd=0.0, adam=False) optimizer_group = OptimizerGroup(optimizer, optimizer_skip) # ** LR scheduler ** @@ -478,7 +479,7 @@ def train_bert(): BEAM.value = TRAIN_BEAM st = time.perf_counter() GlobalCounters.reset() - loss = train_step_bert(model, optimizer_group, scheduler, + loss = train_step_bert(model, optimizer_group, scheduler, loss_scaler, train_data["input_ids"], train_data["segment_ids"], train_data["input_mask"], train_data["masked_lm_positions"], \ train_data["masked_lm_ids"], train_data["masked_lm_weights"], train_data["next_sentence_labels"]) diff --git a/tinygrad/tensor.py b/tinygrad/tensor.py index 4c45b25bba..a9614e6521 100644 --- a/tinygrad/tensor.py +++ b/tinygrad/tensor.py @@ -1321,7 +1321,7 @@ class Tensor: def dropout(self, p=0.5) -> Tensor: if not Tensor.training or p == 0: return self - return self * (Tensor.rand(*self.shape, requires_grad=False, device=self.device) >= p) * (1/(1.0 - p)) + return self * (Tensor.rand(*self.shape, requires_grad=False, dtype=dtypes.default_float, device=self.device) >= p) * (1/(1.0 - p)) def one_hot(self, num_classes:int) -> Tensor: return (self[..., None] == Tensor.arange(num_classes, requires_grad=False, device=self.device)).where(1, 0)