""" train_gpt_simple.py This file descends from the [NanoGPT speedrun](https://github.com/KellerJordan/modded-nanogpt). It was prepared as a simplified version of the speedrun for use in neural net optimization research. """ import os import sys with open(sys.argv[0]) as f: code = f.read() # read the code of this file ASAP, for logging import uuid import time from pathlib import Path import torch from torch import Tensor, nn from torch.optim import AdamW import torch.nn.functional as F import torch.distributed as dist ######################################## # Dataloader # ######################################## def _load_data_shard(file: Path): header = torch.from_file(str(file), False, 256, dtype=torch.int32) # header is 256 int32 assert header[0] == 20240520, "magic number mismatch in the data .bin file" assert header[1] == 1, "unsupported version" num_tokens = int(header[2]) # number of tokens (claimed) with file.open("rb", buffering=0) as f: tokens = torch.empty(num_tokens, dtype=torch.uint16, pin_memory=True) f.seek(256 * 4) nbytes = f.readinto(tokens.numpy()) # avoid bytes->array copy assert nbytes == 2 * num_tokens, "number of tokens read does not match header" return tokens def distributed_data_generator(filename_pattern: str, batch_size: int, seq_len=1024): files = sorted(Path.cwd().glob(filename_pattern)) assert batch_size % dist.get_world_size() == 0 local_batch_size = batch_size // dist.get_world_size() file_iter = iter(files) tokens, pos = _load_data_shard(next(file_iter)), 0 while True: if pos + batch_size + 1 >= len(tokens): tokens, pos = _load_data_shard(next(file_iter)), 0 buf = tokens[pos + dist.get_rank() * local_batch_size:][:local_batch_size + 1] inputs = buf[:-1].to(device="cuda", dtype=torch.int32, non_blocking=True) targets = buf[1:].to(device="cuda", dtype=torch.int64, non_blocking=True) pos += batch_size yield inputs.view(-1, seq_len), targets.view(-1, seq_len) ######################################## # Architecture # ######################################## class RMSNorm(nn.Module): def __init__(self, dim): super().__init__() self.gains = nn.Parameter(torch.ones(dim)) def forward(self, x): return F.rms_norm(x, (x.size(-1),), weight=self.gains.type_as(x)) class Linear(nn.Linear): def __init__(self, in_features, out_features): super().__init__(in_features, out_features, bias=True) def forward(self, x): return F.linear(x, self.weight.type_as(x), self.bias.type_as(x)) class Rotary(nn.Module): def __init__(self, dim: int): super().__init__() # half-truncate RoPE (w/ base freq tuning) angular_freq = (1 / 1024) ** torch.linspace(0, 1, steps=dim//4, dtype=torch.float32) self.register_buffer("angular_freq", torch.cat([angular_freq, angular_freq.new_zeros(dim//4)])) def forward(self, x_BTHD: Tensor): pos = torch.arange(x_BTHD.size(1), dtype=torch.float32, device=x_BTHD.device) theta = torch.outer(pos, self.angular_freq)[None, :, None, :] cos, sin = theta.cos(), theta.sin() x1, x2 = x_BTHD.to(dtype=torch.float32).chunk(2, dim=-1) y1 = x1 * cos + x2 * sin y2 = x1 * (-sin) + x2 * cos return torch.cat((y1, y2), 3).type_as(x_BTHD) class CausalSelfAttention(nn.Module): def __init__(self, dim: int, head_dim=128): super().__init__() self.num_heads = dim // head_dim self.head_dim = head_dim hdim = self.num_heads * self.head_dim self.q = Linear(dim, hdim) self.k = Linear(dim, hdim) self.v = Linear(dim, hdim) self.proj = Linear(hdim, dim) self.rotary = Rotary(head_dim) def forward(self, x: Tensor): B, T = x.size(0), x.size(1) q = self.q(x).view(B, T, self.num_heads, self.head_dim) k = self.k(x).view(B, T, self.num_heads, self.head_dim) v = self.v(x).view(B, T, self.num_heads, self.head_dim) q, k = F.rms_norm(q, (q.size(-1),)), F.rms_norm(k, (k.size(-1),)) q, k = self.rotary(q), self.rotary(k) y = F.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2), scale=0.12, is_causal=True).transpose(1, 2) y = y.contiguous().view(B, T, self.num_heads * self.head_dim) y = self.proj(y) return y class MLP(nn.Module): def __init__(self, dim: int): super().__init__() hdim = 4 * dim self.fc = Linear(dim, hdim) self.proj = Linear(hdim, dim) def forward(self, x: Tensor): x = self.fc(x) x = x.relu().square() x = self.proj(x) return x class Block(nn.Module): def __init__(self, dim: int): super().__init__() self.attn = CausalSelfAttention(dim) self.mlp = MLP(dim) self.norm1 = RMSNorm(dim) self.norm2 = RMSNorm(dim) def forward(self, x: Tensor): x = x + self.attn(self.norm1(x)) x = x + self.mlp(self.norm2(x)) return x class GPT(nn.Module): def __init__(self, vocab_size: int, num_layers: int, model_dim: int): super().__init__() self.embed = nn.Embedding(vocab_size, model_dim).bfloat16() self.blocks = nn.ModuleList([Block(model_dim) for _ in range(num_layers)]) self.proj = Linear(model_dim, vocab_size) self.norm1 = RMSNorm(model_dim) self.norm2 = RMSNorm(model_dim) def forward(self, inputs: Tensor, targets: Tensor): x = self.norm1(self.embed(inputs)) for block in self.blocks: x = block(x) logits = self.proj(self.norm2(x)).float() logits = 15 * logits * (logits.square() + 15**2).rsqrt() return F.cross_entropy(logits.view(targets.numel(), -1), targets.view(-1), reduction="sum") ######################################## # Optimizer # ######################################## def zeropower_via_newtonschulz5(G: Tensor) -> Tensor: assert G.ndim >= 2 X = G.bfloat16() if G.size(-2) > G.size(-1): X = X.mT # Ensure spectral norm is at most 1 X = X / (X.norm(dim=(-2, -1), keepdim=True) + 1e-7) # Perform the NS iterations, not optimizing for wallclock speed a, b, c = 2, -1.5, 0.5 for _ in range(12): A = X @ X.mT B = b * A + c * A @ A X = a * X + B @ X if G.size(-2) > G.size(-1): X = X.mT return X @torch.compile def muon_update(grad, momentum, mu=0.95, nesterov=True): momentum.lerp_(grad, 1 - mu) update = grad.lerp_(momentum, mu) if nesterov else momentum update = zeropower_via_newtonschulz5(update) update *= max(1, grad.size(-2) / grad.size(-1))**0.5 return update class Muon(torch.optim.Optimizer): def __init__(self, params, lr=0.02, weight_decay=0, mu=0.95): assert isinstance(params, list) and len(params) >= 1 and isinstance(params[0], torch.nn.Parameter) params = sorted(params, key=lambda x: x.size(), reverse=True) defaults = dict(lr=lr, weight_decay=weight_decay, mu=mu) super().__init__(params, defaults) @torch.no_grad() def step(self): world_size = dist.get_world_size() rank = dist.get_rank() for group in self.param_groups: params = group["params"] params_pad = params + [torch.empty_like(params[-1])] * (world_size - len(params) % world_size) for base_i in range(0, len(params), world_size): if base_i + rank < len(params): p = params[base_i + rank] state = self.state[p] if len(state) == 0: state["momentum"] = torch.zeros_like(p) update = muon_update(p.grad, state["momentum"], mu=group["mu"]) p.mul_(1 - group["lr"] * group["weight_decay"]) p.add_(update, alpha=-group["lr"]) dist.all_gather(params_pad[base_i:base_i + world_size], params_pad[base_i + rank]) ######################################## # Setup # ######################################## # torchrun sets these env variables device = torch.device("cuda", int(os.environ["LOCAL_RANK"])) torch.cuda.set_device(device) dist.init_process_group(backend="nccl", device_id=device) dist.barrier() # this code can be run equivalently with 1, 2, 4, or 8 gpus. assert 8 % dist.get_world_size() == 0 # logging setup if dist.get_rank() == 0: os.makedirs("logs", exist_ok=True) logfile = f"logs/{uuid.uuid4()}.txt" print(logfile) def print0(s, console=False, log=True): if dist.get_rank() == 0: if console: print(s) if log: with open(logfile, "a") as f: print(s, file=f) # we begin by logging this file itself print0(code) print0("="*100) print0(f"Running PyTorch {torch.version.__version__} compiled for CUDA {torch.version.cuda}" + f" on {torch.cuda.get_device_name(device)} with world_size {dist.get_world_size()}") print0("="*100) val_tokens = 20 * 524288 batch_size = 8 * 64 * 1024 mbs = 64 val_inputs, val_targets = next(distributed_data_generator("data/fineweb10B/fineweb_val_*.bin", val_tokens)) model = GPT(vocab_size=50304, num_layers=12, model_dim=768).cuda() model.compile(dynamic=False) num_trials = int(sys.argv[-1]) if len(sys.argv) > 1 else 1 for _ in range(num_trials): ######################################## # Init & Optim Hyperparams # ######################################## # we want to minimize this while still reaching 3.28 val loss train_steps = 3150 # initialize model parameters for name, p in model.named_parameters(): w = p.data if name.endswith("weight"): if "proj" in name: w.zero_() elif "embed" in name: w.normal_() # default torch init else: w.normal_(std=0.33**0.5 / w.size(-1)**0.5) # default torch init elif name.endswith("bias"): w.zero_() elif name.endswith("gains"): w.normal_(mean=1, std=0) else: raise Exception(f"Uninitialized parameter: {name}") # create the optimizer(s) optimizer1 = AdamW([dict(params=[model.embed.weight], lr=0.3), dict(params=[model.proj.weight], lr=1/320), dict(params=[p for p in model.parameters() if p.ndim < 2], lr=0.01)], betas=(0.8, 0.95), eps=1e-10, weight_decay=0, fused=True) optimizer2 = Muon([p for p in model.blocks.parameters() if p.ndim >= 2], lr=0.035, weight_decay=0.025) optimizers = [optimizer1, optimizer2] assert set(p for opt in optimizers for group in opt.param_groups for p in group["params"]) == set(model.parameters()) for opt in optimizers: for group in opt.param_groups: group["initial_lr"] = group["lr"] # learning rate schedule: stable then decay def set_hparams(step, cooldown_frac=0.7): progress = step / train_steps assert 0 <= progress < 1 if progress < 1 - cooldown_frac: eta = 1.0 else: eta = (1 - progress) / cooldown_frac for opt in optimizers: for group in opt.param_groups: group["lr"] = group["initial_lr"] * eta ######################################## # Training and Validation # ######################################## train_loader = distributed_data_generator("data/fineweb10B/fineweb_train_*.bin", batch_size) for p in model.parameters(): dist.broadcast(p.detach(), 0) # start the clock training_time = 0 last_val_step = 0 dist.barrier() t0 = time.perf_counter() for step in range(train_steps + 1): # --------------- VALIDATION SECTION ----------------- val_step_freq = 125 if step / train_steps < 0.9 else 25 if step == train_steps or step % val_step_freq == 0: # stop the clock dist.barrier() time_since_last_val = time.perf_counter() - t0 step_avg = time_since_last_val / (step - last_val_step) if step > 0 else float("nan") last_val_step = step training_time += time_since_last_val model.eval() val_loss = 0 with torch.no_grad(): assert len(val_inputs) % mbs == 0 for i in range(len(val_inputs) // mbs): val_loss += model(val_inputs[i*mbs:(i+1)*mbs], val_targets[i*mbs:(i+1)*mbs]) dist.all_reduce(val_loss, op=dist.ReduceOp.SUM) val_loss /= val_tokens print0(f"step:{step}/{train_steps} val_loss:{val_loss:.5f} train_time:{training_time:.3f}s" + f" step_avg:{1000*step_avg:.2f}ms", console=True) model.train() # start the clock again dist.barrier() t0 = time.perf_counter() if step == train_steps: break # --------------- TRAINING SECTION ----------------- inputs, targets = next(train_loader) # accumulate across microbatches in case we are running with fewer than 8 gpus assert len(inputs) % mbs == 0 for i in range(len(inputs) // mbs): model(inputs[i*mbs:(i+1)*mbs], targets[i*mbs:(i+1)*mbs]).backward() for name, p in model.named_parameters(): assert p.grad is not None, name dist.all_reduce(p.grad, op=dist.ReduceOp.SUM) # set optimization hyperparameters and take a step set_hparams(step) for opt in optimizers: opt.step() model.zero_grad(set_to_none=True) approx_training_time = training_time + (time.perf_counter() - t0) print0(f"step:{step+1}/{train_steps} train_time:{approx_training_time:.3f}s" + f" step_avg:{1000*approx_training_time/(step + 1):.2f}ms", console=True, log=False) dist.destroy_process_group() ==================================================================================================== Running PyTorch 2.11.0+cu128 compiled for CUDA 12.8 on NVIDIA H100 80GB HBM3 with world_size 8 ==================================================================================================== step:0/3150 val_loss:10.82584 train_time:0.000s step_avg:nanms step:125/3150 val_loss:4.67334 train_time:22.955s step_avg:183.64ms step:250/3150 val_loss:4.12866 train_time:41.441s step_avg:147.89ms step:375/3150 val_loss:3.94055 train_time:59.789s step_avg:146.78ms step:500/3150 val_loss:3.83494 train_time:78.162s step_avg:146.98ms step:625/3150 val_loss:3.76399 train_time:96.540s step_avg:147.03ms step:750/3150 val_loss:3.71678 train_time:114.906s step_avg:146.93ms step:875/3150 val_loss:3.67816 train_time:133.300s step_avg:147.15ms step:1000/3150 val_loss:3.63895 train_time:151.696s step_avg:147.16ms step:1125/3150 val_loss:3.60852 train_time:170.052s step_avg:146.85ms step:1250/3150 val_loss:3.57628 train_time:188.425s step_avg:146.98ms step:1375/3150 val_loss:3.54910 train_time:206.807s step_avg:147.06ms step:1500/3150 val_loss:3.52081 train_time:225.178s step_avg:146.96ms step:1625/3150 val_loss:3.50028 train_time:243.558s step_avg:147.04ms step:1750/3150 val_loss:3.47812 train_time:262.023s step_avg:147.72ms step:1875/3150 val_loss:3.45665 train_time:280.368s step_avg:146.76ms step:2000/3150 val_loss:3.43638 train_time:298.760s step_avg:147.14ms step:2125/3150 val_loss:3.41764 train_time:317.146s step_avg:147.09ms step:2250/3150 val_loss:3.39949 train_time:335.508s step_avg:146.89ms step:2375/3150 val_loss:3.38198 train_time:353.882s step_avg:147.00ms step:2500/3150 val_loss:3.36425 train_time:372.247s step_avg:146.92ms step:2625/3150 val_loss:3.34638 train_time:390.587s step_avg:146.72ms step:2750/3150 val_loss:3.32953 train_time:408.943s step_avg:146.84ms step:2850/3150 val_loss:3.31663 train_time:423.612s step_avg:146.69ms step:2875/3150 val_loss:3.31367 train_time:427.308s step_avg:147.84ms step:2900/3150 val_loss:3.31060 train_time:430.984s step_avg:147.03ms step:2925/3150 val_loss:3.30729 train_time:434.660s step_avg:147.03ms step:2950/3150 val_loss:3.30456 train_time:438.335s step_avg:147.02ms step:2975/3150 val_loss:3.30182 train_time:442.014s step_avg:147.13ms step:3000/3150 val_loss:3.29936 train_time:445.685s step_avg:146.85ms step:3025/3150 val_loss:3.29656 train_time:449.358s step_avg:146.92ms step:3050/3150 val_loss:3.29409 train_time:453.049s step_avg:147.63ms step:3075/3150 val_loss:3.29217 train_time:456.724s step_avg:147.00ms step:3100/3150 val_loss:3.29045 train_time:460.399s step_avg:147.01ms step:3125/3150 val_loss:3.28914 train_time:464.081s step_avg:147.27ms step:3150/3150 val_loss:3.28850 train_time:467.759s step_avg:147.14ms step:0/3150 val_loss:10.82584 train_time:0.000s step_avg:nanms step:125/3150 val_loss:4.66504 train_time:18.450s step_avg:147.60ms step:250/3150 val_loss:4.12423 train_time:36.847s step_avg:147.17ms step:375/3150 val_loss:3.93878 train_time:55.241s step_avg:147.16ms step:500/3150 val_loss:3.83191 train_time:73.649s step_avg:147.26ms step:625/3150 val_loss:3.76386 train_time:92.056s step_avg:147.25ms step:750/3150 val_loss:3.71540 train_time:110.416s step_avg:146.88ms step:875/3150 val_loss:3.67593 train_time:128.806s step_avg:147.12ms step:1000/3150 val_loss:3.63955 train_time:147.185s step_avg:147.03ms step:1125/3150 val_loss:3.61170 train_time:165.557s step_avg:146.98ms step:1250/3150 val_loss:3.57663 train_time:183.936s step_avg:147.03ms step:1375/3150 val_loss:3.54938 train_time:202.311s step_avg:147.00ms step:1500/3150 val_loss:3.52130 train_time:220.663s step_avg:146.82ms step:1625/3150 val_loss:3.50218 train_time:239.045s step_avg:147.06ms step:1750/3150 val_loss:3.47826 train_time:257.423s step_avg:147.03ms step:1875/3150 val_loss:3.45732 train_time:275.794s step_avg:146.96ms step:2000/3150 val_loss:3.43711 train_time:294.182s step_avg:147.11ms step:2125/3150 val_loss:3.41856 train_time:312.557s step_avg:147.00ms step:2250/3150 val_loss:3.39978 train_time:330.911s step_avg:146.84ms step:2375/3150 val_loss:3.38217 train_time:349.346s step_avg:147.47ms step:2500/3150 val_loss:3.36473 train_time:367.713s step_avg:146.94ms step:2625/3150 val_loss:3.34676 train_time:386.080s step_avg:146.94ms step:2750/3150 val_loss:3.33013 train_time:404.451s step_avg:146.97ms step:2850/3150 val_loss:3.31694 train_time:419.131s step_avg:146.80ms step:2875/3150 val_loss:3.31398 train_time:422.838s step_avg:148.28ms step:2900/3150 val_loss:3.31093 train_time:426.514s step_avg:147.05ms step:2925/3150 val_loss:3.30769 train_time:430.189s step_avg:146.99ms step:2950/3150 val_loss:3.30492 train_time:433.867s step_avg:147.12ms step:2975/3150 val_loss:3.30201 train_time:437.546s step_avg:147.17ms step:3000/3150 val_loss:3.29959 train_time:441.225s step_avg:147.18ms step:3025/3150 val_loss:3.29679 train_time:444.904s step_avg:147.14ms step:3050/3150 val_loss:3.29438 train_time:448.603s step_avg:147.95ms step:3075/3150 val_loss:3.29248 train_time:452.284s step_avg:147.24ms step:3100/3150 val_loss:3.29073 train_time:455.957s step_avg:146.94ms step:3125/3150 val_loss:3.28942 train_time:459.631s step_avg:146.96ms step:3150/3150 val_loss:3.28879 train_time:463.304s step_avg:146.93ms step:0/3150 val_loss:10.82584 train_time:0.000s step_avg:nanms step:125/3150 val_loss:4.69016 train_time:18.472s step_avg:147.77ms step:250/3150 val_loss:4.12766 train_time:36.894s step_avg:147.38ms step:375/3150 val_loss:3.94146 train_time:55.318s step_avg:147.39ms step:500/3150 val_loss:3.83209 train_time:73.706s step_avg:147.10ms step:625/3150 val_loss:3.76454 train_time:92.086s step_avg:147.05ms step:750/3150 val_loss:3.71676 train_time:110.457s step_avg:146.97ms step:875/3150 val_loss:3.67577 train_time:128.856s step_avg:147.19ms step:1000/3150 val_loss:3.64021 train_time:147.251s step_avg:147.16ms step:1125/3150 val_loss:3.60741 train_time:165.624s step_avg:146.98ms step:1250/3150 val_loss:3.57558 train_time:183.992s step_avg:146.94ms step:1375/3150 val_loss:3.54994 train_time:202.377s step_avg:147.08ms step:1500/3150 val_loss:3.51944 train_time:220.764s step_avg:147.10ms step:1625/3150 val_loss:3.49928 train_time:239.162s step_avg:147.18ms step:1750/3150 val_loss:3.47653 train_time:257.555s step_avg:147.15ms step:1875/3150 val_loss:3.45559 train_time:275.930s step_avg:147.00ms step:2000/3150 val_loss:3.43568 train_time:294.340s step_avg:147.28ms step:2125/3150 val_loss:3.41716 train_time:312.721s step_avg:147.05ms step:2250/3150 val_loss:3.39856 train_time:331.088s step_avg:146.93ms step:2375/3150 val_loss:3.38080 train_time:349.479s step_avg:147.13ms step:2500/3150 val_loss:3.36314 train_time:367.865s step_avg:147.08ms step:2625/3150 val_loss:3.34536 train_time:386.234s step_avg:146.95ms step:2750/3150 val_loss:3.32853 train_time:404.618s step_avg:147.07ms step:2850/3150 val_loss:3.31552 train_time:419.309s step_avg:146.92ms step:2875/3150 val_loss:3.31234 train_time:423.013s step_avg:148.13ms step:2900/3150 val_loss:3.30917 train_time:426.686s step_avg:146.92ms step:2925/3150 val_loss:3.30602 train_time:430.356s step_avg:146.83ms step:2950/3150 val_loss:3.30325 train_time:434.022s step_avg:146.63ms step:2975/3150 val_loss:3.30056 train_time:437.698s step_avg:147.01ms step:3000/3150 val_loss:3.29803 train_time:441.388s step_avg:147.64ms step:3025/3150 val_loss:3.29523 train_time:445.082s step_avg:147.75ms step:3050/3150 val_loss:3.29284 train_time:448.773s step_avg:147.62ms step:3075/3150 val_loss:3.29084 train_time:452.447s step_avg:146.98ms step:3100/3150 val_loss:3.28910 train_time:456.120s step_avg:146.93ms step:3125/3150 val_loss:3.28779 train_time:459.792s step_avg:146.87ms step:3150/3150 val_loss:3.28717 train_time:463.469s step_avg:147.05ms step:0/3150 val_loss:10.82584 train_time:0.000s step_avg:nanms step:125/3150 val_loss:4.66189 train_time:18.447s step_avg:147.58ms step:250/3150 val_loss:4.12548 train_time:36.872s step_avg:147.40ms step:375/3150 val_loss:3.93869 train_time:55.272s step_avg:147.20ms step:500/3150 val_loss:3.83318 train_time:73.680s step_avg:147.26ms step:625/3150 val_loss:3.76475 train_time:92.101s step_avg:147.37ms step:750/3150 val_loss:3.71519 train_time:110.481s step_avg:147.04ms step:875/3150 val_loss:3.67737 train_time:128.881s step_avg:147.20ms step:1000/3150 val_loss:3.63990 train_time:147.274s step_avg:147.14ms step:1125/3150 val_loss:3.60777 train_time:165.660s step_avg:147.09ms step:1250/3150 val_loss:3.57631 train_time:184.063s step_avg:147.22ms step:1375/3150 val_loss:3.54992 train_time:202.448s step_avg:147.08ms step:1500/3150 val_loss:3.52233 train_time:220.831s step_avg:147.07ms step:1625/3150 val_loss:3.50186 train_time:239.227s step_avg:147.16ms step:1750/3150 val_loss:3.47761 train_time:257.620s step_avg:147.15ms step:1875/3150 val_loss:3.45581 train_time:276.001s step_avg:147.05ms step:2000/3150 val_loss:3.43646 train_time:294.400s step_avg:147.19ms step:2125/3150 val_loss:3.41811 train_time:312.796s step_avg:147.17ms step:2250/3150 val_loss:3.39904 train_time:331.172s step_avg:147.01ms step:2375/3150 val_loss:3.38176 train_time:349.575s step_avg:147.22ms step:2500/3150 val_loss:3.36424 train_time:367.972s step_avg:147.18ms step:2625/3150 val_loss:3.34623 train_time:386.360s step_avg:147.10ms step:2750/3150 val_loss:3.32936 train_time:404.760s step_avg:147.20ms step:2850/3150 val_loss:3.31648 train_time:419.459s step_avg:147.00ms step:2875/3150 val_loss:3.31341 train_time:423.156s step_avg:147.86ms step:2900/3150 val_loss:3.31043 train_time:426.830s step_avg:146.97ms step:2925/3150 val_loss:3.30705 train_time:430.507s step_avg:147.06ms step:2950/3150 val_loss:3.30429 train_time:434.177s step_avg:146.81ms step:2975/3150 val_loss:3.30150 train_time:437.848s step_avg:146.83ms step:3000/3150 val_loss:3.29906 train_time:441.526s step_avg:147.15ms step:3025/3150 val_loss:3.29625 train_time:445.204s step_avg:147.09ms step:3050/3150 val_loss:3.29374 train_time:448.896s step_avg:147.68ms step:3075/3150 val_loss:3.29181 train_time:452.570s step_avg:146.97ms step:3100/3150 val_loss:3.29008 train_time:456.244s step_avg:146.96ms step:3125/3150 val_loss:3.28878 train_time:459.921s step_avg:147.09ms step:3150/3150 val_loss:3.28815 train_time:463.596s step_avg:147.01ms step:0/3150 val_loss:10.82584 train_time:0.000s step_avg:nanms step:125/3150 val_loss:4.68075 train_time:18.447s step_avg:147.58ms step:250/3150 val_loss:4.12675 train_time:36.866s step_avg:147.35ms step:375/3150 val_loss:3.94234 train_time:55.265s step_avg:147.19ms step:500/3150 val_loss:3.83887 train_time:73.741s step_avg:147.81ms step:625/3150 val_loss:3.76494 train_time:92.141s step_avg:147.20ms step:750/3150 val_loss:3.71828 train_time:110.502s step_avg:146.89ms step:875/3150 val_loss:3.67681 train_time:128.903s step_avg:147.20ms step:1000/3150 val_loss:3.63987 train_time:147.293s step_avg:147.12ms step:1125/3150 val_loss:3.61066 train_time:165.656s step_avg:146.91ms step:1250/3150 val_loss:3.57570 train_time:184.046s step_avg:147.12ms step:1375/3150 val_loss:3.55066 train_time:202.434s step_avg:147.10ms step:1500/3150 val_loss:3.51966 train_time:220.788s step_avg:146.84ms step:1625/3150 val_loss:3.50024 train_time:239.162s step_avg:146.99ms step:1750/3150 val_loss:3.47751 train_time:257.547s step_avg:147.08ms step:1875/3150 val_loss:3.45669 train_time:275.914s step_avg:146.93ms step:2000/3150 val_loss:3.43580 train_time:294.308s step_avg:147.15ms step:2125/3150 val_loss:3.41750 train_time:312.699s step_avg:147.13ms step:2250/3150 val_loss:3.39891 train_time:331.084s step_avg:147.08ms step:2375/3150 val_loss:3.38121 train_time:349.486s step_avg:147.22ms step:2500/3150 val_loss:3.36362 train_time:367.887s step_avg:147.21ms step:2625/3150 val_loss:3.34589 train_time:386.261s step_avg:146.99ms step:2750/3150 val_loss:3.32914 train_time:404.661s step_avg:147.21ms step:2850/3150 val_loss:3.31593 train_time:419.374s step_avg:147.13ms step:2875/3150 val_loss:3.31274 train_time:423.080s step_avg:148.23ms step:2900/3150 val_loss:3.30990 train_time:426.759s step_avg:147.17ms step:2925/3150 val_loss:3.30643 train_time:430.436s step_avg:147.07ms step:2950/3150 val_loss:3.30369 train_time:434.118s step_avg:147.29ms step:2975/3150 val_loss:3.30105 train_time:437.798s step_avg:147.18ms step:3000/3150 val_loss:3.29851 train_time:441.472s step_avg:146.97ms step:3025/3150 val_loss:3.29574 train_time:445.152s step_avg:147.22ms step:3050/3150 val_loss:3.29328 train_time:448.850s step_avg:147.92ms step:3075/3150 val_loss:3.29134 train_time:452.536s step_avg:147.43ms step:3100/3150 val_loss:3.28957 train_time:456.214s step_avg:147.12ms step:3125/3150 val_loss:3.28828 train_time:459.891s step_avg:147.07ms step:3150/3150 val_loss:3.28765 train_time:463.572s step_avg:147.23ms step:0/3150 val_loss:10.82584 train_time:0.000s step_avg:nanms step:125/3150 val_loss:4.68110 train_time:18.487s step_avg:147.90ms step:250/3150 val_loss:4.13439 train_time:36.930s step_avg:147.54ms step:375/3150 val_loss:3.94315 train_time:55.346s step_avg:147.32ms step:500/3150 val_loss:3.83463 train_time:73.775s step_avg:147.43ms step:625/3150 val_loss:3.76575 train_time:92.179s step_avg:147.23ms step:750/3150 val_loss:3.71709 train_time:110.568s step_avg:147.11ms step:875/3150 val_loss:3.67712 train_time:128.983s step_avg:147.32ms step:1000/3150 val_loss:3.63910 train_time:147.396s step_avg:147.31ms step:1125/3150 val_loss:3.60790 train_time:165.775s step_avg:147.03ms step:1250/3150 val_loss:3.57668 train_time:184.260s step_avg:147.88ms step:1375/3150 val_loss:3.54864 train_time:202.672s step_avg:147.29ms step:1500/3150 val_loss:3.52007 train_time:221.067s step_avg:147.16ms step:1625/3150 val_loss:3.49983 train_time:239.457s step_avg:147.12ms step:1750/3150 val_loss:3.47608 train_time:257.855s step_avg:147.18ms step:1875/3150 val_loss:3.45474 train_time:276.238s step_avg:147.06ms step:2000/3150 val_loss:3.43480 train_time:294.638s step_avg:147.20ms step:2125/3150 val_loss:3.41689 train_time:313.032s step_avg:147.15ms step:2250/3150 val_loss:3.39786 train_time:331.404s step_avg:146.98ms step:2375/3150 val_loss:3.37998 train_time:349.811s step_avg:147.25ms step:2500/3150 val_loss:3.36243 train_time:368.209s step_avg:147.19ms step:2625/3150 val_loss:3.34473 train_time:386.562s step_avg:146.82ms step:2750/3150 val_loss:3.32794 train_time:404.945s step_avg:147.07ms step:2850/3150 val_loss:3.31471 train_time:419.640s step_avg:146.95ms step:2875/3150 val_loss:3.31174 train_time:423.345s step_avg:148.24ms step:2900/3150 val_loss:3.30879 train_time:427.021s step_avg:147.04ms step:2925/3150 val_loss:3.30538 train_time:430.701s step_avg:147.16ms step:2950/3150 val_loss:3.30268 train_time:434.376s step_avg:147.04ms step:2975/3150 val_loss:3.30000 train_time:438.057s step_avg:147.21ms step:3000/3150 val_loss:3.29749 train_time:441.731s step_avg:146.97ms step:3025/3150 val_loss:3.29460 train_time:445.410s step_avg:147.17ms step:3050/3150 val_loss:3.29222 train_time:449.102s step_avg:147.67ms step:3075/3150 val_loss:3.29025 train_time:452.780s step_avg:147.12ms step:3100/3150 val_loss:3.28853 train_time:456.457s step_avg:147.08ms step:3125/3150 val_loss:3.28721 train_time:460.132s step_avg:146.98ms step:3150/3150 val_loss:3.28658 train_time:463.810s step_avg:147.12ms step:0/3150 val_loss:10.82584 train_time:0.000s step_avg:nanms step:125/3150 val_loss:4.67479 train_time:18.440s step_avg:147.52ms step:250/3150 val_loss:4.12818 train_time:36.855s step_avg:147.32ms step:375/3150 val_loss:3.94323 train_time:55.266s step_avg:147.29ms step:500/3150 val_loss:3.83524 train_time:73.687s step_avg:147.36ms step:625/3150 val_loss:3.76755 train_time:92.100s step_avg:147.31ms step:750/3150 val_loss:3.72052 train_time:110.489s step_avg:147.11ms step:875/3150 val_loss:3.67796 train_time:128.894s step_avg:147.24ms step:1000/3150 val_loss:3.64028 train_time:147.314s step_avg:147.36ms step:1125/3150 val_loss:3.60926 train_time:165.719s step_avg:147.24ms step:1250/3150 val_loss:3.57796 train_time:184.122s step_avg:147.22ms step:1375/3150 val_loss:3.55110 train_time:202.524s step_avg:147.22ms step:1500/3150 val_loss:3.52239 train_time:220.901s step_avg:147.01ms step:1625/3150 val_loss:3.50166 train_time:239.297s step_avg:147.17ms step:1750/3150 val_loss:3.47796 train_time:257.685s step_avg:147.10ms step:1875/3150 val_loss:3.45733 train_time:276.116s step_avg:147.45ms step:2000/3150 val_loss:3.43788 train_time:294.508s step_avg:147.14ms step:2125/3150 val_loss:3.41949 train_time:312.908s step_avg:147.20ms step:2250/3150 val_loss:3.40105 train_time:331.283s step_avg:147.00ms step:2375/3150 val_loss:3.38386 train_time:349.656s step_avg:146.98ms step:2500/3150 val_loss:3.36596 train_time:368.036s step_avg:147.04ms step:2625/3150 val_loss:3.34800 train_time:386.387s step_avg:146.81ms step:2750/3150 val_loss:3.33103 train_time:404.765s step_avg:147.02ms step:2850/3150 val_loss:3.31817 train_time:419.463s step_avg:146.98ms step:2875/3150 val_loss:3.31521 train_time:423.158s step_avg:147.83ms step:2900/3150 val_loss:3.31223 train_time:426.834s step_avg:147.03ms step:2925/3150 val_loss:3.30879 train_time:430.510s step_avg:147.03ms step:2950/3150 val_loss:3.30600 train_time:434.186s step_avg:147.04ms step:2975/3150 val_loss:3.30326 train_time:437.864s step_avg:147.12ms step:3000/3150 val_loss:3.30081 train_time:441.546s step_avg:147.31ms step:3025/3150 val_loss:3.29797 train_time:445.229s step_avg:147.31ms step:3050/3150 val_loss:3.29558 train_time:448.929s step_avg:147.97ms step:3075/3150 val_loss:3.29361 train_time:452.607s step_avg:147.15ms step:3100/3150 val_loss:3.29186 train_time:456.286s step_avg:147.16ms step:3125/3150 val_loss:3.29055 train_time:459.963s step_avg:147.05ms step:3150/3150 val_loss:3.28992 train_time:463.636s step_avg:146.93ms step:0/3150 val_loss:10.82584 train_time:0.000s step_avg:nanms step:125/3150 val_loss:4.68238 train_time:18.471s step_avg:147.77ms step:250/3150 val_loss:4.12620 train_time:36.905s step_avg:147.47ms step:375/3150 val_loss:3.94133 train_time:55.315s step_avg:147.28ms step:500/3150 val_loss:3.83476 train_time:73.744s step_avg:147.43ms step:625/3150 val_loss:3.76258 train_time:92.151s step_avg:147.26ms step:750/3150 val_loss:3.71732 train_time:110.553s step_avg:147.22ms step:875/3150 val_loss:3.67605 train_time:128.964s step_avg:147.29ms step:1000/3150 val_loss:3.63871 train_time:147.360s step_avg:147.17ms step:1125/3150 val_loss:3.61077 train_time:165.755s step_avg:147.16ms step:1250/3150 val_loss:3.57728 train_time:184.152s step_avg:147.17ms step:1375/3150 val_loss:3.55035 train_time:202.543s step_avg:147.13ms step:1500/3150 val_loss:3.52142 train_time:220.900s step_avg:146.86ms step:1625/3150 val_loss:3.50084 train_time:239.276s step_avg:147.00ms step:1750/3150 val_loss:3.47769 train_time:257.643s step_avg:146.94ms step:1875/3150 val_loss:3.45746 train_time:275.996s step_avg:146.83ms step:2000/3150 val_loss:3.43620 train_time:294.383s step_avg:147.10ms step:2125/3150 val_loss:3.41783 train_time:312.755s step_avg:146.97ms step:2250/3150 val_loss:3.39891 train_time:331.098s step_avg:146.75ms step:2375/3150 val_loss:3.38160 train_time:349.472s step_avg:146.99ms step:2500/3150 val_loss:3.36386 train_time:367.919s step_avg:147.58ms step:2625/3150 val_loss:3.34623 train_time:386.297s step_avg:147.02ms step:2750/3150 val_loss:3.32942 train_time:404.694s step_avg:147.17ms step:2850/3150 val_loss:3.31634 train_time:419.417s step_avg:147.23ms step:2875/3150 val_loss:3.31316 train_time:423.114s step_avg:147.86ms step:2900/3150 val_loss:3.31020 train_time:426.792s step_avg:147.14ms step:2925/3150 val_loss:3.30693 train_time:430.465s step_avg:146.92ms step:2950/3150 val_loss:3.30416 train_time:434.140s step_avg:147.00ms step:2975/3150 val_loss:3.30149 train_time:437.818s step_avg:147.13ms step:3000/3150 val_loss:3.29892 train_time:441.493s step_avg:146.97ms step:3025/3150 val_loss:3.29606 train_time:445.171s step_avg:147.15ms step:3050/3150 val_loss:3.29367 train_time:448.866s step_avg:147.77ms step:3075/3150 val_loss:3.29172 train_time:452.542s step_avg:147.05ms step:3100/3150 val_loss:3.28999 train_time:456.222s step_avg:147.23ms step:3125/3150 val_loss:3.28871 train_time:459.900s step_avg:147.12ms step:3150/3150 val_loss:3.28807 train_time:463.576s step_avg:147.04ms