import os import sys # Read the current file and the kernels file code ASAP, for logging with open(sys.argv[0], 'r') as f: code = f.read() with open(os.path.join(os.path.dirname(sys.argv[0]), 'triton_kernels.py'), 'r') as f: code += f"\n\n{'-'*40}\n# triton_kernels.py\n{'-'*40}\n\n" code += f.read() import copy import glob import math import threading import time import uuid from dataclasses import dataclass from itertools import accumulate from pathlib import Path import gc os.environ["PYTORCH_ALLOC_CONF"] = "expandable_segments:True" import torch import triton torch.empty( 1, device=f"cuda:{os.environ['LOCAL_RANK']}", requires_grad=True ).backward() # prevents a bug on some systems import torch._dynamo as dynamo import torch.distributed as dist import torch.nn.functional as F # torch._inductor.config.coordinate_descent_tuning = True # we have banned this flag for new records because it causes compilation to take 30min from kernels import get_kernel from torch import Tensor, nn from triton_kernels import XXT, ba_plus_cAA, FusedLinearReLUSquareFunction, FusedSoftcappedCrossEntropy dynamo.config.recompile_limit = 64 # ----------------------------------------------------------------------------- # Custom operators: FP8 matmul by @YouJiacheng # Transposed layout by @ChrisJMcCormick allows for faster gradient accumulation. @torch.library.custom_op("nanogpt::mm_t", mutates_args=()) def mm_t_op(x: Tensor, w: Tensor, x_s: float, w_s: float, grad_s: float) -> tuple[Tensor, Tensor, Tensor]: """Computes y = x @ w with F8 weights stored as (in_features, out_features).""" @torch.compile def impl(x: Tensor, w: Tensor): assert x.is_contiguous() and w.is_contiguous() assert x.shape[1] == w.shape[0] # x: (batch, in), w: (in, out) x_f8 = x.div(x_s).to(torch.float8_e4m3fn) w_f8 = w.div(w_s).to(torch.float8_e4m3fn) # _scaled_mm requires column-major B. w_f8 is row-major (in, out). # .T.contiguous().T creates a column-major view without changing logical shape. w_f8_col_major = w_f8.T.contiguous().T out = torch._scaled_mm( x_f8, w_f8_col_major, out_dtype=torch.bfloat16, scale_a=x.new_tensor(x_s, dtype=torch.float32), scale_b=x.new_tensor(w_s, dtype=torch.float32), use_fast_accum=True, ) return out, x_f8, w_f8 return impl(x, w) @mm_t_op.register_fake def _(x: Tensor, w: Tensor, *_): assert x.ndim == w.ndim == 2 assert x.shape[1] == w.shape[0] assert x.device == w.device assert x.is_contiguous() and w.is_contiguous() return x @ w, x.to(torch.float8_e4m3fn), w.to(torch.float8_e4m3fn) @torch.library.custom_op("nanogpt::mm_t_backward", mutates_args=()) def mm_t_backward_op(g: Tensor, x_f8: Tensor, w_f8: Tensor, x_s: float, w_s: float, grad_s: float) -> tuple[Tensor, Tensor]: @torch.compile def impl(grad: Tensor, x_f8: Tensor, w_f8: Tensor): assert grad.is_contiguous() x_scale = grad.new_tensor(x_s, dtype=torch.float32) w_scale = grad.new_tensor(w_s, dtype=torch.float32) grad_scale = grad.new_tensor(grad_s, dtype=torch.float32) grad_f8 = grad.div(grad_s).to(torch.float8_e5m2) # grad_x = grad @ w.T grad_x = torch._scaled_mm( grad_f8, w_f8.T, out_dtype=torch.bfloat16, scale_a=grad_scale, scale_b=w_scale, use_fast_accum=False, ) # grad_w = x.T @ grad # Result is (in, out), naturally matching weight storage. No final .T needed. grad_w = torch._scaled_mm( x_f8.T.contiguous(), grad_f8.T.contiguous().T, out_dtype=torch.float32, scale_a=x_scale, scale_b=grad_scale, use_fast_accum=False, ) return grad_x, grad_w grad_x, grad_w = impl(g, x_f8, w_f8) return grad_x, grad_w @mm_t_backward_op.register_fake def _(g: Tensor, x_f8: Tensor, w_f8: Tensor, *_): return x_f8.to(torch.bfloat16), w_f8.to(torch.float32) def backward_t(ctx, grad_out: Tensor, *_): x_f8, w_f8 = ctx.saved_tensors x_s, w_s, grad_s = ctx.scales grad_x, grad_w = torch.ops.nanogpt.mm_t_backward( grad_out, x_f8, w_f8, x_s, w_s, grad_s ) return grad_x, grad_w, None, None, None def setup_context_t(ctx: torch.autograd.function.FunctionCtx, inputs, output): *_, x_s, w_s, grad_s = inputs _, x_f8, w_f8 = output ctx.save_for_backward(x_f8, w_f8) ctx.scales = x_s, w_s, grad_s ctx.set_materialize_grads(False) mm_t_op.register_autograd(backward_t, setup_context=setup_context_t) # ----------------------------------------------------------------------------- # Polar Express # Computed for num_iters=5, safety_factor=2e-2, cushion=2 polar_express_coeffs = [ (8.156554524902461, -22.48329292557795, 15.878769915207462), (4.042929935166739, -2.808917465908714, 0.5000178451051316), (3.8916678022926607, -2.772484153217685, 0.5060648178503393), (3.285753657755655, -2.3681294933425376, 0.46449024233003106), (2.3465413258596377, -1.7097828382687081, 0.42323551169305323) ] @torch.compile(dynamic=False, fullgraph=True) # Must use dynamic=False or else it's much slower def polar_express(G: torch.Tensor, split_baddbmm: bool = False): """ Polar Express Sign Method: https://arxiv.org/pdf/2505.16932 by Noah Amsel, David Persson, Christopher Musco, Robert M. Gower. """ 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) * (1 + 2e-2) + 1e-6) # Allocate buffers X = X.contiguous() A = torch.empty((*X.shape[:-1], X.size(-2)), device=X.device, dtype=X.dtype) B = torch.empty_like(A) C = torch.empty_like(X) # Select batched vs unbatched if split_baddbmm: BX_matmul = torch.bmm if X.ndim > 2 else torch.mm else: aX_plus_BX = torch.baddbmm if X.ndim > 2 else torch.addmm # Perform the iterations for a, b, c in polar_express_coeffs: XXT(X, out=A) # A = X @ X.mT ba_plus_cAA(A, alpha=c, beta=b, out=B) # B = b * A + c * A @ A # Referencing X twice causes pytorch to make a defensive copy, # resulting in a cudaMemcpyAsync in baddbmm. # For large matrices (i.e., the mlp weights), it's faster to split # the operation into two kernels to avoid this. if split_baddbmm: BX_matmul(B, X, out=C) # C = B @ X C.add_(X, alpha=a) # C = C + a*X (in-place, X only read) else: aX_plus_BX(X, B, X, beta=a, out=C) # C = a * X + B @ X X, C = C, X # Swap references to avoid unnecessary copies if G.size(-2) > G.size(-1): X = X.mT return X # ----------------------------------------------------------------------------- # Combined NorMuon + Adam Optimizer @dataclass class ParamConfig: """Per-parameter configuration for NorMuonAndAdam optimizer.""" label: str optim: str # "adam" or "normuon" comms: str # "none", "replicated", or "sharded" adam_betas: tuple[float, float] | None lr_mul: float wd_mul: float lr: float initial_lr: float weight_decay: float # Adam-specific eps: float | None = None # NorMuon-specific reshape: tuple | None = None chunk_size: int | None = None momentum: float | None = None beta2: float | None = None per_matrix_lr_mul: list[float] | None = None class NorMuonAndAdam: """ Combined optimizer that handles both NorMuon (for projection matrices) and Adam (for embeddings/scalars/gate weights). Muon - MomentUm Orthogonalized by Newton-schulz https://kellerjordan.github.io/posts/muon/ Muon internally runs standard SGD-momentum, and then performs an orthogonalization post- processing step, in which each 2D parameter's update is replaced with the nearest orthogonal matrix. To efficiently orthogonalize each update, Muon uses a Newton-Schulz iteration (replaced here with Polar Express), which has the advantage that it can be stably run in bfloat16 on the GPU. Muon is applied only to the projection matrices in the attention and MLP layers, and is not recommended for embeddings, scalars, or individual weight vectors (e.g., bias terms or gate weights). Differences from standard Muon: - Newton-Shulz is replaced with Polar Express for the orthogonalization step - NorMuon adds a low-rank variance estimator similar to Adafactor. https://arxiv.org/pdf/2510.05491 - Cautious weight decay, a gated version of decoupled weight decay - Mantissa tracking for precision Adam (for embeddings/scalars/gates): - Standard Adam with bias correction - Cautious weight decay Configuration: Unlike torch.optim.Optimizer, this class uses per-parameter configs from a `param_table` dict and does not include parameter "groups". All parameters require a .label attribute, and a corresponding entry in the param_table to specify their hyperparameters (lr_mul, wd_mul, adam_betas, etc.). Communication and ordering: Gradient communication is explicitly scheduled rather than hook-driven. Reductions are launched in `scatter_order`, while update math and final gathers are executed in `work_order`. These orders are independent and must each contain every parameter label exactly once. Two communication modes are supported per parameter: - 'replicated': Gradients are all-reduced and each rank computes the full update. - 'sharded': Gradients are reduce-scattered, each rank updates its shard, and results are all-gathered. Adam parameters may be freely sharded. NorMuon operates on full matrices; sharding is supported by grouping matrices into parameter banks. NorMuon parameters must have a `.reshape` attribute that reshapes the bank so that the leading dimension is divisible by world_size. # Contributors include @YouJiacheng, @KonstantinWilleke, @alexrgilbert, @adricarda, # @tuttyfrutyee, @vdlad, @ryanyang0, @vagrawal, @varunneal, @chrisjmccormick """ def __init__(self, named_params, param_table: dict, scatter_order: list, work_order: list, adam_defaults: dict, normuon_defaults: dict): self.world_size = dist.get_world_size() if dist.is_initialized() else 1 # Store defaults for each optimizer type self.adam_defaults = adam_defaults self.normuon_defaults = normuon_defaults self.param_table = param_table self.scatter_order = scatter_order self.work_order = work_order # Collect params by label and build config self.param_cfgs: dict[nn.Parameter, ParamConfig] = {} self.param_states: dict[nn.Parameter, dict] = {} self._param_by_label: dict[str, nn.Parameter] = {} for name, param in named_params: label = getattr(param, "label", None) assert label is not None and label in param_table # all params must have valid label assert label not in self._param_by_label # exactly one param per label self._param_by_label[label] = param self._build_param_cfg(param, label) # Assert scatter_order and work_order match present labels exactly present = set(self._param_by_label.keys()) assert set(scatter_order) == present and set(work_order) == present # Handle world_size=1: overwrite comms to "none" if self.world_size == 1: for p_cfg in self.param_cfgs.values(): p_cfg.comms = "none" # Initialize state for all params self._init_state() # 0-D CPU tensors to avoid recompilation self._step_size_t = torch.tensor(0.0, dtype=torch.float32, device="cpu") self._eff_wd_t = torch.tensor(0.0, dtype=torch.float32, device="cpu") self._eff_lr_t = torch.tensor(0.0, dtype=torch.float32, device="cpu") # Track async operations self._reduce_futures: dict[nn.Parameter, tuple] = {} # Embed/lm_head tying state self.split_embed = False self._lm_head_param = self._param_by_label.get("lm_head") self._embed_param = self._param_by_label.get("embed") def _build_param_cfg(self, param: nn.Parameter, label: str): """Build config for a single parameter from param_table.""" table_entry = self.param_table[label] optim = table_entry["optim"] comms = table_entry["comms"] adam_betas = table_entry.get("adam_betas") lr_mul = table_entry.get("lr_mul", 1.0) wd_mul = table_entry.get("wd_mul", 1.0) if optim == "adam": chunk_size = param.shape[0] // self.world_size if comms == "sharded" else None p_cfg = ParamConfig( label=label, optim=optim, comms=comms, adam_betas=tuple(adam_betas) if adam_betas else None, lr_mul=lr_mul, wd_mul=wd_mul, lr=self.adam_defaults["lr"], initial_lr=self.adam_defaults["lr"], weight_decay=self.adam_defaults["weight_decay"], eps=self.adam_defaults["eps"], chunk_size=chunk_size, ) elif optim == "normuon": reshape = getattr(param, "reshape", None) if reshape is None: raise ValueError(f"NorMuon param {label} must have .reshape attribute") if reshape[0] % self.world_size != 0: raise ValueError(f"reshape[0]={reshape[0]} must be divisible by world_size") chunk_size = reshape[0] // self.world_size chunk_shape = (chunk_size, *reshape[1:]) # Shape-based LR multiplier for NorMuon shape_mult = max(1.0, chunk_shape[-2] / chunk_shape[-1]) ** 0.5 if len(chunk_shape) >= 2 else 1.0 lr_mul = shape_mult * lr_mul # Per-matrix LR multipliers for MLP c_proj (2x LR on odd indices) per_matrix_lr_mul = None if label == "mlp": rank = dist.get_rank() if dist.is_initialized() else 0 start_idx = rank * chunk_size per_matrix_lr_mul = [] for i in range(chunk_size): global_idx = start_idx + i is_c_proj = (global_idx % 2 == 1) per_matrix_lr_mul.append(2.0 if is_c_proj else 1.0) p_cfg = ParamConfig( label=label, optim=optim, comms=comms, adam_betas=tuple(adam_betas) if adam_betas else None, lr_mul=lr_mul, wd_mul=wd_mul, lr=self.normuon_defaults["lr"], initial_lr=self.normuon_defaults["lr"], weight_decay=self.normuon_defaults["weight_decay"], reshape=reshape, chunk_size=chunk_size, momentum=self.normuon_defaults["momentum"], beta2=self.normuon_defaults["beta2"], per_matrix_lr_mul=per_matrix_lr_mul, ) else: raise ValueError(f"Unknown optim type: {optim}") self.param_cfgs[param] = p_cfg def _init_state(self): """Initialize optimizer state for all parameters.""" for param, p_cfg in self.param_cfgs.items(): if p_cfg.optim == "adam": # Sharded params use chunk state, replicated use full state if p_cfg.comms == "sharded": chunk = param[:p_cfg.chunk_size] else: chunk = param exp_avg = torch.zeros_like(chunk, dtype=torch.float32, device=param.device) self.param_states[param] = dict(step=0, exp_avg=exp_avg, exp_avg_sq=torch.zeros_like(exp_avg)) elif p_cfg.optim == "normuon": chunk_shape = (p_cfg.chunk_size, *p_cfg.reshape[1:]) # Momentum buffer (FP32 for precision) momentum_buffer = torch.zeros( chunk_shape, dtype=torch.float32, device=param.device ) # Second momentum buffer - reduced along one dimension if chunk_shape[-2] >= chunk_shape[-1]: second_mom_shape = (*chunk_shape[:-1], 1) else: second_mom_shape = (*chunk_shape[:-2], 1, chunk_shape[-1]) second_momentum_buffer = torch.zeros( second_mom_shape, dtype=torch.float32, device=param.device ) # Mantissa buffer for precision tracking mantissa = torch.zeros( chunk_shape, dtype=torch.uint16, device=param.device ) self.param_states[param] = dict( momentum_buffer=momentum_buffer, second_momentum_buffer=second_momentum_buffer, mantissa=mantissa, ) # ----------------------------------- # Reduce/Gather operations def _launch_reduce(self, param: nn.Parameter, grad: Tensor): """Launch async reduce for a parameter based on its comms policy.""" p_cfg = self.param_cfgs[param] if p_cfg.comms == "none": if p_cfg.optim == "normuon": # NorMuon needs reshaped gradient even without communication grad = grad.view(p_cfg.reshape) self._reduce_futures[param] = (None, grad) elif p_cfg.comms == "replicated": future = dist.all_reduce(grad, op=dist.ReduceOp.AVG, async_op=True).get_future() self._reduce_futures[param] = (future, grad) elif p_cfg.comms == "sharded": if p_cfg.optim == "normuon": # NorMuon: reshape before reduce_scatter grad_reshaped = grad.view(p_cfg.reshape) grad_chunk = torch.empty( (p_cfg.chunk_size, *grad_reshaped.shape[1:]), dtype=grad.dtype, device=grad.device ) future = dist.reduce_scatter_tensor( grad_chunk, grad_reshaped.contiguous(), op=dist.ReduceOp.AVG, async_op=True ).get_future() self._reduce_futures[param] = (future, grad_chunk) else: # Adam: simple reduce_scatter grad_chunk = torch.empty_like(grad[:p_cfg.chunk_size]) future = dist.reduce_scatter_tensor( grad_chunk, grad, op=dist.ReduceOp.AVG, async_op=True ).get_future() self._reduce_futures[param] = (future, grad_chunk) def _launch_gather(self, param: nn.Parameter, p_slice: Tensor) -> "torch.futures.Future": """Launch async all_gather for a sharded parameter.""" p_cfg = self.param_cfgs[param] if p_cfg.optim == "normuon": full_param = param.data.view(p_cfg.reshape) assert full_param.is_contiguous() return dist.all_gather_into_tensor( full_param, p_slice.contiguous(), async_op=True ).get_future() else: return dist.all_gather_into_tensor( param, p_slice.contiguous(), async_op=True ).get_future() # ----------------------------------- # State management def reset(self): """Reset NorMuon momentum buffers and split_embed state (called on training reset).""" self.split_embed = False for param, p_cfg in self.param_cfgs.items(): if p_cfg.optim == "normuon": p_state = self.param_states[param] p_state["momentum_buffer"].zero_() p_state["mantissa"].zero_() p_state["second_momentum_buffer"].zero_() def copy_lm_state_to_embed(self): """ Copy the optimizer state from the lm_head to the embed at the untie point. This requires an all-gather + reshard because of different sharding: - lm_head (768, 50304) is sharded to (96, 50304) per rank (along model_dim) - embed (50304, 768) is sharded to (6288, 768) per rank (along vocab_size) We all-gather the lm_head momentum, transpose it, then each rank takes their embed shard to get the correct momentum state. """ lm_head = self._lm_head_param embed = self._embed_param lm_state = self.param_states[lm_head] embed_state = self.param_states[embed] lm_cfg = self.param_cfgs[lm_head] embed_cfg = self.param_cfgs[embed] embed_state['step'] = lm_state['step'] # Preserve step count for bias correction # Copy optimizer state with all-gather + transpose + reshard if self.world_size > 1: rank = dist.get_rank() lm_chunk_size = lm_cfg.chunk_size # 96 embed_chunk_size = embed_cfg.chunk_size # 6288 # All-gather lm_head momentum to get full (768, 50304) tensor for key in ["exp_avg", "exp_avg_sq"]: lm_chunk = lm_state[key] # (96, 50304) full_lm = torch.empty(lm_head.shape[0], lm_head.shape[1], dtype=lm_chunk.dtype, device=lm_chunk.device) dist.all_gather_into_tensor(full_lm, lm_chunk.contiguous()) embed_state[key].copy_(full_lm.T[rank * embed_chunk_size:(rank + 1) * embed_chunk_size]) else: # Single GPU: simple transpose for key in ["exp_avg", "exp_avg_sq"]: embed_state[key].copy_(lm_state[key].T) # Mark as split self.split_embed = True def state_dict(self): """Return the optimizer state as a dict.""" return { "param_states": {id(p): s for p, s in self.param_states.items()}, "param_cfgs": {id(p): s for p, s in self.param_cfgs.items()}, } def load_state_dict(self, state_dict): """Load optimizer state from a dict.""" # Build id->param mapping id_to_param = {id(p): p for p in self.param_cfgs.keys()} # Load state, preserving dtypes for param_id, saved_p_state in state_dict["param_states"].items(): if param_id in id_to_param: param = id_to_param[param_id] p_state = self.param_states[param] for k, v in saved_p_state.items(): if isinstance(v, torch.Tensor) and k in p_state: target_dtype = p_state[k].dtype p_state[k] = v.to(dtype=target_dtype, device=p_state[k].device) else: p_state[k] = v # ----------------------------------- # Unified optimizer step with explicit ordering @torch.no_grad() def step(self, do_adam: bool = True): """ Combined optimizer step with explicit ordering. Args: do_adam: If True, update Adam params. NorMuon params always updated. Flow: 1. Scatter phase: Launch reduces in scatter_order 2. Work phase: Process updates in work_order - Wait for reduce, compute update, launch gather 3. Finalize phase: Wait for gathers While the embeddings are tied: - Comms and update math are only done on lm_head. - We add embed.grad.T into lm_head.grad before comms. - After lm_head gather, we copy lm_head.data.T --> embed.data """ rank = dist.get_rank() if dist.is_initialized() else 0 lm_param, embed_param = self._lm_head_param, self._embed_param # ===== Phase 1: Launch reduces in scatter_order ===== for label in self.scatter_order: param = self._param_by_label[label] p_cfg = self.param_cfgs[param] if p_cfg.optim == "adam" and not do_adam: continue if param.grad is None: continue # lm_head when tied: aggregate embed.grad.T (transposed shapes) if label == "lm_head" and do_adam and not self.split_embed: if embed_param is not None and embed_param.grad is not None: param.grad.add_(embed_param.grad.T) # Skip embed when tied (copied from lm_head after gather) if label == "embed" and not self.split_embed: continue self._launch_reduce(param, param.grad) # ===== Phase 2: Process updates in work_order ===== gather_futures = [] lm_head_gather_future = None for label in self.work_order: param = self._param_by_label[label] if param not in self._reduce_futures: continue p_cfg = self.param_cfgs[param] if p_cfg.optim == "adam" and not do_adam: continue # Wait for reduce future, grad_chunk = self._reduce_futures[param] if future is not None: future.wait() # Apply update based on optim type if p_cfg.optim == "adam": p_slice = self._adam_update(param, grad_chunk, p_cfg, rank) else: p_slice = self._normuon_update(param, grad_chunk, p_cfg, rank) # Launch gather for sharded params if p_cfg.comms == "sharded" and self.world_size > 1: gather_fut = self._launch_gather(param, p_slice) if label == "lm_head": lm_head_gather_future = gather_fut else: gather_futures.append(gather_fut) # ===== Phase 3: Wait for gathers, sync embed if tied ===== # Wait for lm_head gather first so we can copy to embed while other gathers complete if lm_head_gather_future is not None: lm_head_gather_future.wait() # When tied: copy lm_head.T to embed if do_adam and not self.split_embed and embed_param is not None and lm_param is not None: embed_param.data.copy_(lm_param.data.T) # Wait for remaining gathers for fut in gather_futures: fut.wait() self._reduce_futures.clear() # Clear grads for updated params for param, p_cfg in self.param_cfgs.items(): if p_cfg.optim == "adam" and not do_adam: continue # Don't clear Adam grads on even steps param.grad = None # ----------------------------------- # Adam update def _adam_update(self, param: nn.Parameter, grad_chunk: Tensor, p_cfg: ParamConfig, rank: int) -> Tensor: """Apply Adam update to a parameter. Returns the updated p_slice.""" beta1, beta2 = p_cfg.adam_betas lr = p_cfg.lr * p_cfg.lr_mul # Get parameter slice if p_cfg.comms == "sharded": p_slice = param[rank * p_cfg.chunk_size:(rank + 1) * p_cfg.chunk_size] else: p_slice = param p_state = self.param_states[param] p_state["step"] += 1 t = p_state["step"] bias1, bias2 = 1 - beta1 ** t, 1 - beta2 ** t self._step_size_t.fill_(lr * (bias2 ** 0.5 / bias1)) self._eff_wd_t.fill_(lr * lr * p_cfg.weight_decay * p_cfg.wd_mul) NorMuonAndAdam._adam_update_step( p_slice, grad_chunk, p_state["exp_avg"], p_state["exp_avg_sq"], beta1, beta2, p_cfg.eps, self._step_size_t, self._eff_wd_t ) return p_slice @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _adam_update_step(p_slice, g_slice, exp_avg, exp_avg_sq, beta1, beta2, eps, step_size_t, eff_wd_t): """Compiled Adam update step.""" exp_avg.mul_(beta1).add_(g_slice, alpha=1 - beta1) exp_avg_sq.mul_(beta2).addcmul_(g_slice, g_slice, value=1 - beta2) update = exp_avg.div(exp_avg_sq.sqrt().add_(eps)).mul_(step_size_t) # Cautious weight decay mask = (update * p_slice) > 0 update.addcmul_(p_slice, mask, value=eff_wd_t) p_slice.add_(other=update, alpha=-1.0) # ----------------------------------- # NorMuon update def _normuon_update(self, param: nn.Parameter, grad_chunk: Tensor, p_cfg: ParamConfig, rank: int) -> Tensor: """Apply NorMuon update to a parameter. Returns the updated p_slice.""" chunk_shape = grad_chunk.shape p_state = self.param_states[param] grad_chunk = grad_chunk.float() # FP32 for momentum # Momentum update momentum_buffer = p_state["momentum_buffer"] momentum_buffer.lerp_(grad_chunk, 1 - p_cfg.momentum) updated_grads = grad_chunk.lerp_(momentum_buffer, p_cfg.momentum) self._eff_lr_t.fill_(p_cfg.lr_mul * p_cfg.lr) self._eff_wd_t.fill_(p_cfg.wd_mul * p_cfg.weight_decay * p_cfg.lr) # Polar Express orthogonalization is_large_matrix = chunk_shape[-2] > 1024 v_chunk = polar_express(updated_grads, split_baddbmm=is_large_matrix) # Variance reduction red_dim = -1 if chunk_shape[-2] >= chunk_shape[-1] else -2 v_chunk = NorMuonAndAdam._apply_normuon_variance_reduction( v_chunk, p_state["second_momentum_buffer"], p_cfg.beta2, red_dim ) # Update parameter, in place, with cautious weight decay param_view = param.data.view(p_cfg.reshape) p_slice = param_view[rank * p_cfg.chunk_size:(rank + 1) * p_cfg.chunk_size] # MLP has per-matrix LR multipliers (c_proj gets 2x LR) if p_cfg.per_matrix_lr_mul is not None: for mat_idx in range(p_cfg.chunk_size): self._eff_lr_t.fill_(p_cfg.lr_mul * p_cfg.per_matrix_lr_mul[mat_idx] * p_cfg.lr) self._eff_wd_t.fill_(p_cfg.wd_mul * p_cfg.weight_decay * p_cfg.lr) NorMuonAndAdam._cautious_wd_and_update_inplace( p_slice[mat_idx].view(torch.uint16), p_state["mantissa"][mat_idx], v_chunk[mat_idx], self._eff_wd_t, self._eff_lr_t ) else: NorMuonAndAdam._cautious_wd_and_update_inplace( p_slice.view(torch.uint16), p_state["mantissa"], v_chunk, self._eff_wd_t, self._eff_lr_t ) return p_slice @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _cautious_wd_and_update_inplace(p, mantissa, grad, wd_tensor, lr_tensor): """ Cautious weight decay + parameter update. wd_tensor and lr_tensor are 0-D CPU tensors. Mantissa is tracked to enable higher precision updates on bfloat16 parameters. bfloat16 format: 1 sign bit + 8 exponent bits + 7 mantissa bits = 16 bits total float32 format: 1 sign bit + 8 exponent bits + 23 mantissa bits = 32 bits total """ assert p.dtype == mantissa.dtype == torch.uint16 grad = grad.float() wd_factor = wd_tensor.to(torch.float32) lr_factor = lr_tensor.to(torch.float32) p_precise_raw = (p.to(torch.uint32) << 16) | mantissa.to(torch.uint32) p_precise = p_precise_raw.view(torch.float32) mask = (grad * p_precise) >= 0 p_precise.copy_(p_precise - (p_precise * mask * wd_factor * lr_factor) - (grad * lr_factor)) p.copy_((p_precise_raw >> 16).to(torch.uint16)) mantissa.copy_(p_precise_raw.to(torch.uint16)) @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _apply_normuon_variance_reduction(v_chunk, second_momentum_buffer, beta2, red_dim): """NorMuon variance reduction. Algebraically fuses the normalization steps to minimize memory ops.""" v_mean = v_chunk.float().square().mean(dim=red_dim, keepdim=True) red_dim_size = v_chunk.size(red_dim) v_norm_sq = v_mean.sum(dim=(-2, -1), keepdim=True).mul_(red_dim_size) v_norm = v_norm_sq.sqrt_() second_momentum_buffer.lerp_(v_mean.to(dtype=second_momentum_buffer.dtype), 1 - beta2) step_size = second_momentum_buffer.clamp_min(1e-10).rsqrt_() scaled_sq_sum = (v_mean * red_dim_size) * step_size.float().square() v_norm_new = scaled_sq_sum.sum(dim=(-2, -1), keepdim=True).sqrt_() final_scale = step_size * (v_norm / v_norm_new.clamp_min_(1e-10)) return v_chunk.mul_(final_scale.type_as(v_chunk)) # ----------------------------------------------------------------------------- # PyTorch nn.Module definitions for the model def norm(x: Tensor): return F.rms_norm(x, (x.size(-1),)) class CastedLinearT(nn.Module): """ Linear layer with transposed weight storage (in_features, out_features) which addresses the slow kernel that was used for gradient accumulation. @chrisjmccormick """ def __init__(self, in_features: int, out_features: int, use_fp8=False, x_s=1.0, w_s=1.0, grad_s=1.0): super().__init__() self.in_features = in_features self.out_features = out_features self.use_fp8 = use_fp8 self.x_s = x_s self.w_s = w_s self.grad_s = grad_s self.weight = nn.Parameter(torch.empty(in_features, out_features, dtype=torch.bfloat16)) self.reset_parameters() def reset_parameters(self) -> None: with torch.no_grad(): nn.init.zeros_(self.weight) # @Grad62304977 and others def forward(self, x: Tensor): if self.use_fp8 and self.training: _x = x.flatten(0, -2) out = torch.ops.nanogpt.mm_t(_x, self.weight, x_s=self.x_s, w_s=self.w_s, grad_s=self.grad_s)[0] return out.reshape(*x.shape[:-1], -1) else: return x @ self.weight.type_as(x) # ----------------------------------------------------------------------------- # PyTorch nn.Module definitions for the model class Yarn(nn.Module): def __init__(self, head_dim, max_seq_len): super().__init__() self.head_dim = head_dim self.max_seq_len = max_seq_len self.reset() def rotary(self, x_BTHD): assert self.factor1.size(0) >= x_BTHD.size(-3) factor1, factor2 = ( self.factor1[None, : x_BTHD.size(-3), None, :], self.factor2[None, : x_BTHD.size(-3), None, :], ) x_flip = x_BTHD.view(*x_BTHD.shape[:-1], x_BTHD.shape[-1] // 2, 2).flip(-1).view(x_BTHD.shape) return factor1 * x_BTHD + factor2 * x_flip def reset(self): angular_freq = (1 / 1024) ** torch.linspace(0, 1, steps=self.head_dim//4, dtype=torch.float32, device=device) angular_freq = angular_freq.repeat_interleave(2) # half-truncate RoPE by @YouJiacheng (w/ base freq tuning) angular_freq = torch.cat([angular_freq, angular_freq.new_zeros(self.head_dim//2)]) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=device) theta = torch.outer(t, angular_freq) self.factor1 = nn.Buffer( theta.cos().to(torch.bfloat16), persistent=False ) self.factor2 = nn.Buffer( theta.sin().to(torch.bfloat16), persistent=False ) self.factor2[..., 1::2] *= -1 self.angular_freq = angular_freq # start with 0.1, inspired by 0.12 from @leloykun and learnable scalars used by @brendanh0gan https://x.com/hi_tysam/status/1879693583898591283 self.attn_scale = 0.1 def apply(self, old_window: int, new_window: int, alpha: int=1, beta: int=32): rotations = args.block_size * old_window * self.angular_freq / (2 * torch.pi) scaling_factor = old_window / new_window interpolation_weight = torch.clamp((rotations - alpha) / (beta - alpha), 0, 1) self.angular_freq *= scaling_factor + interpolation_weight * (1 - scaling_factor) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=self.angular_freq.device) theta = torch.outer(t, self.angular_freq) self.factor1.copy_(theta.cos()) self.factor2.copy_(theta.sin()) self.factor2[..., 1::2] *= -1 self.attn_scale *= 0.2 * math.log(new_window / old_window) + 1 class YarnPairedHead(nn.Module): def __init__(self, head_dim, max_seq_len): super().__init__() self.head_dim = head_dim self.max_seq_len = max_seq_len self.reset() def rotary(self, x_BTHD): assert self.factor1.size(0) >= x_BTHD.size(-3) factor1, factor2 = ( self.factor1[None, : x_BTHD.size(-3), None, :], self.factor2[None, : x_BTHD.size(-3), None, :], ) x_flip = x_BTHD.view(*x_BTHD.shape[:-1], x_BTHD.shape[-1] // 2, 2).flip(-1).view(x_BTHD.shape) return factor1 * x_BTHD + factor2 * x_flip def reset(self): angular_freq = (1 / 1024) ** torch.linspace(0, 1, steps=self.head_dim//4, dtype=torch.float32, device=device) angular_freq = angular_freq.repeat_interleave(2) angular_freq = torch.cat([angular_freq, angular_freq.new_zeros(self.head_dim//2)]) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=device) t_even = 2 * t t_odd = 2 * t + 1 theta1 = torch.outer(t_even, angular_freq) theta2 = torch.outer(t_odd, angular_freq) self.factor1 = nn.Buffer( torch.cat((theta1.cos(),theta2.cos()), dim=-1).to(torch.bfloat16), persistent=False ) self.factor2 = nn.Buffer( torch.cat((theta1.sin(),theta2.sin()), dim=-1).to(torch.bfloat16), persistent=False ) self.factor2[..., 1::2] *= -1 self.angular_freq = angular_freq # start with 0.1, inspired by 0.12 from @leloykun and learnable scalars used by @brendanh0gan https://x.com/hi_tysam/status/1879693583898591283 self.attn_scale = 0.1 def apply(self, old_window: int, new_window: int, alpha: int=1, beta: int=32): rotations = args.block_size * old_window * self.angular_freq / (2 * torch.pi) scaling_factor = old_window / new_window interpolation_weight = torch.clamp((rotations - alpha) / (beta - alpha), 0, 1) self.angular_freq *= scaling_factor + interpolation_weight * (1 - scaling_factor) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=self.angular_freq.device) t_even = 2 * t t_odd = 2 * t + 1 theta1 = torch.outer(t_even, self.angular_freq) theta2 = torch.outer(t_odd, self.angular_freq) self.factor1.copy_(torch.cat((theta1.cos(),theta2.cos()), dim=-1)) self.factor2.copy_( torch.cat((theta1.sin(),theta2.sin()), dim=-1)) self.factor2[..., 1::2] *= -1 self.attn_scale *= 0.2 * math.log(new_window / old_window) + 1 @dataclass class AttnArgs: ve: torch.Tensor sa_lambdas: torch.Tensor seqlens: torch.Tensor bm_size: int yarn: Yarn key_offset: bool attn_gate_w: torch.Tensor ve_gate_w: torch.Tensor flash_attn_interface = get_kernel('varunneal/flash-attention-3').flash_attn_interface class CausalSelfAttention(nn.Module): def __init__(self, dim: int, head_dim: int, num_heads: int): super().__init__() self.num_heads = num_heads self.head_dim = head_dim self.dim = dim self.hdim = num_heads * head_dim assert self.hdim == self.dim, "num_heads * head_dim must equal model_dim" # Weights are stored in parameter banks and passed via forward() def forward(self, x: Tensor, attn_args: AttnArgs, qkvo_w: Tensor): B, T = x.size(0), x.size(1) # batch size, sequence length assert B == 1, "varlen sequences requires B == 1" assert T % 16 == 0 # unpack attention args yarn = attn_args.yarn ve, sa_lambdas, key_offset = attn_args.ve, attn_args.sa_lambdas, attn_args.key_offset seqlens, bm_size = attn_args.seqlens, attn_args.bm_size # sparse gated attention to enable context based no-op by @classiclarryd # only include gates on layers with value embeds used on forward pass attn_gate_w, ve_gate_w = attn_args.attn_gate_w, attn_args.ve_gate_w q, k, v = F.linear(x, sa_lambdas[0] * qkvo_w[:self.dim * 3].type_as(x)).view(B, T, 3 * self.num_heads, self.head_dim).chunk(3, dim=-2) q, k = norm(q), norm(k) # QK norm @Grad62304977 q, k = yarn.rotary(q), yarn.rotary(k) if key_offset: # shift keys forward for the stationary head dims. Enables 1-layer induction. k[:, 1:, :, self.head_dim // 2:] = k[:, :-1, :, self.head_dim // 2:] if ve is not None: ve_gate_out = 2 * torch.sigmoid(F.linear(x[..., :12], ve_gate_w)).view(B, T, self.num_heads, 1) v = v + ve_gate_out * ve.view_as(v) # @ KoszarskyB & @Grad62304977 max_len = args.train_max_seq_len if self.training else (args.val_batch_size // (grad_accum_steps * world_size)) # use flash_attn over flex_attn @varunneal. flash_attn_varlen suggested by @YouJiacheng y = flash_attn_interface.flash_attn_varlen_func(q[0], k[0], v[0], cu_seqlens_q=seqlens, cu_seqlens_k=seqlens, max_seqlen_q=max_len, max_seqlen_k=max_len, causal=True, softmax_scale=yarn.attn_scale, window_size=(bm_size, 0)) y = y.view(B, T, self.num_heads, self.head_dim) y = y * torch.sigmoid(F.linear(x[..., :12], attn_gate_w)).view(B, T, self.num_heads, 1) y = y.contiguous().view(B, T, self.num_heads * self.head_dim) # re-assemble all head outputs side by side y = F.linear(y, sa_lambdas[1] * qkvo_w[self.dim * 3:].type_as(y)) # sa_lambdas[1] pre-multiplied to O @shenberg return y class PairedHeadCausalSelfAttention(nn.Module): """ Pairs up attention heads such that queries from head 1 can attend to keys in head 2, and vice-versa. Implemented by interleaving the k, q, and v for pairs of heads to form twice as long sequences EG [k1_h1, k2_h1, k3_h1], [k1_h2, k2_h2, k3_h2] -> [k1_h1, k1_h2, k2_h1, k2_h2, k3_h1, k3_h2], repeat for q and v """ def __init__(self, dim: int, head_dim: int, num_heads: int): super().__init__() self.num_heads = num_heads self.head_dim = head_dim self.dim = dim self.hdim = num_heads * head_dim assert self.hdim == self.dim, "num_heads * head_dim must equal model_dim" # Weights are stored in parameter banks and passed via forward() def forward(self, x: Tensor, attn_args: AttnArgs, qkvo_w: Tensor): B, T = x.size(0), x.size(1) # batch size, sequence length assert B == 1, "varlen sequences requires B == 1" assert T % 16 == 0 # unpack attention args yarn = attn_args.yarn ve, sa_lambdas = attn_args.ve, attn_args.sa_lambdas seqlens, bm_size = attn_args.seqlens, attn_args.bm_size attn_gate_w, ve_gate_w = attn_args.attn_gate_w, attn_args.ve_gate_w q, k, v = F.linear(x, sa_lambdas[0] * qkvo_w[:self.dim * 3].type_as(x)).view(B, T, 3 * self.num_heads, self.head_dim).chunk(3, dim=-2) q, k = norm(q), norm(k) # delay q,k reshape until rotary makes data contiguous, to enable view (non-copy) q = q.view(B, T, self.num_heads // 2, self.head_dim * 2) k = k.view(B, T, self.num_heads // 2, self.head_dim * 2) v = v.reshape(B, T*2, self.num_heads//2, self.head_dim) q, k = yarn.rotary(q), yarn.rotary(k) q = q.view(B, T*2, self.num_heads//2, self.head_dim) k = k.view(B, T*2, self.num_heads//2, self.head_dim) if ve is not None: ve_gate_out = 2 * torch.sigmoid(F.linear(x[..., :12], ve_gate_w)).view(B, T*2, self.num_heads//2, 1) v = v + ve_gate_out * ve.view_as(v) max_len = args.train_max_seq_len if self.training else (args.val_batch_size // (grad_accum_steps * world_size)) # paired head correction seqlens = 2 * seqlens max_len = 2 * max_len y = flash_attn_interface.flash_attn_varlen_func(q[0], k[0], v[0], cu_seqlens_q=seqlens, cu_seqlens_k=seqlens, max_seqlen_q=max_len, max_seqlen_k=max_len, causal=True, softmax_scale=yarn.attn_scale, window_size=(bm_size, 0)) y = y.view(B, T, self.num_heads, self.head_dim) y = y * torch.sigmoid(F.linear(x[..., :12], attn_gate_w)).view(B, T, self.num_heads, 1) y = y.contiguous().view(B, T, self.num_heads * self.head_dim) y = F.linear(y, sa_lambdas[1] * qkvo_w[self.dim * 3:].type_as(y)) return y class MLP(nn.Module): def __init__(self): super().__init__() # Weights are stored in parameter banks and passed via forward() def forward(self, x: Tensor, c_fc: Tensor, c_proj: Tensor): # relu(x)^2: # https://arxiv.org/abs/2109.08668v2; ~1-2% better than GELU; suggested by @SKYLINEZ007 and @Grad62304977 # Fused triton kernel for relu(x @ W1.T)^2 @ W2.T return FusedLinearReLUSquareFunction.apply(x, c_fc, c_proj) class Block(nn.Module): def __init__(self, dim: int, head_dim: int, num_heads: int, has_attn: bool, has_mlp: bool, use_paired_head: bool): super().__init__() # skip attention of blocks.6 (the 7th layer) by @YouJiacheng if has_attn: if use_paired_head: self.attn = PairedHeadCausalSelfAttention(dim, head_dim, num_heads) else: self.attn = CausalSelfAttention(dim, head_dim, num_heads) else: self.attn = None # skip MLP blocks for first MLP layer by @EmelyanenkoK self.mlp = MLP() if has_mlp else None def forward(self, x: Tensor, attn_args: AttnArgs, qkvo_w: Tensor = None, c_fc: Tensor = None, c_proj: Tensor = None): if self.attn is not None: x = x + self.attn(norm(x), attn_args, qkvo_w) if self.mlp is not None: x = x + self.mlp(norm(x), c_fc, c_proj) return x # ----------------------------------------------------------------------------- # The main model def next_multiple_of_n(v: float | int, *, n: int): return next(x for x in range(n, int(v) + 1 + n, n) if x >= v) @dataclass class ForwardScheduleConfig: mtp_weights: torch.Tensor ws_short: int ws_long: int class GPT(nn.Module): def __init__(self, vocab_size: int, num_layers: int, num_heads: int, head_dim: int, model_dim: int, max_seq_len: int): super().__init__() self.num_layers = num_layers vocab_size = next_multiple_of_n(vocab_size, n=128) self.smear_gate = nn.Linear(12, 1, bias=False) nn.init.zeros_(self.smear_gate.weight) self.smear_gate.weight.label = 'smear_gate' self.skip_gate = nn.Linear(12, 1, bias=False) nn.init.zeros_(self.skip_gate.weight) self.skip_gate.weight.label = 'skip_gate' # token value embeddings by @KoszarskyB - inspired by @Grad62304977's value residual implementation following https://arxiv.org/abs/2410.17897 # value embedding code simplification inspired by @ragulpr https://github.com/KellerJordan/modded-nanogpt/pull/78 self.value_embeds = nn.ModuleList([nn.Embedding(vocab_size, model_dim) for _ in range(5)]) for embed in self.value_embeds: nn.init.zeros_(embed.weight) for i, ve in enumerate(self.value_embeds): ve.weight.label = f've{i}' # ve0, ve1, ve2, ve3, ve4 # parameter banks for attention and value embedding gate weights self.attn_gate_bank = nn.Parameter(torch.zeros(10, num_heads, 12)) # 10 layers self.attn_gate_bank.label = 'attn_gate_bank' self.ve_gate_bank = nn.Parameter(torch.zeros(5, num_heads, 12)) # 5 unique gates self.ve_gate_bank.label = 've_gate_bank' # ----------------------------------- # Parameter banks for sharded optimization, by @chrisjmccormick # Identify which layers have attention/MLP # Attention is skipped in layer 6 by @YouJiacheng self.attn_layer_indices = [i for i in range(num_layers) if i != 6] # All layers have MLP (At 11 layers--dropped first layer @EmelyanenkoK) self.mlp_layer_indices = list(range(num_layers)) hdim = num_heads * head_dim mlp_hdim = 4 * model_dim # Create index mappings: layer_idx -> bank_idx self.layer_to_attn_idx = {layer_idx: bank_idx for bank_idx, layer_idx in enumerate(self.attn_layer_indices)} self.layer_to_mlp_idx = {layer_idx: bank_idx for bank_idx, layer_idx in enumerate(self.mlp_layer_indices)} # Attention bank: stores QKVO weights for all attention layers # merged QKVO weights: suggested by many, implemented by @fernbear.bsky.social, and further improved by @YouJiacheng # https://x.com/hi_tysam/status/1879699187107033311 # Simplified layout by @chrisjmccormick # Shape: (num_attn_layers, 4*model_dim, hdim) = (10, 3072, 768) # Reshape for sharding: (40, 768, 768) for even distribution across 8 GPUs self.attn_bank = nn.Parameter(torch.empty(len(self.attn_layer_indices), 4 * model_dim, hdim)) self.attn_bank.label = 'attn' self.attn_bank.reshape = (len(self.attn_layer_indices) * 4, hdim, hdim) # (40, 768, 768) # MLP bank: stores c_fc and c_proj for all MLP layers # Shape: (num_mlp_layers + padding, 2, mlp_hdim, model_dim) = (12, 2, 3072, 768) # We add 1 padding layer (index 11) to get 12*2=24 matrices for even distribution across 8 GPUs # Reshape for sharding: (24, 3072, 768) num_mlp_with_padding = len(self.mlp_layer_indices) + 1 # 11 + 1 = 12 self.mlp_bank = nn.Parameter(torch.empty(num_mlp_with_padding, 2, mlp_hdim, model_dim)) self.mlp_bank.label = 'mlp' self.mlp_bank.reshape = (num_mlp_with_padding * 2, mlp_hdim, model_dim) # (24, 3072, 768) # improved init scale by @YouJiacheng # Attention uses dim^-0.5, MLP uses 0.5 * dim^-0.5 attn_std = model_dim ** -0.5 attn_bound = (3 ** 0.5) * attn_std mlp_std = 0.5 * (model_dim ** -0.5) mlp_bound = (3 ** 0.5) * mlp_std with torch.no_grad(): # Init attention bank (QKV uniform, O zero) self.attn_bank[:, :model_dim * 3, :].uniform_(-attn_bound, attn_bound) self.attn_bank[:, model_dim * 3:, :].zero_() # Init MLP bank (c_fc uniform, c_proj zero) self.mlp_bank[:, 0, :, :].uniform_(-mlp_bound, mlp_bound) # c_fc self.mlp_bank[:, 1, :, :].zero_() # c_proj - zero init suggested by @Grad62304977 # Create blocks with has_attn/has_mlp flags self.paired_head_layers = [0, 2, 5, 9] self.blocks = nn.ModuleList([ Block(model_dim, head_dim, num_heads, has_attn=(i in self.layer_to_attn_idx), has_mlp=(i in self.layer_to_mlp_idx), use_paired_head=(i in self.paired_head_layers)) for i in range(num_layers) ]) self.yarn = Yarn(head_dim, max_seq_len) self.yarn_paired_head = YarnPairedHead(head_dim, max_seq_len) # there are only 50257 unique GPT-2 tokens; we extend to nearest multiple of 128 for efficiency. # suggested to me by @Grad62304977. this originates from Karpathy's experiments. use_fp8 = not os.environ.get("DISABLE_FP8", False) # Transposed weight storage for faster gradient accumulation self.lm_head = CastedLinearT(model_dim, vocab_size, use_fp8=use_fp8, x_s=100/448, w_s=1.6/448, grad_s=0.75/448) nn.init.normal_(self.lm_head.weight, mean=0, std=0.005) self.lm_head.weight.label = 'lm_head' self.embed = nn.Embedding(vocab_size, model_dim) self.embed.weight.label = 'embed' with torch.no_grad(): self.embed.weight.copy_(self.lm_head.weight.T) self.bigram_embed = nn.Embedding(args.bigram_vocab_size, model_dim) self.bigram_embed.weight.label = 'bigram_embed' nn.init.zeros_(self.bigram_embed.weight) # x0_lambdas separated out for different optimizer treatment (no beta smoothing) self.x0_lambdas = nn.Parameter(torch.zeros(num_layers)) self.x0_lambdas.label = 'x0_lambdas' pad = (-num_layers * 3 - 3) % dist.get_world_size() # updated: 3*num_layers instead of 4* self.scalars = nn.Parameter( torch.cat( [ 1.1 * torch.ones(num_layers), # resid lambdas. 1.1 init such that layer i weight is i^(num_layers-i). *[torch.tensor([0.5, 1.0]) for _ in range(num_layers)], # SA lambdas 0.1 * torch.ones(num_layers), # bigram lambdas torch.zeros(1), # smear_lambda 0.5*torch.ones(1), # backout_lambda -1.5 * torch.ones(1), # skip_lambda -> σ(-1.5) ≈ 0.18 torch.ones(pad), ] ) ) self.scalars.label = 'scalars' def forward(self, input_seq: Tensor, target_seq: Tensor, seqlens: Tensor, bigram_input_seq: Tensor, schedule_cfg: ForwardScheduleConfig): assert input_seq.ndim == 1 # unpack schedule_cfg mtp_weights, ws_short, ws_long = schedule_cfg.mtp_weights, schedule_cfg.ws_short, schedule_cfg.ws_long # set configs skip_connections = [] skip_in = [3] # long attention window on layer 3 skip_out = [6] # no attn op on layer 6 x_backout = None backout_layer = 7 # set lambdas resid_lambdas = self.scalars[: 1 * self.num_layers] x0_lambdas = self.x0_lambdas sa_lambdas = self.scalars[1 * self.num_layers: 3 * self.num_layers].view(-1, 2) bigram_lambdas = self.scalars[3 * self.num_layers: 4 * self.num_layers] smear_lambda = self.scalars[4 * self.num_layers] backout_lambda = self.scalars[4 * self.num_layers+1] skip_lambda = self.scalars[4 * self.num_layers+2] # set block masks and key shift short_bm = ws_short * args.block_size long_bm = ws_long * args.block_size bm_sizes = [short_bm, short_bm, short_bm, long_bm, short_bm, short_bm, None, short_bm, short_bm, short_bm, long_bm] assert len(bm_sizes) == self.num_layers key_offset = [b==long_bm for b in bm_sizes] # apply partial key offset to long windows # Embedding lookup - embed is synced from lm_head during tied phase by optimizer x = self.embed(input_seq) x0_bigram = self.bigram_embed(bigram_input_seq)[None] # Value embeddings - always computed (not precomputed) ve = [value_embed(input_seq) for value_embed in self.value_embeds] # 01 ... 01 structure on token value embeddings by @YouJiacheng, improved on @leloykun's U-net structure # shifting first layer updates this to 01 ... 01 @photomz ve = [ve[0], ve[1]] + [None] * (self.num_layers - 5) + [ve[2], ve[3], ve[4]] assert len(ve) == self.num_layers # smear token embed forward 1 position @classiclarryd smear_gate_out = smear_lambda * torch.sigmoid(self.smear_gate(x[1:, :self.smear_gate.weight.size(-1)])) x = torch.cat([x[:1], x[1:] + smear_gate_out * x[:-1]]) x = x0 = norm(x[None]) # unbind gate banks to avoid select_backwards kernel ag = [w.bfloat16() for w in self.attn_gate_bank.unbind(0)] veg = [w.bfloat16() for w in self.ve_gate_bank.unbind(0)] attn_gates = ag[:6] + [None] + ag[6:] ve_gates = [veg[0], veg[1]] + [None] * (self.num_layers - 5) + [veg[2], veg[3], veg[4]] assert len(attn_gates) == self.num_layers assert len(ve_gates) == self.num_layers # unbind weight banks to avoid select_backwards kernel attn_weights = self.attn_bank.unbind(0) # tuple of [4*dim, hdim] tensors mlp_fcs = self.mlp_bank[:, 0, :, :].unbind(0) # tuple of [mlp_hdim, dim] tensors mlp_projs = self.mlp_bank[:, 1, :, :].unbind(0) # tuple of [mlp_hdim, dim] tensors for i in range(self.num_layers): yarn = self.yarn_paired_head if i in self.paired_head_layers else self.yarn attn_args = AttnArgs( ve=ve[i], sa_lambdas=sa_lambdas[i], seqlens=seqlens, bm_size=bm_sizes[i], yarn=yarn, key_offset=key_offset[i], attn_gate_w=attn_gates[i], ve_gate_w=ve_gates[i] ) if i in skip_out: skip_gate_out = torch.sigmoid(skip_lambda) * 2 * torch.sigmoid(self.skip_gate(x0[..., :self.skip_gate.weight.size(-1)])) x = x + skip_gate_out * skip_connections.pop() if i == 0: x = (resid_lambdas[0] + x0_lambdas[0]) * x + bigram_lambdas[0] * x0_bigram else: x = resid_lambdas[i] * x + x0_lambdas[i] * x0 + bigram_lambdas[i] * x0_bigram # Get weights for this layer from banks qkvo_w = attn_weights[self.layer_to_attn_idx[i]] if i in self.layer_to_attn_idx else None c_fc = mlp_fcs[self.layer_to_mlp_idx[i]] if i in self.layer_to_mlp_idx else None c_proj = mlp_projs[self.layer_to_mlp_idx[i]] if i in self.layer_to_mlp_idx else None x = self.blocks[i](x, attn_args, qkvo_w, c_fc, c_proj) if i in skip_in: skip_connections.append(x) if i == backout_layer: x_backout = x # back out contributions from first 7 layers that are only required for downstream context and not direct prediction x -= backout_lambda * x_backout x = norm(x) logits = self.lm_head(x) # @Grad62304977 added tanh softcapping following Gemma 2 paper, @KoszarskyB reduced it from 30 to 15 # @YouJiacheng shifted it by +15 (2*sigmoid(2*x)=tanh(x)+1). @classiclarryd updated to 23*sigmoid((logits+5)/7.5) if self.training: losses = FusedSoftcappedCrossEntropy.apply(logits.view(-1, logits.size(-1)), target_seq, mtp_weights) loss = losses.sum() else: logits = 23 * torch.sigmoid((logits + 5) / 7.5) logits_for_loss = logits.float() loss = F.cross_entropy(logits_for_loss.view(-1, logits_for_loss.size(-1)), target_seq, reduction="mean") return loss # ----------------------------------------------------------------------------- # Distributed data loader 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) # avoid pin_memory copy by @YouJiacheng f.seek(256 * 4) nbytes = f.readinto(tokens.numpy()) # avoid bytes->array copy by @YouJiacheng assert nbytes == 2 * num_tokens, "number of tokens read does not match header" return tokens BOS_ID = 50256 class BOSFinder: # Helper for getting sequences that start at the beginning of documents by @varunneal based on work by @classiclarryd def __init__(self, tokens: Tensor, world_size: int = 1, quickload: bool = False): # Precompute BOS positions once per shard self.tokens=tokens self.size = tokens.numel() self.quickload = quickload if quickload: # only scan first 4 million tokens, then kickoff async thread to scan rest self.bos_idx = (tokens[:4_000_000] == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self.thread = None self.ready = threading.Event() self.start() else: self.bos_idx = (tokens == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self.i = 0 self.world_size = world_size self.batch_iter = 0 def _load(self): self.bos_idx_async = (self.tokens == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self.ready.set() def start(self): self.ready.clear() self.thread = threading.Thread(target=self._load) self.thread.start() def get(self): if self.thread: self.ready.wait() self.thread.join() self.bos_idx = self.bos_idx_async def next_batch(self, num_tokens_local: int, max_seq_len: int): # if quickload was used, repoint to the full dataset after 5 batches if self.quickload and self.batch_iter==5: self.get() n = len(self.bos_idx) starts = [[] for _ in range(self.world_size)] ends = [[] for _ in range(self.world_size)] idx = self.i for r in range(self.world_size): cur_len = 0 while cur_len <= num_tokens_local: if idx >= n: raise StopIteration(f"Insufficient BOS ahead; hit tail of shard.") cur = self.bos_idx[idx] starts[r].append(cur) end = min(self.bos_idx[idx + 1] if idx + 1 < n else self.size, cur + max_seq_len, cur + num_tokens_local - cur_len + 1) ends[r].append(end) cur_len += end - cur idx += 1 assert cur_len == num_tokens_local + 1 self.i = idx self.batch_iter+=1 return starts, ends class DataPreloader: # Helper for asynchronously loading next shard and indexing bos tokens def __init__(self, file_iter, world_size: int = 1): self.file_iter = file_iter self.world_size = world_size self.thread = None self.data = None self.ready = threading.Event() def _load(self): tokens = _load_data_shard(next(self.file_iter)) self.data = (tokens, BOSFinder(tokens, self.world_size)) self.ready.set() def start(self): self.ready.clear() self.thread = threading.Thread(target=self._load) self.thread.start() def get(self): if self.thread: self.ready.wait() self.thread.join() return self.data def get_bigram_hash(x): """ Computes bigram hash for each position using [prev_token, curr_token]. Multiply by arbitary large ints to get even spread over int32 range. Position 0 is mapped to the reserved index (vocab_size - 1). BOS_tokens within the batch will hash based on last token of prior doc. Masking this ran slower and showed no improvement. """ rand_int_1 = 36313 rand_int_2 = 27191 mod = args.bigram_vocab_size-1 x = x.to(torch.int32).clone() x[0] = mod x[1:] = torch.bitwise_xor(rand_int_1 * x[1:], rand_int_2 * x[:-1]) % mod return x def distributed_data_generator(filename_pattern: str, num_tokens: int, max_seq_len: int, grad_accum_steps: int = 1, align_to_bos: bool = True): # align_to_bos: each sequence begins with Beginning of Sequence token, sequences truncated to max_seq_len rank = dist.get_rank() if dist.is_initialized() else 0 world_size = dist.get_world_size() if dist.is_initialized() else 1 assert num_tokens % (world_size * grad_accum_steps) == 0, "Batch size must be divisible by world size" num_tokens = num_tokens // grad_accum_steps files = [Path(file) for file in sorted(glob.glob(filename_pattern))] if not files: raise FileNotFoundError(f"No files found for pattern: {filename_pattern}") file_iter = iter(files) # Use itertools.cycle(files) for multi-epoch training tokens = _load_data_shard(next(file_iter)) if align_to_bos: finder = BOSFinder(tokens, world_size=world_size, quickload=True) preloader = DataPreloader(file_iter, world_size) preloader.start() else: pos = 0 # for unaligned case while True: num_tokens_local = num_tokens // world_size max_num_docs = next_multiple_of_n(num_tokens_local // 300, n=128) # median doc length is ~400 if align_to_bos: try: seq_starts, seq_ends = finder.next_batch(num_tokens_local, max_seq_len) start_idxs, end_idxs = torch.tensor(seq_starts[rank]), torch.tensor(seq_ends[rank]) except StopIteration: # This shard is exhausted, load the next one in the next loop iteration. tokens, finder = preloader.get() preloader.start() continue buf = torch.cat([tokens[i:j] for i, j in zip(start_idxs, end_idxs)]) _inputs = buf[:-1] _targets = buf[1:] end_idxs[-1] -= 1 # last document was too long to account for _targets offset cum_lengths = (end_idxs - start_idxs).cumsum(0) else: if pos + num_tokens + 1 >= len(tokens): # should not occur for val data tokens, pos = _load_data_shard(next(file_iter)), 0 pos_local = pos + rank * num_tokens_local buf = tokens[pos_local: pos_local + num_tokens_local + 1] _inputs = buf[:-1].view(num_tokens_local, ) _targets = buf[1:].view(num_tokens_local, ) cum_lengths = torch.nonzero(_inputs == BOS_ID)[:, 0] pos += num_tokens _cum_lengths = torch.full((max_num_docs,), num_tokens_local) _cum_lengths[0] = 0 _cum_lengths[1:len(cum_lengths) + 1] = cum_lengths # Cast to int32 on CPU before transfer to avoid dtype conversion during .to() _inputs = _inputs.to(dtype=torch.int32) _targets = _targets.to(dtype=torch.int64) _cum_lengths = _cum_lengths.to(dtype=torch.int32) _bigram_inputs = get_bigram_hash(_inputs) new_params = yield ( _inputs.to(device="cuda", non_blocking=True), _targets.to(device="cuda", non_blocking=True), _cum_lengths.to(device="cuda", non_blocking=True), _bigram_inputs.to(device="cuda", non_blocking=True) ) if new_params is not None: # makes it possible for generator to receive new (num_tokens, max_seq_len, grad_accum_steps) via .send() new_num_tokens, new_max_seq_len, new_grad_accum_steps = new_params assert new_num_tokens % (world_size * new_grad_accum_steps) == 0, "Num tokens must be divisible by world size" num_tokens = new_num_tokens // new_grad_accum_steps max_seq_len = new_max_seq_len # ----------------------------------------------------------------------------- # Training Management def get_bs(step: int): if step >= args.num_scheduled_iterations: return args.train_bs_extension x = step / args.num_scheduled_iterations bs_idx = int(len(args.train_bs_schedule) * x) return args.train_bs_schedule[bs_idx] def get_ws(step: int): # set short window size to half of long window size # Higher ws on "extension" steps if step >= args.num_scheduled_iterations: return args.ws_final // 2, args.ws_final x = step / args.num_scheduled_iterations assert 0 <= x < 1 ws_idx = int(len(args.ws_schedule) * x) return args.ws_schedule[ws_idx] // 2, args.ws_schedule[ws_idx] # learning rate schedule: tied to batch size schedule, with cooldown at the end. def get_lr(step: int): if step > args.num_scheduled_iterations: return 0.1 lr_max = 1.0 x = step / args.num_scheduled_iterations if x > 1/3: lr_max = 1.52 # (16/8)**0.6 if x > 2/3: lr_max = 1.73 # (24/8)**0.5 if x >= 1 - args.cooldown_frac: w = (1 - x) / args.cooldown_frac lr = lr_max * w + (1 - w) * 0.1 return lr return lr_max def get_muon_momentum(step: int, muon_warmup_steps=300, muon_cooldown_steps=50, momentum_min=0.85, momentum_max=0.95): # warmup phase: linearly increase momentum from min to max # cooldown phase: linearly decrease momentum from max to min momentum_cd_start = args.num_iterations - muon_cooldown_steps if step < muon_warmup_steps: frac = step / muon_warmup_steps momentum = momentum_min + frac * (momentum_max - momentum_min) elif step > momentum_cd_start: frac = (step - momentum_cd_start) / muon_cooldown_steps momentum = momentum_max - frac * (momentum_max - momentum_min) else: momentum = momentum_max return momentum class TrainingManager(): """ Manages the NorMuonAndAdam for all parameters with explicit ordering. Notable Features: 1. Scalars are given higher momentum terms to smooth learning @ChrisJMcCormick 2. Adam optimizers are only stepped on odd steps @classiclarryd 3. Explicit scatter_order and work_order for communication scheduling (no backward hooks) 4. Muon has a linear momentum warmup and cooldown schedule 5. Learning rates follow a linear decay schedule 6. Embed is tied to lm_head until split step (2/3 of training), then untied @classiclarryd Manages model architecture, data, and target that changes during training Notable Features: 1. Multi Token Prediction schedule of [1, 0.5, 0.25->0] -> [1, 0.5->0] -> [1] @varunneal 2. Sliding Attention window schedule of [1,3] -> [3,7] -> [5,11] -> [6,13] 3. YaRN updates to RoPE on window changes 4. Split embed and lm_head at 2/3 of training (weights and optimizer state copied) 5. Batch size schedule of 8 -> 16 -> 24 6. Post training extension of long windows from 13 to 20 """ def __init__(self, model): self.mtp_weights_schedule = self._build_mtp_schedule() self.model = model # - Ordering dictates when to launch reduce/reduce_scatter operations # - "sharded" parameters use reduce_scatter/all_gather and "replicated" ones use all_reduce # - lr_mul and wd_mul are per-parameter learning rate and weight decay multipliers self.param_table = { "attn": {"optim": "normuon", "comms": "sharded", "adam_betas": None}, "mlp": {"optim": "normuon", "comms": "sharded", "adam_betas": None}, "scalars": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99], "lr_mul": 5.0, "wd_mul": 0.0}, "ve0": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve1": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve2": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve3": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve4": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "bigram_embed": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "smear_gate": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99], "lr_mul": 0.01, "wd_mul": 0.0}, "skip_gate": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99], "lr_mul": 0.05, "wd_mul": 0.0}, "attn_gate_bank": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99]}, "ve_gate_bank": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99]}, "x0_lambdas": {"optim": "adam", "comms": "replicated", "adam_betas": [0.65, 0.95], "lr_mul": 5.0, "wd_mul": 0.0}, "lm_head": {"optim": "adam", "comms": "sharded", "adam_betas": [0.5, 0.95], "wd_mul": 150.}, "embed": {"optim": "adam", "comms": "sharded", "adam_betas": [0.5, 0.95], "wd_mul": 150.}, } # - Process smaller/faster params first while large reduces complete # - lm_head must complete before embed sync (when tied) self.work_order = [ "scalars", "smear_gate", "skip_gate", "attn_gate_bank", "ve_gate_bank", "x0_lambdas", # Small, fast "ve0", "ve1", "ve2", "ve3", "ve4", "bigram_embed", # Medium "lm_head", "embed", # lm_head must complete before embed sync (when tied) "attn", "mlp", # Large, polar express - process last to maximize overlap ] adam_defaults = dict( lr=0.008, eps=1e-10, weight_decay=0.005, ) normuon_defaults = dict( lr=0.023, momentum=0.95, beta2=0.95, weight_decay=1.2, ) self.optimizer = NorMuonAndAdam( model.named_parameters(), param_table=self.param_table, scatter_order=list(self.param_table.keys()), # Dict order defines scatter priority work_order=self.work_order, adam_defaults=adam_defaults, normuon_defaults=normuon_defaults, ) # Split embed from lm_head at 2/3 of training (on an odd step so Adam updates) self.split_step = math.ceil(args.split_embed_frac * args.num_scheduled_iterations) | 1 self.reset() def _build_mtp_schedule(self): # Precompute MTP weights for all steps to avoid tensor allocation during training # Schedule: [1, 0.5, 0.25->0] -> [1, 0.5->0] -> [1] mtp_weights_schedule = [] for s in range(args.num_iterations + 1): x = s / args.num_scheduled_iterations if x < 1/3: w = [1.0, 0.5, 0.25 * (1 - 3*x)] elif x < 2/3: w = [1.0, 0.5 * (1 - (3*x - 1))] else: w = [1.0] mtp_weights_schedule.append(torch.tensor(w, device=device)) return mtp_weights_schedule def apply_final_ws_ext(self): self.ws_long = args.ws_validate_post_yarn_ext def get_forward_args(self): return ForwardScheduleConfig( mtp_weights = self.mtp_weights, ws_short = self.ws_short, ws_long = self.ws_long ) def _is_adam_step(self, step: int): """Adam params are only updated on odd steps.""" return step % 2 == 1 def get_transition_steps(self): transition_steps = [] ws_short, ws_long = get_ws(0) for step in range(1, args.num_iterations): ws_short, new_ws_long = get_ws(step) if new_ws_long != ws_long: transition_steps.append(step) ws_long = new_ws_long return transition_steps def advance_schedule(self, step: int): self.ws_short, new_ws_long = get_ws(step) if new_ws_long != self.ws_long: self.model.yarn.apply(self.ws_long, new_ws_long) self.model.yarn_paired_head.apply(self.ws_long, new_ws_long) new_batch_size = get_bs(step) if new_batch_size != self.batch_size: self.train_loader_send_args = (new_batch_size, args.train_max_seq_len, grad_accum_steps) self.batch_size = new_batch_size else: self.train_loader_send_args = None self.ws_long = new_ws_long self.mtp_weights = self.mtp_weights_schedule[step] def step_optimizers(self, step: int): step_lr = get_lr(step) muon_momentum = get_muon_momentum(step) do_adam = self._is_adam_step(step) # Update learning rates and momentum for all params for param, p_cfg in self.optimizer.param_cfgs.items(): p_cfg.lr = p_cfg.initial_lr * step_lr if p_cfg.optim == "normuon": p_cfg.momentum = muon_momentum # Step optimizer with do_adam flag self.optimizer.step(do_adam=do_adam) # At split step: copy lm_head optimizer state to embed and mark as split if step == self.split_step: self.optimizer.copy_lm_state_to_embed() def reset(self, state=None): if state is not None: self.optimizer.load_state_dict(state) # Reset NorMuon momentum buffers and split_embed state self.optimizer.reset() self.ws_short, self.ws_long = get_ws(0) self.batch_size = get_bs(0) self.model.yarn.reset() self.model.yarn_paired_head.reset() def get_state(self): return copy.deepcopy(self.optimizer.state_dict()) # ----------------------------------------------------------------------------- # int main @dataclass class Hyperparameters: # data train_files: str = "data/fineweb10B/fineweb_train_*.bin" # input .bin to train on val_files: str = "data/fineweb10B/fineweb_val_*.bin" # input .bin to eval validation loss on val_tokens: int = 10485760 # how many tokens of validation data? it's important to keep this fixed for consistent comparisons # batch sizes train_bs_schedule: tuple = (8 * 2048 * 8, 16 * 2048 * 8, 24 * 2048 * 8) train_bs_extension: int = 24 * 2048 * 8 train_max_seq_len: int = 128 * 16 val_batch_size: int = 4 * 64 * 1024 * 8 # optimization num_scheduled_iterations: int = 1535 # number of steps to complete lr and ws schedule num_extension_iterations: int = 40 # number of steps to continue training at final lr and ws num_iterations: int = num_scheduled_iterations + num_extension_iterations cooldown_frac: float = 0.55 # fraction of num_scheduled_iterations spent cooling down the learning rate split_embed_frac: float = 2/3 # fraction of training when embeddings split from lm_head # evaluation and logging run_id: str = f"{uuid.uuid4()}" val_loss_every: int = 250 # every how many steps to evaluate val loss? 0 for only at the end save_checkpoint: bool = False # attention masking block_size: int = 128 ws_schedule: tuple = (3, 7, 11) ws_final: int = 13 # increase final validation ws, used for YaRN extension and short window size @classiclarryd ws_validate_post_yarn_ext: int = 20 # extend long windows out even further after applying YaRN # bigram hash embedding bigram_vocab_size = 50304 * 5 args = Hyperparameters() data_path = os.environ.get("DATA_PATH", ".") args.train_files = os.path.join(data_path, args.train_files) args.val_files = os.path.join(data_path, args.val_files) # torchrun sets these env variables rank = int(os.environ["RANK"]) world_size = int(os.environ["WORLD_SIZE"]) assert 8 % world_size == 0, "world_size must be a divisor of 8" grad_accum_steps = 8 // world_size assert torch.cuda.is_available() 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() master_process = (rank == 0) # this process will do logging, checkpointing etc. # begin logging logfile = None if master_process: run_id = args.run_id os.makedirs("logs", exist_ok=True) logfile = f"logs/{run_id}.txt" print(logfile) def print0(s, console=False): if master_process: with open(logfile, "a") as f: if console: print(s) print(s, file=f) # begin by printing this file (the Python code) print0(code) print0("="*100) # log information about the hardware/software environment this is running on print0(f"Running Python {sys.version}") print0(f"Running PyTorch {torch.version.__version__} compiled for CUDA {torch.version.cuda}") print0(f"Running Triton version {triton.__version__}") def nvidia_smi(): import subprocess # avoid top level import return subprocess.run(["nvidia-smi"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True).stdout print0(nvidia_smi()) print0("="*100) model: nn.Module = GPT( vocab_size=50257, num_layers=11, num_heads=6, head_dim=128, model_dim=768, max_seq_len=args.val_batch_size // (grad_accum_steps * world_size) ).cuda() for m in model.modules(): if isinstance(m, (nn.Embedding, nn.Linear)): m.weight.data = m.weight.data.bfloat16() model.attn_gate_bank.data = model.attn_gate_bank.data.bfloat16() model.ve_gate_bank.data = model.ve_gate_bank.data.bfloat16() model.attn_bank.data = model.attn_bank.data.bfloat16() model.mlp_bank.data = model.mlp_bank.data.bfloat16() for param in model.parameters(): dist.broadcast(param.detach(), 0) model: nn.Module = torch.compile(model, dynamic=False, fullgraph=True) training_manager = TrainingManager(model) ######################################## # Warmup kernels # ######################################## print0("Compiling model and warming up kernels (~7 minutes on first execution)", console=True) # Warmup the training kernels, then re-initialize the state so we aren't cheating initial_state = dict(model=copy.deepcopy(model.state_dict()), optimizer=training_manager.get_state()) # save the initial state train_loader = distributed_data_generator(args.train_files, args.train_bs_schedule[0], args.train_max_seq_len, grad_accum_steps=grad_accum_steps) val_loader = distributed_data_generator(args.val_files, args.val_batch_size, -1, grad_accum_steps=grad_accum_steps, align_to_bos=False) transition_steps = training_manager.get_transition_steps() # first few steps plus transitions warmup_steps = sorted({0, 1, 2} | set(s + offset for s in transition_steps for offset in [-1, 0, 1] if s + offset >= 0)) print0(f"Sampling steps {warmup_steps} for warmup", console=True) for step in warmup_steps: training_manager.advance_schedule(step) model.eval() with torch.no_grad(): inputs, targets, cum_seqlens, bigram_inputs = next(val_loader) model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) model.train() for idx in range(grad_accum_steps): send_args = training_manager.train_loader_send_args inputs, targets, cum_seqlens, bigram_inputs = train_loader.send(send_args) (model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) / grad_accum_steps).backward() training_manager.step_optimizers(step) print0("Resetting Model", console=True) model.zero_grad(set_to_none=True) model.load_state_dict(initial_state["model"]) training_manager.reset(initial_state["optimizer"]) del val_loader, train_loader, initial_state model.train() ######################################## # Training and validation # ######################################## train_loader = distributed_data_generator(args.train_files, args.train_bs_schedule[0], args.train_max_seq_len, grad_accum_steps=grad_accum_steps) gc.collect() training_time_ms = 0 # start the clock torch.cuda.synchronize() t0 = time.perf_counter() # begin training train_steps = args.num_iterations for step in range(train_steps + 1): last_step = (step == train_steps) training_manager.advance_schedule(step) # --------------- VALIDATION SECTION ----------------- if last_step or (args.val_loss_every > 0 and step % args.val_loss_every == 0): if last_step: training_manager.apply_final_ws_ext() # stop the clock torch.cuda.synchronize() training_time_ms += 1000 * (time.perf_counter() - t0) model.eval() assert args.val_tokens % args.val_batch_size == 0 val_steps = grad_accum_steps * args.val_tokens // args.val_batch_size val_loader = distributed_data_generator(args.val_files, args.val_batch_size, -1, grad_accum_steps=grad_accum_steps, align_to_bos=False) val_loss = 0 with torch.no_grad(): for _ in range(val_steps): inputs, targets, cum_seqlens, bigram_inputs = next(val_loader) val_loss += model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) val_loss /= val_steps del val_loader dist.reduce(val_loss, 0, op=dist.ReduceOp.AVG) print0(f"step:{step}/{train_steps} val_loss:{val_loss:.4f} train_time:{training_time_ms:.0f}ms step_avg:{training_time_ms/max(step, 1):.2f}ms", console=True) model.train() # start the clock again torch.cuda.synchronize() t0 = time.perf_counter() if last_step: if master_process and args.save_checkpoint: log = dict(step=step, code=code, model=model.state_dict(), optimizer=training_manager.get_state()) os.makedirs(f"logs/{run_id}", exist_ok=True) torch.save(log, f"logs/{run_id}/state_step{step:06d}.pt") # the last step only has the validation loop, so break to avoid training break # --------------- TRAINING SECTION ----------------- for idx in range(grad_accum_steps): inputs, targets, cum_seqlens, bigram_inputs = train_loader.send(training_manager.train_loader_send_args) (model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) / grad_accum_steps).backward() training_manager.step_optimizers(step) # logging approx_training_time_ms = training_time_ms + 1000 * (time.perf_counter() - t0) print0(f"step:{step+1}/{train_steps} train_time:{approx_training_time_ms:.0f}ms step_avg:{approx_training_time_ms/(step + 1):.2f}ms", console=True) print0(f"peak memory allocated: {torch.cuda.max_memory_allocated() // 1024 // 1024} MiB " f"reserved: {torch.cuda.max_memory_reserved() // 1024 // 1024} MiB", console=True) dist.destroy_process_group() ---------------------------------------- # triton_kernels.py ---------------------------------------- import torch import triton import triton.language as tl from triton.tools.tensor_descriptor import TensorDescriptor # ----------------------------------------------------------------------------- # Triton kernel for symmetric matrix multiplication by @byronxu99 def _get_autotune_configs(): return [ triton.Config( { "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk, "GROUP_SIZE_M": 8, "LOWER_UPPER": 1, }, num_stages=stages, num_warps=warps, ) for bm in [64, 128] for bn in [64, 128, 256] for bk in [64, 128] for stages, warps in [(3, 4), (3, 8), (4, 4)] if bm // bn <= 2 and bn // bm <= 2 ] @triton.jit def _pid_to_block( pid, M, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, GROUP_SIZE_M: tl.constexpr, ): # Split output matrix into blocks of size (BLOCK_SIZE_M, BLOCK_SIZE_N) num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) num_pid_n = tl.cdiv(M, BLOCK_SIZE_N) # Map PID to a single matrix in batch batch_idx = pid // (num_pid_m * num_pid_n) pid = pid % (num_pid_m * num_pid_n) # Map PID to 2D grid of blocks pid_m = pid // num_pid_n pid_n = pid % num_pid_n pid_m, pid_n = tl.swizzle2d(pid_m, pid_n, num_pid_m, num_pid_n, GROUP_SIZE_M) m_idx = pid_m * BLOCK_SIZE_M n_idx = pid_n * BLOCK_SIZE_N return batch_idx, m_idx, n_idx @triton.autotune( configs=_get_autotune_configs(), key=["M", "K", "a_stride_r", "a_stride_c", "c_stride_r", "c_stride_c"], ) @triton.jit def XXT_kernel( A_ptr, C_ptr, M, K, a_stride_b, a_stride_r, a_stride_c, c_stride_b, c_stride_r, c_stride_c, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, LOWER_UPPER: tl.constexpr, ): pid = tl.program_id(axis=0) batch_idx, m_idx, n_idx = _pid_to_block( pid, M, BLOCK_SIZE_M, BLOCK_SIZE_N, GROUP_SIZE_M ) # Skip blocks that don't need to be computed skip_block_below_diag = (LOWER_UPPER == 0) and (n_idx + BLOCK_SIZE_N <= m_idx) skip_block_above_diag = (LOWER_UPPER != 0) and (m_idx + BLOCK_SIZE_M <= n_idx) if skip_block_below_diag or skip_block_above_diag: return # Index into one matrix of batch A_ptr += batch_idx * a_stride_b C_ptr += batch_idx * c_stride_b # Create pointer arrays for A and A.T offs_m = (m_idx + tl.arange(0, BLOCK_SIZE_M)) % M offs_n = (n_idx + tl.arange(0, BLOCK_SIZE_N)) % M offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = A_ptr + (offs_m[:, None] * a_stride_r + offs_k[None, :] * a_stride_c) at_ptrs = A_ptr + (offs_k[:, None] * a_stride_c + offs_n[None, :] * a_stride_r) accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) # Accumulate over blocks of K for k in tl.range(0, tl.cdiv(K, BLOCK_SIZE_K)): a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0) at = tl.load(at_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0) accumulator = tl.dot(a, at, accumulator) a_ptrs += BLOCK_SIZE_K * a_stride_c at_ptrs += BLOCK_SIZE_K * a_stride_c out_dtype = C_ptr.dtype.element_ty output = accumulator.to(out_dtype) # Store block of C offs_cm = m_idx + tl.arange(0, BLOCK_SIZE_M) offs_cn = n_idx + tl.arange(0, BLOCK_SIZE_N) c_ptrs = C_ptr + (offs_cm[:, None] * c_stride_r + offs_cn[None, :] * c_stride_c) c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < M) tl.store(c_ptrs, output, mask=c_mask) # Store block of C mirrored across the diagonal c_ptrs_t = C_ptr + (offs_cn[:, None] * c_stride_r + offs_cm[None, :] * c_stride_c) c_mask_t = (offs_cn[:, None] < M) & (offs_cm[None, :] < M) tl.store(c_ptrs_t, output.T, mask=c_mask_t) def XXT(A: torch.Tensor, out: torch.Tensor): """ Launch Triton kernel to compute C = A @ A.T """ assert A.ndim == 2 or A.ndim == 3 M, K = A.shape[-2:] assert out.size(-2) == M, "Output matrix has incorrect shape" assert out.size(-1) == M, "Output matrix has incorrect shape" batch_size = A.size(0) if A.ndim == 3 else 1 input_batch_stride = A.stride(0) if A.ndim == 3 else 0 output_batch_stride = out.stride(0) if out.ndim == 3 else 0 grid = lambda meta: ( batch_size * triton.cdiv(M, meta["BLOCK_SIZE_M"]) * triton.cdiv(M, meta["BLOCK_SIZE_N"]), ) XXT_kernel[grid]( A_ptr=A, C_ptr=out, M=M, K=K, a_stride_b=input_batch_stride, a_stride_r=A.stride(-2), a_stride_c=A.stride(-1), c_stride_b=output_batch_stride, c_stride_r=out.stride(-2), c_stride_c=out.stride(-1), ) return out @triton.autotune( configs=_get_autotune_configs(), key=["M", "a_stride_r", "a_stride_c", "c_stride_r", "c_stride_c"], ) @triton.jit def ba_plus_cAA_kernel( A_ptr, C_ptr, M, a_stride_b, a_stride_r, a_stride_c, c_stride_b, c_stride_r, c_stride_c, alpha, beta, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, LOWER_UPPER: tl.constexpr, ): # This is mostly duplicated from XXT_kernel, but also loads and adds a block of A # Performance is slightly slower than XXT_kernel, so we use two separate kernels pid = tl.program_id(axis=0) batch_idx, m_idx, n_idx = _pid_to_block( pid, M, BLOCK_SIZE_M, BLOCK_SIZE_N, GROUP_SIZE_M ) # Skip blocks that don't need to be computed skip_block_below_diag = (LOWER_UPPER == 0) and (n_idx + BLOCK_SIZE_N <= m_idx) skip_block_above_diag = (LOWER_UPPER != 0) and (m_idx + BLOCK_SIZE_M <= n_idx) if skip_block_below_diag or skip_block_above_diag: return # Index into one matrix of batch A_ptr += batch_idx * a_stride_b C_ptr += batch_idx * c_stride_b # Create pointer arrays for A and A.T offs_m = (m_idx + tl.arange(0, BLOCK_SIZE_M)) % M offs_n = (n_idx + tl.arange(0, BLOCK_SIZE_N)) % M offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = A_ptr + (offs_m[:, None] * a_stride_r + offs_k[None, :] * a_stride_c) at_ptrs = A_ptr + (offs_k[:, None] * a_stride_c + offs_n[None, :] * a_stride_r) accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) # Accumulate over blocks of K for k in tl.range(0, tl.cdiv(M, BLOCK_SIZE_K)): a = tl.load(a_ptrs, mask=offs_k[None, :] < M - k * BLOCK_SIZE_K, other=0.0) at = tl.load(at_ptrs, mask=offs_k[:, None] < M - k * BLOCK_SIZE_K, other=0.0) accumulator = tl.dot(a, at, accumulator) a_ptrs += BLOCK_SIZE_K * a_stride_c at_ptrs += BLOCK_SIZE_K * a_stride_c # Load block of A to add (corresponds to the current block of C) offs_am = m_idx + tl.arange(0, BLOCK_SIZE_M) offs_an = n_idx + tl.arange(0, BLOCK_SIZE_N) a_add_ptrs = A_ptr + (offs_am[:, None] * a_stride_r + offs_an[None, :] * a_stride_c) a_add_mask = (offs_am[:, None] < M) & (offs_an[None, :] < M) a_add = tl.load(a_add_ptrs, mask=a_add_mask, other=0.0).to(tl.float32) # Apply alpha and beta accumulator *= alpha accumulator += a_add * beta out_dtype = C_ptr.dtype.element_ty output = accumulator.to(out_dtype) # Store block of C offs_cm = m_idx + tl.arange(0, BLOCK_SIZE_M) offs_cn = n_idx + tl.arange(0, BLOCK_SIZE_N) c_ptrs = C_ptr + (offs_cm[:, None] * c_stride_r + offs_cn[None, :] * c_stride_c) c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < M) tl.store(c_ptrs, output, mask=c_mask) # Store block of C mirrored across the diagonal c_ptrs_t = C_ptr + (offs_cn[:, None] * c_stride_r + offs_cm[None, :] * c_stride_c) c_mask_t = (offs_cn[:, None] < M) & (offs_cm[None, :] < M) tl.store(c_ptrs_t, output.T, mask=c_mask_t) def ba_plus_cAA(A: torch.Tensor, alpha: float, beta: float, out: torch.Tensor): """ Launch Triton kernel to compute C = alpha * A @ A.T + beta * A """ assert A.ndim == 2 or A.ndim == 3 M, K = A.shape[-2:] assert M == K, "Input matrix must be square" assert out.size(-2) == M assert out.size(-1) == M batch_size = A.size(0) if A.ndim == 3 else 1 input_batch_stride = A.stride(0) if A.ndim == 3 else 0 output_batch_stride = out.stride(0) if out.ndim == 3 else 0 grid = lambda meta: ( batch_size * triton.cdiv(M, meta["BLOCK_SIZE_M"]) * triton.cdiv(M, meta["BLOCK_SIZE_N"]), ) ba_plus_cAA_kernel[grid]( A_ptr=A, C_ptr=out, M=M, a_stride_b=input_batch_stride, a_stride_r=A.stride(-2), a_stride_c=A.stride(-1), c_stride_b=output_batch_stride, c_stride_r=out.stride(-2), c_stride_c=out.stride(-1), alpha=alpha, beta=beta, ) return out # ----------------------------------------------------------------------------- # Triton kernel for MLP: relu(x @ W1.T)^2, by @andrewbriand, @jrauvola @triton.jit def linear_relu_square_kernel(a_desc, b_desc, c_desc, aux_desc, M, N, K, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, NUM_SMS: tl.constexpr, FORWARD: tl.constexpr, ): dtype = tl.bfloat16 start_pid = tl.program_id(axis=0) num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) k_tiles = tl.cdiv(K, BLOCK_SIZE_K) num_tiles = num_pid_m * num_pid_n tile_id_c = start_pid - NUM_SMS num_pid_in_group = GROUP_SIZE_M * num_pid_n for tile_id in tl.range(start_pid, num_tiles, NUM_SMS, flatten=True): pid_m = tile_id // num_pid_n pid_n = tile_id % num_pid_n offs_am = pid_m * BLOCK_SIZE_M offs_bn = pid_n * BLOCK_SIZE_N accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) for ki in range(k_tiles): offs_k = ki * BLOCK_SIZE_K a = a_desc.load([offs_am, offs_k]) b = b_desc.load([offs_bn, offs_k]) accumulator = tl.dot(a, b.T, accumulator) tile_id_c += NUM_SMS pid_m = tile_id // num_pid_n pid_n = tile_id % num_pid_n offs_am_c = pid_m * BLOCK_SIZE_M offs_bn_c = pid_n * BLOCK_SIZE_N acc = tl.reshape(accumulator, (BLOCK_SIZE_M, 2, BLOCK_SIZE_N // 2)) acc = tl.permute(acc, (0, 2, 1)) acc0, acc1 = tl.split(acc) c0 = acc0.to(dtype) if not FORWARD: c0_pre = aux_desc.load([offs_am_c, offs_bn_c]) c0 = 2 * c0 * tl.where(c0_pre > 0, c0_pre, 0) c_desc.store([offs_am_c, offs_bn_c], c0) if FORWARD: c0_post = tl.maximum(c0, 0) c0_post = c0_post * c0_post aux_desc.store([offs_am_c, offs_bn_c], c0_post) c1 = acc1.to(dtype) if not FORWARD: c1_pre = aux_desc.load([offs_am_c, offs_bn_c + BLOCK_SIZE_N // 2]) c1 = 2 * c1 * tl.where(c1_pre > 0, c1_pre, 0) c_desc.store([offs_am_c, offs_bn_c + BLOCK_SIZE_N // 2], c1) if FORWARD: c1_post = tl.maximum(c1, 0) c1_post = c1_post * c1_post aux_desc.store([offs_am_c, offs_bn_c + BLOCK_SIZE_N // 2], c1_post) def linear_relu_square(a, b, aux=None): M, K = a.shape N, K = b.shape dtype = a.dtype c = torch.empty((M, N), device=a.device, dtype=dtype) FORWARD = False if aux is None: FORWARD = True aux = torch.empty((M, N), device=a.device, dtype=dtype) NUM_SMS = torch.cuda.get_device_properties("cuda").multi_processor_count BLOCK_SIZE_M = 128 BLOCK_SIZE_N = 256 BLOCK_SIZE_K = 64 num_stages = 4 if FORWARD else 3 num_warps = 8 a_desc = TensorDescriptor.from_tensor(a, [BLOCK_SIZE_M, BLOCK_SIZE_K]) b_desc = TensorDescriptor.from_tensor(b, [BLOCK_SIZE_N, BLOCK_SIZE_K]) c_desc = TensorDescriptor.from_tensor(c, [BLOCK_SIZE_M, BLOCK_SIZE_N // 2]) aux_desc = TensorDescriptor.from_tensor(aux, [BLOCK_SIZE_M, BLOCK_SIZE_N // 2]) def grid(META): return (min( NUM_SMS, triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N), ), ) linear_relu_square_kernel[grid]( a_desc, b_desc, c_desc, aux_desc, M, N, K, BLOCK_SIZE_M=BLOCK_SIZE_M, BLOCK_SIZE_N=BLOCK_SIZE_N, BLOCK_SIZE_K=BLOCK_SIZE_K, GROUP_SIZE_M=1, NUM_SMS=NUM_SMS, FORWARD=FORWARD, num_stages=num_stages, num_warps=num_warps ) if FORWARD: return c, aux else: return c class FusedLinearReLUSquareFunction(torch.autograd.Function): @staticmethod def forward(ctx, x, W1, W2): pre, post = linear_relu_square(x.view((-1, x.shape[-1])), W1) x3 = post @ W2 ctx.save_for_backward(x, W1, W2, pre, post) return x3.view(x.shape) @staticmethod def backward(ctx, grad_output): x, W1, W2, pre, post = ctx.saved_tensors dW2 = post.T @ grad_output dpre = linear_relu_square(grad_output.view((-1, grad_output.shape[-1])), W2, aux=pre) dW1 = dpre.T @ x dx = dpre @ W1 return dx.view(x.shape), dW1, dW2 # ----------------------------------------------------------------------------- # Fused Softcapped Cross Entropy @triton.jit def fused_softcapped_entropy_fwd_kernel( logits_ptr, losses_ptr, lse_ptr, targets_ptr, mtp_weights_ptr, stride_logits_n, stride_logits_v, n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE: tl.constexpr ): row_idx = tl.program_id(0).to(tl.int64) logits_row_ptr = logits_ptr + row_idx * stride_logits_n max_val = -float('inf') sum_exp = 0.0 for off in range(0, n_cols, BLOCK_SIZE): cols = off + tl.arange(0, BLOCK_SIZE) mask = cols < n_cols val = tl.load(logits_row_ptr + cols, mask=mask, other=-float('inf')).to(tl.float32) z = A * tl.sigmoid((val + B) / C) z = tl.where(mask, z, -float('inf')) curr_max = tl.max(z, axis=0) new_max = tl.maximum(max_val, curr_max) sum_exp = sum_exp * tl.exp(max_val - new_max) + tl.sum(tl.exp(z - new_max), axis=0) max_val = new_max lse = max_val + tl.log(sum_exp) tl.store(lse_ptr + row_idx, lse) total_loss = 0.0 for k in range(n_predict): target_idx = row_idx + k if target_idx < n_rows: weight = tl.load(mtp_weights_ptr + k) if weight > 0: target = tl.load(targets_ptr + target_idx).to(tl.int32) if target >= 0 and target < n_cols: val_target = tl.load(logits_row_ptr + target).to(tl.float32) z_target = A * tl.sigmoid((val_target + B) / C) total_loss += weight * (lse - z_target) tl.store(losses_ptr + row_idx, total_loss) @triton.jit def fused_softcapped_entropy_bwd_kernel( grad_input_ptr, grad_output_ptr, lse_ptr, logits_ptr, targets_ptr, mtp_weights_ptr, stride_logits_n, stride_logits_v, stride_grad_n, stride_grad_v, n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE: tl.constexpr ): row_idx = tl.program_id(0).to(tl.int64) logits_row_ptr = logits_ptr + row_idx * stride_logits_n grad_row_ptr = grad_input_ptr + row_idx * stride_grad_n lse = tl.load(lse_ptr + row_idx) grad_loss = tl.load(grad_output_ptr + row_idx) S_w = 0.0 for k in range(n_predict): if row_idx + k < n_rows: S_w += tl.load(mtp_weights_ptr + k) for off in range(0, n_cols, BLOCK_SIZE): cols = off + tl.arange(0, BLOCK_SIZE) mask = cols < n_cols val = tl.load(logits_row_ptr + cols, mask=mask, other=0.0).to(tl.float32) u = (val + B) / C sigmoid_u = tl.sigmoid(u) z = A * sigmoid_u p = tl.exp(z - lse) term1 = S_w * p term2 = tl.zeros([BLOCK_SIZE], dtype=tl.float32) for k in range(n_predict): if row_idx + k < n_rows: target = tl.load(targets_ptr + row_idx + k).to(tl.int32) weight = tl.load(mtp_weights_ptr + k) term2 += tl.where(cols == target, weight, 0.0) grad_z = grad_loss * (term1 - term2) dz_dx = (1.0 / C) * z * (1.0 - sigmoid_u) grad_x = grad_z * dz_dx tl.store(grad_row_ptr + cols, grad_x.to(tl.bfloat16), mask=mask) class FusedSoftcappedCrossEntropy(torch.autograd.Function): @staticmethod def forward(ctx, logits, targets, mtp_weights, A=23.0, B=5.0, C=7.5): n_rows, n_cols = logits.shape if mtp_weights is None: mtp_weights = torch.tensor([1.0], device=logits.device, dtype=torch.float32) n_predict = mtp_weights.shape[0] losses = torch.empty(n_rows, dtype=torch.float32, device=logits.device) lse = torch.empty(n_rows, dtype=torch.float32, device=logits.device) logits = logits.contiguous() targets = targets.contiguous() mtp_weights = mtp_weights.contiguous() grid = (n_rows,) fused_softcapped_entropy_fwd_kernel[grid]( logits, losses, lse, targets, mtp_weights, logits.stride(0), logits.stride(1), n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE=1024, num_warps=8, num_stages=4 ) ctx.save_for_backward(logits, targets, mtp_weights, lse) ctx.params = (A, B, C) return losses @staticmethod def backward(ctx, grad_output): logits, targets, mtp_weights, lse = ctx.saved_tensors A, B, C = ctx.params n_rows, n_cols = logits.shape n_predict = mtp_weights.shape[0] grad_input = torch.empty((n_rows, n_cols), dtype=torch.bfloat16, device=logits.device) grad_output = grad_output.contiguous() grid = (n_rows,) fused_softcapped_entropy_bwd_kernel[grid]( grad_input, grad_output, lse, logits, targets, mtp_weights, logits.stride(0), logits.stride(1), grad_input.stride(0), grad_input.stride(1), n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE=1024, num_warps=8, num_stages=4 ) return grad_input, None, None, None, None, None ==================================================================================================== Running Python 3.10.12 (main, May 27 2025, 17:12:29) [GCC 11.4.0] Running PyTorch 2.10.0.dev20251210+cu126 compiled for CUDA 12.6 Running Triton version 3.6.0 Mon Jan 26 02:20:09 2026 +-----------------------------------------------------------------------------------------+ | NVIDIA-SMI 570.148.08 Driver Version: 570.148.08 CUDA Version: 12.8 | |-----------------------------------------+------------------------+----------------------+ | GPU Name Persistence-M | Bus-Id Disp.A | Volatile Uncorr. ECC | | Fan Temp Perf Pwr:Usage/Cap | Memory-Usage | GPU-Util Compute M. | | | | MIG M. | |=========================================+========================+======================| | 0 NVIDIA H100 80GB HBM3 On | 00000000:61:00.0 Off | 0 | | N/A 35C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:62:00.0 Off | 0 | | N/A 38C P0 122W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 40C P0 122W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 35C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:6A:00.0 Off | 0 | | N/A 35C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 41C P0 130W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 39C P0 124W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 36C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 240581 C /usr/bin/python3 1510MiB | | 1 N/A N/A 240582 C /usr/bin/python3 1510MiB | | 2 N/A N/A 240583 C /usr/bin/python3 1510MiB | | 3 N/A N/A 240584 C /usr/bin/python3 1510MiB | | 4 N/A N/A 240585 C /usr/bin/python3 1510MiB | | 5 N/A N/A 240586 C /usr/bin/python3 1510MiB | | 6 N/A N/A 240587 C /usr/bin/python3 1510MiB | | 7 N/A N/A 240588 C /usr/bin/python3 1510MiB | +-----------------------------------------------------------------------------------------+ ==================================================================================================== Compiling model and warming up kernels (~7 minutes on first execution) Sampling steps [0, 1, 2, 511, 512, 513, 1023, 1024, 1025, 1534, 1535, 1536] for warmup Resetting Model step:0/1575 val_loss:10.8351 train_time:0ms step_avg:0.05ms step:1/1575 train_time:84ms step_avg:84.10ms step:2/1575 train_time:108ms step_avg:53.87ms step:3/1575 train_time:129ms step_avg:42.85ms step:4/1575 train_time:155ms step_avg:38.69ms step:5/1575 train_time:184ms step_avg:36.79ms step:6/1575 train_time:306ms step_avg:50.95ms step:7/1575 train_time:325ms step_avg:46.37ms step:8/1575 train_time:349ms step_avg:43.58ms step:9/1575 train_time:379ms step_avg:42.13ms step:10/1575 train_time:418ms step_avg:41.77ms step:11/1575 train_time:448ms step_avg:40.76ms step:12/1575 train_time:487ms step_avg:40.60ms step:13/1575 train_time:518ms step_avg:39.85ms step:14/1575 train_time:558ms step_avg:39.85ms step:15/1575 train_time:588ms step_avg:39.21ms step:16/1575 train_time:627ms step_avg:39.21ms step:17/1575 train_time:658ms step_avg:38.73ms step:18/1575 train_time:697ms step_avg:38.73ms step:19/1575 train_time:728ms step_avg:38.32ms step:20/1575 train_time:767ms step_avg:38.35ms step:21/1575 train_time:798ms step_avg:38.00ms step:22/1575 train_time:837ms step_avg:38.04ms step:23/1575 train_time:868ms step_avg:37.72ms step:24/1575 train_time:906ms step_avg:37.76ms step:25/1575 train_time:937ms step_avg:37.48ms step:26/1575 train_time:976ms step_avg:37.54ms step:27/1575 train_time:1007ms step_avg:37.30ms step:28/1575 train_time:1046ms step_avg:37.35ms step:29/1575 train_time:1077ms step_avg:37.13ms step:30/1575 train_time:1116ms step_avg:37.19ms step:31/1575 train_time:1147ms step_avg:36.99ms step:32/1575 train_time:1186ms step_avg:37.07ms step:33/1575 train_time:1217ms step_avg:36.87ms step:34/1575 train_time:1256ms step_avg:36.94ms step:35/1575 train_time:1287ms step_avg:36.78ms step:36/1575 train_time:1327ms step_avg:36.85ms step:37/1575 train_time:1358ms step_avg:36.69ms step:38/1575 train_time:1396ms step_avg:36.75ms step:39/1575 train_time:1427ms step_avg:36.59ms step:40/1575 train_time:1466ms step_avg:36.64ms step:41/1575 train_time:1497ms step_avg:36.50ms step:42/1575 train_time:1536ms step_avg:36.57ms step:43/1575 train_time:1567ms step_avg:36.44ms step:44/1575 train_time:1606ms step_avg:36.49ms step:45/1575 train_time:1637ms step_avg:36.38ms step:46/1575 train_time:1676ms step_avg:36.43ms step:47/1575 train_time:1707ms step_avg:36.31ms step:48/1575 train_time:1745ms step_avg:36.36ms step:49/1575 train_time:1776ms step_avg:36.25ms step:50/1575 train_time:1815ms step_avg:36.31ms step:51/1575 train_time:1846ms step_avg:36.19ms step:52/1575 train_time:1885ms step_avg:36.25ms step:53/1575 train_time:1916ms step_avg:36.15ms step:54/1575 train_time:1955ms step_avg:36.21ms step:55/1575 train_time:1986ms step_avg:36.11ms step:56/1575 train_time:2025ms step_avg:36.15ms step:57/1575 train_time:2055ms step_avg:36.06ms step:58/1575 train_time:2095ms step_avg:36.12ms step:59/1575 train_time:2125ms step_avg:36.03ms step:60/1575 train_time:2164ms step_avg:36.07ms step:61/1575 train_time:2195ms step_avg:35.98ms step:62/1575 train_time:2233ms step_avg:36.02ms step:63/1575 train_time:2264ms step_avg:35.94ms step:64/1575 train_time:2303ms step_avg:35.98ms step:65/1575 train_time:2334ms step_avg:35.91ms step:66/1575 train_time:2374ms step_avg:35.96ms step:67/1575 train_time:2404ms step_avg:35.88ms step:68/1575 train_time:2443ms step_avg:35.92ms step:69/1575 train_time:2473ms step_avg:35.84ms step:70/1575 train_time:2512ms step_avg:35.89ms step:71/1575 train_time:2543ms step_avg:35.81ms step:72/1575 train_time:2581ms step_avg:35.85ms step:73/1575 train_time:2612ms step_avg:35.78ms step:74/1575 train_time:2651ms step_avg:35.82ms step:75/1575 train_time:2682ms step_avg:35.75ms step:76/1575 train_time:2720ms step_avg:35.79ms step:77/1575 train_time:2751ms step_avg:35.73ms step:78/1575 train_time:2790ms step_avg:35.77ms step:79/1575 train_time:2820ms step_avg:35.70ms step:80/1575 train_time:2859ms step_avg:35.74ms step:81/1575 train_time:2890ms step_avg:35.68ms step:82/1575 train_time:2929ms step_avg:35.71ms step:83/1575 train_time:2959ms step_avg:35.65ms step:84/1575 train_time:2998ms step_avg:35.69ms step:85/1575 train_time:3029ms step_avg:35.63ms step:86/1575 train_time:3068ms step_avg:35.67ms step:87/1575 train_time:3099ms step_avg:35.62ms step:88/1575 train_time:3137ms step_avg:35.65ms step:89/1575 train_time:3168ms step_avg:35.60ms step:90/1575 train_time:3207ms step_avg:35.64ms step:91/1575 train_time:3238ms step_avg:35.58ms step:92/1575 train_time:3277ms step_avg:35.62ms step:93/1575 train_time:3308ms step_avg:35.57ms step:94/1575 train_time:3346ms step_avg:35.60ms step:95/1575 train_time:3377ms step_avg:35.55ms step:96/1575 train_time:3416ms step_avg:35.59ms step:97/1575 train_time:3447ms step_avg:35.54ms step:98/1575 train_time:3487ms step_avg:35.58ms step:99/1575 train_time:3517ms step_avg:35.52ms step:100/1575 train_time:3556ms step_avg:35.56ms step:101/1575 train_time:3586ms step_avg:35.50ms step:102/1575 train_time:3626ms step_avg:35.54ms step:103/1575 train_time:3656ms step_avg:35.49ms step:104/1575 train_time:3694ms step_avg:35.52ms step:105/1575 train_time:3725ms step_avg:35.48ms step:106/1575 train_time:3764ms step_avg:35.51ms step:107/1575 train_time:3795ms step_avg:35.46ms step:108/1575 train_time:3834ms step_avg:35.50ms step:109/1575 train_time:3864ms step_avg:35.45ms step:110/1575 train_time:3903ms step_avg:35.48ms step:111/1575 train_time:3934ms step_avg:35.44ms step:112/1575 train_time:3973ms step_avg:35.48ms step:113/1575 train_time:4004ms step_avg:35.43ms step:114/1575 train_time:4042ms step_avg:35.46ms step:115/1575 train_time:4073ms step_avg:35.42ms step:116/1575 train_time:4112ms step_avg:35.45ms step:117/1575 train_time:4143ms step_avg:35.41ms step:118/1575 train_time:4181ms step_avg:35.44ms step:119/1575 train_time:4212ms step_avg:35.39ms step:120/1575 train_time:4250ms step_avg:35.42ms step:121/1575 train_time:4281ms step_avg:35.38ms step:122/1575 train_time:4320ms step_avg:35.41ms step:123/1575 train_time:4351ms step_avg:35.37ms step:124/1575 train_time:4390ms step_avg:35.40ms step:125/1575 train_time:4420ms step_avg:35.36ms step:126/1575 train_time:4459ms step_avg:35.39ms step:127/1575 train_time:4490ms step_avg:35.36ms step:128/1575 train_time:4529ms step_avg:35.39ms step:129/1575 train_time:4561ms step_avg:35.35ms step:130/1575 train_time:4600ms step_avg:35.38ms step:131/1575 train_time:4630ms step_avg:35.35ms step:132/1575 train_time:4669ms step_avg:35.37ms step:133/1575 train_time:4700ms step_avg:35.34ms step:134/1575 train_time:4739ms step_avg:35.36ms step:135/1575 train_time:4770ms step_avg:35.33ms step:136/1575 train_time:4808ms step_avg:35.35ms step:137/1575 train_time:4839ms step_avg:35.32ms step:138/1575 train_time:4878ms step_avg:35.35ms step:139/1575 train_time:4909ms step_avg:35.31ms step:140/1575 train_time:4948ms step_avg:35.34ms step:141/1575 train_time:4979ms step_avg:35.31ms step:142/1575 train_time:5018ms step_avg:35.34ms step:143/1575 train_time:5049ms step_avg:35.31ms step:144/1575 train_time:5088ms step_avg:35.33ms step:145/1575 train_time:5119ms step_avg:35.30ms step:146/1575 train_time:5158ms step_avg:35.33ms step:147/1575 train_time:5189ms step_avg:35.30ms step:148/1575 train_time:5228ms step_avg:35.32ms step:149/1575 train_time:5258ms step_avg:35.29ms step:150/1575 train_time:5298ms step_avg:35.32ms step:151/1575 train_time:5328ms step_avg:35.29ms step:152/1575 train_time:5367ms step_avg:35.31ms step:153/1575 train_time:5398ms step_avg:35.28ms step:154/1575 train_time:5437ms step_avg:35.30ms step:155/1575 train_time:5468ms step_avg:35.27ms step:156/1575 train_time:5506ms step_avg:35.29ms step:157/1575 train_time:5537ms step_avg:35.27ms step:158/1575 train_time:5576ms step_avg:35.29ms step:159/1575 train_time:5606ms step_avg:35.26ms step:160/1575 train_time:5645ms step_avg:35.28ms step:161/1575 train_time:5676ms step_avg:35.25ms step:162/1575 train_time:5715ms step_avg:35.28ms step:163/1575 train_time:5746ms step_avg:35.25ms step:164/1575 train_time:5785ms step_avg:35.27ms step:165/1575 train_time:5815ms step_avg:35.24ms step:166/1575 train_time:5854ms step_avg:35.26ms step:167/1575 train_time:5885ms step_avg:35.24ms step:168/1575 train_time:5923ms step_avg:35.26ms step:169/1575 train_time:5954ms step_avg:35.23ms step:170/1575 train_time:5993ms step_avg:35.25ms step:171/1575 train_time:6024ms step_avg:35.23ms step:172/1575 train_time:6062ms step_avg:35.25ms step:173/1575 train_time:6093ms step_avg:35.22ms step:174/1575 train_time:6132ms step_avg:35.24ms step:175/1575 train_time:6163ms step_avg:35.21ms step:176/1575 train_time:6201ms step_avg:35.23ms step:177/1575 train_time:6231ms step_avg:35.21ms step:178/1575 train_time:6270ms step_avg:35.23ms step:179/1575 train_time:6301ms step_avg:35.20ms step:180/1575 train_time:6340ms step_avg:35.22ms step:181/1575 train_time:6370ms step_avg:35.20ms step:182/1575 train_time:6409ms step_avg:35.21ms step:183/1575 train_time:6440ms step_avg:35.19ms step:184/1575 train_time:6479ms step_avg:35.21ms step:185/1575 train_time:6509ms step_avg:35.18ms step:186/1575 train_time:6548ms step_avg:35.20ms step:187/1575 train_time:6579ms step_avg:35.18ms step:188/1575 train_time:6618ms step_avg:35.20ms step:189/1575 train_time:6648ms step_avg:35.18ms step:190/1575 train_time:6687ms step_avg:35.20ms step:191/1575 train_time:6718ms step_avg:35.17ms step:192/1575 train_time:6757ms step_avg:35.19ms step:193/1575 train_time:6787ms step_avg:35.17ms step:194/1575 train_time:6826ms step_avg:35.18ms step:195/1575 train_time:6857ms step_avg:35.16ms step:196/1575 train_time:6896ms step_avg:35.18ms step:197/1575 train_time:6926ms step_avg:35.16ms step:198/1575 train_time:6965ms step_avg:35.18ms step:199/1575 train_time:6996ms step_avg:35.16ms step:200/1575 train_time:7035ms step_avg:35.17ms step:201/1575 train_time:7066ms step_avg:35.15ms step:202/1575 train_time:7104ms step_avg:35.17ms step:203/1575 train_time:7135ms step_avg:35.15ms step:204/1575 train_time:7174ms step_avg:35.17ms step:205/1575 train_time:7205ms step_avg:35.14ms step:206/1575 train_time:7243ms step_avg:35.16ms step:207/1575 train_time:7274ms step_avg:35.14ms step:208/1575 train_time:7312ms step_avg:35.16ms step:209/1575 train_time:7343ms step_avg:35.13ms step:210/1575 train_time:7382ms step_avg:35.15ms step:211/1575 train_time:7412ms step_avg:35.13ms step:212/1575 train_time:7451ms step_avg:35.15ms step:213/1575 train_time:7482ms step_avg:35.12ms step:214/1575 train_time:7520ms step_avg:35.14ms step:215/1575 train_time:7551ms step_avg:35.12ms step:216/1575 train_time:7590ms step_avg:35.14ms step:217/1575 train_time:7621ms step_avg:35.12ms step:218/1575 train_time:7659ms step_avg:35.13ms step:219/1575 train_time:7690ms step_avg:35.11ms step:220/1575 train_time:7728ms step_avg:35.13ms step:221/1575 train_time:7759ms step_avg:35.11ms step:222/1575 train_time:7798ms step_avg:35.13ms step:223/1575 train_time:7828ms step_avg:35.11ms step:224/1575 train_time:7867ms step_avg:35.12ms step:225/1575 train_time:7898ms step_avg:35.10ms step:226/1575 train_time:7937ms step_avg:35.12ms step:227/1575 train_time:7967ms step_avg:35.10ms step:228/1575 train_time:8006ms step_avg:35.11ms step:229/1575 train_time:8037ms step_avg:35.10ms step:230/1575 train_time:8075ms step_avg:35.11ms step:231/1575 train_time:8106ms step_avg:35.09ms step:232/1575 train_time:8145ms step_avg:35.11ms step:233/1575 train_time:8176ms step_avg:35.09ms step:234/1575 train_time:8215ms step_avg:35.11ms step:235/1575 train_time:8246ms step_avg:35.09ms step:236/1575 train_time:8285ms step_avg:35.11ms step:237/1575 train_time:8315ms step_avg:35.09ms step:238/1575 train_time:8354ms step_avg:35.10ms step:239/1575 train_time:8385ms step_avg:35.08ms step:240/1575 train_time:8423ms step_avg:35.10ms step:241/1575 train_time:8454ms step_avg:35.08ms step:242/1575 train_time:8493ms step_avg:35.09ms step:243/1575 train_time:8523ms step_avg:35.08ms step:244/1575 train_time:8562ms step_avg:35.09ms step:245/1575 train_time:8593ms step_avg:35.07ms step:246/1575 train_time:8631ms step_avg:35.09ms step:247/1575 train_time:8662ms step_avg:35.07ms step:248/1575 train_time:8700ms step_avg:35.08ms step:249/1575 train_time:8731ms step_avg:35.06ms step:250/1575 train_time:8770ms step_avg:35.08ms step:250/1575 val_loss:4.5856 train_time:8818ms step_avg:35.27ms step:251/1575 train_time:8838ms step_avg:35.21ms step:252/1575 train_time:8858ms step_avg:35.15ms step:253/1575 train_time:8876ms step_avg:35.08ms step:254/1575 train_time:8911ms step_avg:35.08ms step:255/1575 train_time:8943ms step_avg:35.07ms step:256/1575 train_time:8984ms step_avg:35.09ms step:257/1575 train_time:9016ms step_avg:35.08ms step:258/1575 train_time:9055ms step_avg:35.10ms step:259/1575 train_time:9086ms step_avg:35.08ms step:260/1575 train_time:9125ms step_avg:35.10ms step:261/1575 train_time:9156ms step_avg:35.08ms step:262/1575 train_time:9195ms step_avg:35.09ms step:263/1575 train_time:9225ms step_avg:35.08ms step:264/1575 train_time:9264ms step_avg:35.09ms step:265/1575 train_time:9294ms step_avg:35.07ms step:266/1575 train_time:9333ms step_avg:35.09ms step:267/1575 train_time:9364ms step_avg:35.07ms step:268/1575 train_time:9402ms step_avg:35.08ms step:269/1575 train_time:9433ms step_avg:35.07ms step:270/1575 train_time:9472ms step_avg:35.08ms step:271/1575 train_time:9502ms step_avg:35.06ms step:272/1575 train_time:9542ms step_avg:35.08ms step:273/1575 train_time:9572ms step_avg:35.06ms step:274/1575 train_time:9610ms step_avg:35.07ms step:275/1575 train_time:9641ms step_avg:35.06ms step:276/1575 train_time:9680ms step_avg:35.07ms step:277/1575 train_time:9710ms step_avg:35.06ms step:278/1575 train_time:9749ms step_avg:35.07ms step:279/1575 train_time:9780ms step_avg:35.05ms step:280/1575 train_time:9818ms step_avg:35.07ms step:281/1575 train_time:9849ms step_avg:35.05ms step:282/1575 train_time:9887ms step_avg:35.06ms step:283/1575 train_time:9918ms step_avg:35.05ms step:284/1575 train_time:9957ms step_avg:35.06ms step:285/1575 train_time:9987ms step_avg:35.04ms step:286/1575 train_time:10026ms step_avg:35.05ms step:287/1575 train_time:10056ms step_avg:35.04ms step:288/1575 train_time:10095ms step_avg:35.05ms step:289/1575 train_time:10126ms step_avg:35.04ms step:290/1575 train_time:10165ms step_avg:35.05ms step:291/1575 train_time:10196ms step_avg:35.04ms step:292/1575 train_time:10235ms step_avg:35.05ms step:293/1575 train_time:10265ms step_avg:35.04ms step:294/1575 train_time:10304ms step_avg:35.05ms step:295/1575 train_time:10335ms step_avg:35.03ms step:296/1575 train_time:10373ms step_avg:35.04ms step:297/1575 train_time:10404ms step_avg:35.03ms step:298/1575 train_time:10443ms step_avg:35.05ms step:299/1575 train_time:10473ms step_avg:35.03ms step:300/1575 train_time:10512ms step_avg:35.04ms step:301/1575 train_time:10543ms step_avg:35.03ms step:302/1575 train_time:10582ms step_avg:35.04ms step:303/1575 train_time:10612ms step_avg:35.02ms step:304/1575 train_time:10651ms step_avg:35.04ms step:305/1575 train_time:10681ms step_avg:35.02ms step:306/1575 train_time:10720ms step_avg:35.03ms step:307/1575 train_time:10750ms step_avg:35.02ms step:308/1575 train_time:10789ms step_avg:35.03ms step:309/1575 train_time:10820ms step_avg:35.02ms step:310/1575 train_time:10859ms step_avg:35.03ms step:311/1575 train_time:10889ms step_avg:35.01ms step:312/1575 train_time:10928ms step_avg:35.03ms step:313/1575 train_time:10959ms step_avg:35.01ms step:314/1575 train_time:10997ms step_avg:35.02ms step:315/1575 train_time:11028ms step_avg:35.01ms step:316/1575 train_time:11066ms step_avg:35.02ms step:317/1575 train_time:11097ms step_avg:35.01ms step:318/1575 train_time:11135ms step_avg:35.02ms step:319/1575 train_time:11166ms step_avg:35.00ms step:320/1575 train_time:11205ms step_avg:35.02ms step:321/1575 train_time:11236ms step_avg:35.00ms step:322/1575 train_time:11274ms step_avg:35.01ms step:323/1575 train_time:11305ms step_avg:35.00ms step:324/1575 train_time:11343ms step_avg:35.01ms step:325/1575 train_time:11374ms step_avg:35.00ms step:326/1575 train_time:11413ms step_avg:35.01ms step:327/1575 train_time:11444ms step_avg:35.00ms step:328/1575 train_time:11483ms step_avg:35.01ms step:329/1575 train_time:11514ms step_avg:35.00ms step:330/1575 train_time:11552ms step_avg:35.01ms step:331/1575 train_time:11583ms step_avg:34.99ms step:332/1575 train_time:11622ms step_avg:35.01ms step:333/1575 train_time:11653ms step_avg:34.99ms step:334/1575 train_time:11691ms step_avg:35.00ms step:335/1575 train_time:11722ms step_avg:34.99ms step:336/1575 train_time:11761ms step_avg:35.00ms step:337/1575 train_time:11791ms step_avg:34.99ms step:338/1575 train_time:11830ms step_avg:35.00ms step:339/1575 train_time:11861ms step_avg:34.99ms step:340/1575 train_time:11899ms step_avg:35.00ms step:341/1575 train_time:11930ms step_avg:34.98ms step:342/1575 train_time:11968ms step_avg:34.99ms step:343/1575 train_time:11999ms step_avg:34.98ms step:344/1575 train_time:12038ms step_avg:34.99ms step:345/1575 train_time:12069ms step_avg:34.98ms step:346/1575 train_time:12107ms step_avg:34.99ms step:347/1575 train_time:12138ms step_avg:34.98ms step:348/1575 train_time:12177ms step_avg:34.99ms step:349/1575 train_time:12207ms step_avg:34.98ms step:350/1575 train_time:12246ms step_avg:34.99ms step:351/1575 train_time:12277ms step_avg:34.98ms step:352/1575 train_time:12316ms step_avg:34.99ms step:353/1575 train_time:12347ms step_avg:34.98ms step:354/1575 train_time:12386ms step_avg:34.99ms step:355/1575 train_time:12416ms step_avg:34.98ms step:356/1575 train_time:12455ms step_avg:34.99ms step:357/1575 train_time:12486ms step_avg:34.97ms step:358/1575 train_time:12524ms step_avg:34.98ms step:359/1575 train_time:12555ms step_avg:34.97ms step:360/1575 train_time:12594ms step_avg:34.98ms step:361/1575 train_time:12624ms step_avg:34.97ms step:362/1575 train_time:12663ms step_avg:34.98ms step:363/1575 train_time:12693ms step_avg:34.97ms step:364/1575 train_time:12732ms step_avg:34.98ms step:365/1575 train_time:12763ms step_avg:34.97ms step:366/1575 train_time:12801ms step_avg:34.98ms step:367/1575 train_time:12832ms step_avg:34.97ms step:368/1575 train_time:12872ms step_avg:34.98ms step:369/1575 train_time:12902ms step_avg:34.97ms step:370/1575 train_time:12941ms step_avg:34.98ms step:371/1575 train_time:12972ms step_avg:34.96ms step:372/1575 train_time:13010ms step_avg:34.97ms step:373/1575 train_time:13041ms step_avg:34.96ms step:374/1575 train_time:13080ms step_avg:34.97ms step:375/1575 train_time:13110ms step_avg:34.96ms step:376/1575 train_time:13149ms step_avg:34.97ms step:377/1575 train_time:13179ms step_avg:34.96ms step:378/1575 train_time:13218ms step_avg:34.97ms step:379/1575 train_time:13249ms step_avg:34.96ms step:380/1575 train_time:13287ms step_avg:34.97ms step:381/1575 train_time:13318ms step_avg:34.95ms step:382/1575 train_time:13357ms step_avg:34.97ms step:383/1575 train_time:13387ms step_avg:34.95ms step:384/1575 train_time:13426ms step_avg:34.96ms step:385/1575 train_time:13457ms step_avg:34.95ms step:386/1575 train_time:13496ms step_avg:34.96ms step:387/1575 train_time:13527ms step_avg:34.95ms step:388/1575 train_time:13565ms step_avg:34.96ms step:389/1575 train_time:13596ms step_avg:34.95ms step:390/1575 train_time:13635ms step_avg:34.96ms step:391/1575 train_time:13666ms step_avg:34.95ms step:392/1575 train_time:13704ms step_avg:34.96ms step:393/1575 train_time:13735ms step_avg:34.95ms step:394/1575 train_time:13774ms step_avg:34.96ms step:395/1575 train_time:13805ms step_avg:34.95ms step:396/1575 train_time:13843ms step_avg:34.96ms step:397/1575 train_time:13874ms step_avg:34.95ms step:398/1575 train_time:13913ms step_avg:34.96ms step:399/1575 train_time:13944ms step_avg:34.95ms step:400/1575 train_time:13982ms step_avg:34.96ms step:401/1575 train_time:14013ms step_avg:34.94ms step:402/1575 train_time:14051ms step_avg:34.95ms step:403/1575 train_time:14082ms step_avg:34.94ms step:404/1575 train_time:14121ms step_avg:34.95ms step:405/1575 train_time:14152ms step_avg:34.94ms step:406/1575 train_time:14190ms step_avg:34.95ms step:407/1575 train_time:14221ms step_avg:34.94ms step:408/1575 train_time:14260ms step_avg:34.95ms step:409/1575 train_time:14290ms step_avg:34.94ms step:410/1575 train_time:14329ms step_avg:34.95ms step:411/1575 train_time:14359ms step_avg:34.94ms step:412/1575 train_time:14398ms step_avg:34.95ms step:413/1575 train_time:14429ms step_avg:34.94ms step:414/1575 train_time:14468ms step_avg:34.95ms step:415/1575 train_time:14498ms step_avg:34.93ms step:416/1575 train_time:14537ms step_avg:34.94ms step:417/1575 train_time:14568ms step_avg:34.93ms step:418/1575 train_time:14606ms step_avg:34.94ms step:419/1575 train_time:14636ms step_avg:34.93ms step:420/1575 train_time:14676ms step_avg:34.94ms step:421/1575 train_time:14706ms step_avg:34.93ms step:422/1575 train_time:14745ms step_avg:34.94ms step:423/1575 train_time:14775ms step_avg:34.93ms step:424/1575 train_time:14814ms step_avg:34.94ms step:425/1575 train_time:14845ms step_avg:34.93ms step:426/1575 train_time:14883ms step_avg:34.94ms step:427/1575 train_time:14914ms step_avg:34.93ms step:428/1575 train_time:14953ms step_avg:34.94ms step:429/1575 train_time:14983ms step_avg:34.93ms step:430/1575 train_time:15022ms step_avg:34.93ms step:431/1575 train_time:15052ms step_avg:34.92ms step:432/1575 train_time:15091ms step_avg:34.93ms step:433/1575 train_time:15122ms step_avg:34.92ms step:434/1575 train_time:15161ms step_avg:34.93ms step:435/1575 train_time:15192ms step_avg:34.92ms step:436/1575 train_time:15230ms step_avg:34.93ms step:437/1575 train_time:15261ms step_avg:34.92ms step:438/1575 train_time:15300ms step_avg:34.93ms step:439/1575 train_time:15331ms step_avg:34.92ms step:440/1575 train_time:15369ms step_avg:34.93ms step:441/1575 train_time:15400ms step_avg:34.92ms step:442/1575 train_time:15439ms step_avg:34.93ms step:443/1575 train_time:15469ms step_avg:34.92ms step:444/1575 train_time:15507ms step_avg:34.93ms step:445/1575 train_time:15538ms step_avg:34.92ms step:446/1575 train_time:15577ms step_avg:34.93ms step:447/1575 train_time:15607ms step_avg:34.92ms step:448/1575 train_time:15645ms step_avg:34.92ms step:449/1575 train_time:15676ms step_avg:34.91ms step:450/1575 train_time:15715ms step_avg:34.92ms step:451/1575 train_time:15746ms step_avg:34.91ms step:452/1575 train_time:15785ms step_avg:34.92ms step:453/1575 train_time:15816ms step_avg:34.91ms step:454/1575 train_time:15854ms step_avg:34.92ms step:455/1575 train_time:15885ms step_avg:34.91ms step:456/1575 train_time:15924ms step_avg:34.92ms step:457/1575 train_time:15955ms step_avg:34.91ms step:458/1575 train_time:15993ms step_avg:34.92ms step:459/1575 train_time:16024ms step_avg:34.91ms step:460/1575 train_time:16063ms step_avg:34.92ms step:461/1575 train_time:16093ms step_avg:34.91ms step:462/1575 train_time:16132ms step_avg:34.92ms step:463/1575 train_time:16163ms step_avg:34.91ms step:464/1575 train_time:16201ms step_avg:34.92ms step:465/1575 train_time:16232ms step_avg:34.91ms step:466/1575 train_time:16270ms step_avg:34.92ms step:467/1575 train_time:16301ms step_avg:34.91ms step:468/1575 train_time:16340ms step_avg:34.91ms step:469/1575 train_time:16371ms step_avg:34.91ms step:470/1575 train_time:16409ms step_avg:34.91ms step:471/1575 train_time:16440ms step_avg:34.90ms step:472/1575 train_time:16479ms step_avg:34.91ms step:473/1575 train_time:16509ms step_avg:34.90ms step:474/1575 train_time:16548ms step_avg:34.91ms step:475/1575 train_time:16578ms step_avg:34.90ms step:476/1575 train_time:16617ms step_avg:34.91ms step:477/1575 train_time:16648ms step_avg:34.90ms step:478/1575 train_time:16686ms step_avg:34.91ms step:479/1575 train_time:16717ms step_avg:34.90ms step:480/1575 train_time:16756ms step_avg:34.91ms step:481/1575 train_time:16786ms step_avg:34.90ms step:482/1575 train_time:16825ms step_avg:34.91ms step:483/1575 train_time:16856ms step_avg:34.90ms step:484/1575 train_time:16894ms step_avg:34.90ms step:485/1575 train_time:16925ms step_avg:34.90ms step:486/1575 train_time:16964ms step_avg:34.90ms step:487/1575 train_time:16994ms step_avg:34.90ms step:488/1575 train_time:17033ms step_avg:34.90ms step:489/1575 train_time:17064ms step_avg:34.90ms step:490/1575 train_time:17103ms step_avg:34.90ms step:491/1575 train_time:17134ms step_avg:34.90ms step:492/1575 train_time:17173ms step_avg:34.90ms step:493/1575 train_time:17203ms step_avg:34.90ms step:494/1575 train_time:17242ms step_avg:34.90ms step:495/1575 train_time:17273ms step_avg:34.89ms step:496/1575 train_time:17312ms step_avg:34.90ms step:497/1575 train_time:17343ms step_avg:34.89ms step:498/1575 train_time:17381ms step_avg:34.90ms step:499/1575 train_time:17412ms step_avg:34.89ms step:500/1575 train_time:17451ms step_avg:34.90ms step:500/1575 val_loss:4.2306 train_time:17500ms step_avg:35.00ms step:501/1575 train_time:17520ms step_avg:34.97ms step:502/1575 train_time:17540ms step_avg:34.94ms step:503/1575 train_time:17558ms step_avg:34.91ms step:504/1575 train_time:17594ms step_avg:34.91ms step:505/1575 train_time:17625ms step_avg:34.90ms step:506/1575 train_time:17666ms step_avg:34.91ms step:507/1575 train_time:17697ms step_avg:34.90ms step:508/1575 train_time:17736ms step_avg:34.91ms step:509/1575 train_time:17766ms step_avg:34.90ms step:510/1575 train_time:17806ms step_avg:34.91ms step:511/1575 train_time:17836ms step_avg:34.90ms step:512/1575 train_time:17874ms step_avg:34.91ms step:513/1575 train_time:17946ms step_avg:34.98ms step:514/1575 train_time:18004ms step_avg:35.03ms step:515/1575 train_time:18067ms step_avg:35.08ms step:516/1575 train_time:18125ms step_avg:35.13ms step:517/1575 train_time:18188ms step_avg:35.18ms step:518/1575 train_time:18248ms step_avg:35.23ms step:519/1575 train_time:18311ms step_avg:35.28ms step:520/1575 train_time:18370ms step_avg:35.33ms step:521/1575 train_time:18432ms step_avg:35.38ms step:522/1575 train_time:18492ms step_avg:35.43ms step:523/1575 train_time:18558ms step_avg:35.48ms step:524/1575 train_time:18618ms step_avg:35.53ms step:525/1575 train_time:18682ms step_avg:35.58ms step:526/1575 train_time:18742ms step_avg:35.63ms step:527/1575 train_time:18806ms step_avg:35.69ms step:528/1575 train_time:18866ms step_avg:35.73ms step:529/1575 train_time:18932ms step_avg:35.79ms step:530/1575 train_time:18991ms step_avg:35.83ms step:531/1575 train_time:19052ms step_avg:35.88ms step:532/1575 train_time:19111ms step_avg:35.92ms step:533/1575 train_time:19175ms step_avg:35.98ms step:534/1575 train_time:19234ms step_avg:36.02ms step:535/1575 train_time:19297ms step_avg:36.07ms step:536/1575 train_time:19356ms step_avg:36.11ms step:537/1575 train_time:19419ms step_avg:36.16ms step:538/1575 train_time:19478ms step_avg:36.20ms step:539/1575 train_time:19541ms step_avg:36.25ms step:540/1575 train_time:19602ms step_avg:36.30ms step:541/1575 train_time:19667ms step_avg:36.35ms step:542/1575 train_time:19726ms step_avg:36.39ms step:543/1575 train_time:19788ms step_avg:36.44ms step:544/1575 train_time:19847ms step_avg:36.48ms step:545/1575 train_time:19912ms step_avg:36.54ms step:546/1575 train_time:19972ms step_avg:36.58ms step:547/1575 train_time:20035ms step_avg:36.63ms step:548/1575 train_time:20094ms step_avg:36.67ms step:549/1575 train_time:20157ms step_avg:36.72ms step:550/1575 train_time:20216ms step_avg:36.76ms step:551/1575 train_time:20279ms step_avg:36.80ms step:552/1575 train_time:20338ms step_avg:36.84ms step:553/1575 train_time:20402ms step_avg:36.89ms step:554/1575 train_time:20461ms step_avg:36.93ms step:555/1575 train_time:20524ms step_avg:36.98ms step:556/1575 train_time:20583ms step_avg:37.02ms step:557/1575 train_time:20648ms step_avg:37.07ms step:558/1575 train_time:20708ms step_avg:37.11ms step:559/1575 train_time:20771ms step_avg:37.16ms step:560/1575 train_time:20830ms step_avg:37.20ms step:561/1575 train_time:20894ms step_avg:37.24ms step:562/1575 train_time:20954ms step_avg:37.28ms step:563/1575 train_time:21017ms step_avg:37.33ms step:564/1575 train_time:21076ms step_avg:37.37ms step:565/1575 train_time:21139ms step_avg:37.41ms step:566/1575 train_time:21198ms step_avg:37.45ms step:567/1575 train_time:21262ms step_avg:37.50ms step:568/1575 train_time:21321ms step_avg:37.54ms step:569/1575 train_time:21389ms step_avg:37.59ms step:570/1575 train_time:21446ms step_avg:37.62ms step:571/1575 train_time:21510ms step_avg:37.67ms step:572/1575 train_time:21568ms step_avg:37.71ms step:573/1575 train_time:21632ms step_avg:37.75ms step:574/1575 train_time:21690ms step_avg:37.79ms step:575/1575 train_time:21754ms step_avg:37.83ms step:576/1575 train_time:21813ms step_avg:37.87ms step:577/1575 train_time:21876ms step_avg:37.91ms step:578/1575 train_time:21936ms step_avg:37.95ms step:579/1575 train_time:21999ms step_avg:37.99ms step:580/1575 train_time:22058ms step_avg:38.03ms step:581/1575 train_time:22121ms step_avg:38.07ms step:582/1575 train_time:22180ms step_avg:38.11ms step:583/1575 train_time:22243ms step_avg:38.15ms step:584/1575 train_time:22302ms step_avg:38.19ms step:585/1575 train_time:22365ms step_avg:38.23ms step:586/1575 train_time:22424ms step_avg:38.27ms step:587/1575 train_time:22487ms step_avg:38.31ms step:588/1575 train_time:22546ms step_avg:38.34ms step:589/1575 train_time:22610ms step_avg:38.39ms step:590/1575 train_time:22671ms step_avg:38.43ms step:591/1575 train_time:22734ms step_avg:38.47ms step:592/1575 train_time:22793ms step_avg:38.50ms step:593/1575 train_time:22857ms step_avg:38.54ms step:594/1575 train_time:22916ms step_avg:38.58ms step:595/1575 train_time:22979ms step_avg:38.62ms step:596/1575 train_time:23038ms step_avg:38.65ms step:597/1575 train_time:23101ms step_avg:38.70ms step:598/1575 train_time:23161ms step_avg:38.73ms step:599/1575 train_time:23223ms step_avg:38.77ms step:600/1575 train_time:23282ms step_avg:38.80ms step:601/1575 train_time:23345ms step_avg:38.84ms step:602/1575 train_time:23404ms step_avg:38.88ms step:603/1575 train_time:23467ms step_avg:38.92ms step:604/1575 train_time:23527ms step_avg:38.95ms step:605/1575 train_time:23590ms step_avg:38.99ms step:606/1575 train_time:23650ms step_avg:39.03ms step:607/1575 train_time:23714ms step_avg:39.07ms step:608/1575 train_time:23773ms step_avg:39.10ms step:609/1575 train_time:23837ms step_avg:39.14ms step:610/1575 train_time:23896ms step_avg:39.17ms step:611/1575 train_time:23960ms step_avg:39.21ms step:612/1575 train_time:24020ms step_avg:39.25ms step:613/1575 train_time:24085ms step_avg:39.29ms step:614/1575 train_time:24144ms step_avg:39.32ms step:615/1575 train_time:24206ms step_avg:39.36ms step:616/1575 train_time:24266ms step_avg:39.39ms step:617/1575 train_time:24328ms step_avg:39.43ms step:618/1575 train_time:24388ms step_avg:39.46ms step:619/1575 train_time:24451ms step_avg:39.50ms step:620/1575 train_time:24510ms step_avg:39.53ms step:621/1575 train_time:24574ms step_avg:39.57ms step:622/1575 train_time:24634ms step_avg:39.60ms step:623/1575 train_time:24697ms step_avg:39.64ms step:624/1575 train_time:24757ms step_avg:39.67ms step:625/1575 train_time:24818ms step_avg:39.71ms step:626/1575 train_time:24878ms step_avg:39.74ms step:627/1575 train_time:24941ms step_avg:39.78ms step:628/1575 train_time:25000ms step_avg:39.81ms step:629/1575 train_time:25063ms step_avg:39.85ms step:630/1575 train_time:25122ms step_avg:39.88ms step:631/1575 train_time:25185ms step_avg:39.91ms step:632/1575 train_time:25244ms step_avg:39.94ms step:633/1575 train_time:25307ms step_avg:39.98ms step:634/1575 train_time:25367ms step_avg:40.01ms step:635/1575 train_time:25430ms step_avg:40.05ms step:636/1575 train_time:25489ms step_avg:40.08ms step:637/1575 train_time:25554ms step_avg:40.12ms step:638/1575 train_time:25612ms step_avg:40.14ms step:639/1575 train_time:25677ms step_avg:40.18ms step:640/1575 train_time:25735ms step_avg:40.21ms step:641/1575 train_time:25798ms step_avg:40.25ms step:642/1575 train_time:25858ms step_avg:40.28ms step:643/1575 train_time:25921ms step_avg:40.31ms step:644/1575 train_time:25980ms step_avg:40.34ms step:645/1575 train_time:26043ms step_avg:40.38ms step:646/1575 train_time:26102ms step_avg:40.41ms step:647/1575 train_time:26166ms step_avg:40.44ms step:648/1575 train_time:26225ms step_avg:40.47ms step:649/1575 train_time:26288ms step_avg:40.51ms step:650/1575 train_time:26347ms step_avg:40.53ms step:651/1575 train_time:26410ms step_avg:40.57ms step:652/1575 train_time:26469ms step_avg:40.60ms step:653/1575 train_time:26533ms step_avg:40.63ms step:654/1575 train_time:26592ms step_avg:40.66ms step:655/1575 train_time:26656ms step_avg:40.70ms step:656/1575 train_time:26715ms step_avg:40.72ms step:657/1575 train_time:26779ms step_avg:40.76ms step:658/1575 train_time:26839ms step_avg:40.79ms step:659/1575 train_time:26902ms step_avg:40.82ms step:660/1575 train_time:26961ms step_avg:40.85ms step:661/1575 train_time:27026ms step_avg:40.89ms step:662/1575 train_time:27084ms step_avg:40.91ms step:663/1575 train_time:27147ms step_avg:40.95ms step:664/1575 train_time:27206ms step_avg:40.97ms step:665/1575 train_time:27269ms step_avg:41.01ms step:666/1575 train_time:27328ms step_avg:41.03ms step:667/1575 train_time:27392ms step_avg:41.07ms step:668/1575 train_time:27451ms step_avg:41.09ms step:669/1575 train_time:27514ms step_avg:41.13ms step:670/1575 train_time:27573ms step_avg:41.15ms step:671/1575 train_time:27637ms step_avg:41.19ms step:672/1575 train_time:27697ms step_avg:41.22ms step:673/1575 train_time:27760ms step_avg:41.25ms step:674/1575 train_time:27819ms step_avg:41.27ms step:675/1575 train_time:27882ms step_avg:41.31ms step:676/1575 train_time:27942ms step_avg:41.33ms step:677/1575 train_time:28005ms step_avg:41.37ms step:678/1575 train_time:28063ms step_avg:41.39ms step:679/1575 train_time:28126ms step_avg:41.42ms step:680/1575 train_time:28186ms step_avg:41.45ms step:681/1575 train_time:28249ms step_avg:41.48ms step:682/1575 train_time:28308ms step_avg:41.51ms step:683/1575 train_time:28372ms step_avg:41.54ms step:684/1575 train_time:28431ms step_avg:41.57ms step:685/1575 train_time:28495ms step_avg:41.60ms step:686/1575 train_time:28554ms step_avg:41.62ms step:687/1575 train_time:28618ms step_avg:41.66ms step:688/1575 train_time:28678ms step_avg:41.68ms step:689/1575 train_time:28741ms step_avg:41.71ms step:690/1575 train_time:28800ms step_avg:41.74ms step:691/1575 train_time:28864ms step_avg:41.77ms step:692/1575 train_time:28923ms step_avg:41.80ms step:693/1575 train_time:28987ms step_avg:41.83ms step:694/1575 train_time:29046ms step_avg:41.85ms step:695/1575 train_time:29109ms step_avg:41.88ms step:696/1575 train_time:29170ms step_avg:41.91ms step:697/1575 train_time:29232ms step_avg:41.94ms step:698/1575 train_time:29290ms step_avg:41.96ms step:699/1575 train_time:29354ms step_avg:41.99ms step:700/1575 train_time:29415ms step_avg:42.02ms step:701/1575 train_time:29478ms step_avg:42.05ms step:702/1575 train_time:29537ms step_avg:42.08ms step:703/1575 train_time:29600ms step_avg:42.11ms step:704/1575 train_time:29659ms step_avg:42.13ms step:705/1575 train_time:29724ms step_avg:42.16ms step:706/1575 train_time:29783ms step_avg:42.19ms step:707/1575 train_time:29847ms step_avg:42.22ms step:708/1575 train_time:29905ms step_avg:42.24ms step:709/1575 train_time:29969ms step_avg:42.27ms step:710/1575 train_time:30028ms step_avg:42.29ms step:711/1575 train_time:30091ms step_avg:42.32ms step:712/1575 train_time:30150ms step_avg:42.35ms step:713/1575 train_time:30214ms step_avg:42.38ms step:714/1575 train_time:30273ms step_avg:42.40ms step:715/1575 train_time:30336ms step_avg:42.43ms step:716/1575 train_time:30396ms step_avg:42.45ms step:717/1575 train_time:30459ms step_avg:42.48ms step:718/1575 train_time:30519ms step_avg:42.50ms step:719/1575 train_time:30582ms step_avg:42.53ms step:720/1575 train_time:30641ms step_avg:42.56ms step:721/1575 train_time:30706ms step_avg:42.59ms step:722/1575 train_time:30765ms step_avg:42.61ms step:723/1575 train_time:30829ms step_avg:42.64ms step:724/1575 train_time:30888ms step_avg:42.66ms step:725/1575 train_time:30952ms step_avg:42.69ms step:726/1575 train_time:31011ms step_avg:42.72ms step:727/1575 train_time:31073ms step_avg:42.74ms step:728/1575 train_time:31133ms step_avg:42.77ms step:729/1575 train_time:31197ms step_avg:42.79ms step:730/1575 train_time:31256ms step_avg:42.82ms step:731/1575 train_time:31318ms step_avg:42.84ms step:732/1575 train_time:31378ms step_avg:42.87ms step:733/1575 train_time:31441ms step_avg:42.89ms step:734/1575 train_time:31500ms step_avg:42.92ms step:735/1575 train_time:31563ms step_avg:42.94ms step:736/1575 train_time:31622ms step_avg:42.96ms step:737/1575 train_time:31686ms step_avg:42.99ms step:738/1575 train_time:31746ms step_avg:43.02ms step:739/1575 train_time:31809ms step_avg:43.04ms step:740/1575 train_time:31868ms step_avg:43.07ms step:741/1575 train_time:31932ms step_avg:43.09ms step:742/1575 train_time:31992ms step_avg:43.12ms step:743/1575 train_time:32055ms step_avg:43.14ms step:744/1575 train_time:32114ms step_avg:43.16ms step:745/1575 train_time:32178ms step_avg:43.19ms step:746/1575 train_time:32237ms step_avg:43.21ms step:747/1575 train_time:32300ms step_avg:43.24ms step:748/1575 train_time:32359ms step_avg:43.26ms step:749/1575 train_time:32422ms step_avg:43.29ms step:750/1575 train_time:32481ms step_avg:43.31ms step:750/1575 val_loss:3.8752 train_time:32528ms step_avg:43.37ms step:751/1575 train_time:32549ms step_avg:43.34ms step:752/1575 train_time:32608ms step_avg:43.36ms step:753/1575 train_time:32673ms step_avg:43.39ms step:754/1575 train_time:32735ms step_avg:43.41ms step:755/1575 train_time:32798ms step_avg:43.44ms step:756/1575 train_time:32857ms step_avg:43.46ms step:757/1575 train_time:32920ms step_avg:43.49ms step:758/1575 train_time:32979ms step_avg:43.51ms step:759/1575 train_time:33042ms step_avg:43.53ms step:760/1575 train_time:33100ms step_avg:43.55ms step:761/1575 train_time:33163ms step_avg:43.58ms step:762/1575 train_time:33222ms step_avg:43.60ms step:763/1575 train_time:33285ms step_avg:43.62ms step:764/1575 train_time:33343ms step_avg:43.64ms step:765/1575 train_time:33407ms step_avg:43.67ms step:766/1575 train_time:33469ms step_avg:43.69ms step:767/1575 train_time:33532ms step_avg:43.72ms step:768/1575 train_time:33591ms step_avg:43.74ms step:769/1575 train_time:33656ms step_avg:43.77ms step:770/1575 train_time:33715ms step_avg:43.79ms step:771/1575 train_time:33778ms step_avg:43.81ms step:772/1575 train_time:33839ms step_avg:43.83ms step:773/1575 train_time:33902ms step_avg:43.86ms step:774/1575 train_time:33961ms step_avg:43.88ms step:775/1575 train_time:34024ms step_avg:43.90ms step:776/1575 train_time:34083ms step_avg:43.92ms step:777/1575 train_time:34146ms step_avg:43.95ms step:778/1575 train_time:34205ms step_avg:43.96ms step:779/1575 train_time:34267ms step_avg:43.99ms step:780/1575 train_time:34326ms step_avg:44.01ms step:781/1575 train_time:34389ms step_avg:44.03ms step:782/1575 train_time:34448ms step_avg:44.05ms step:783/1575 train_time:34511ms step_avg:44.08ms step:784/1575 train_time:34572ms step_avg:44.10ms step:785/1575 train_time:34636ms step_avg:44.12ms step:786/1575 train_time:34695ms step_avg:44.14ms step:787/1575 train_time:34758ms step_avg:44.17ms step:788/1575 train_time:34819ms step_avg:44.19ms step:789/1575 train_time:34882ms step_avg:44.21ms step:790/1575 train_time:34941ms step_avg:44.23ms step:791/1575 train_time:35004ms step_avg:44.25ms step:792/1575 train_time:35063ms step_avg:44.27ms step:793/1575 train_time:35127ms step_avg:44.30ms step:794/1575 train_time:35186ms step_avg:44.31ms step:795/1575 train_time:35248ms step_avg:44.34ms step:796/1575 train_time:35307ms step_avg:44.36ms step:797/1575 train_time:35370ms step_avg:44.38ms step:798/1575 train_time:35429ms step_avg:44.40ms step:799/1575 train_time:35492ms step_avg:44.42ms step:800/1575 train_time:35551ms step_avg:44.44ms step:801/1575 train_time:35615ms step_avg:44.46ms step:802/1575 train_time:35675ms step_avg:44.48ms step:803/1575 train_time:35738ms step_avg:44.51ms step:804/1575 train_time:35797ms step_avg:44.52ms step:805/1575 train_time:35862ms step_avg:44.55ms step:806/1575 train_time:35923ms step_avg:44.57ms step:807/1575 train_time:35985ms step_avg:44.59ms step:808/1575 train_time:36045ms step_avg:44.61ms step:809/1575 train_time:36108ms step_avg:44.63ms step:810/1575 train_time:36167ms step_avg:44.65ms step:811/1575 train_time:36230ms step_avg:44.67ms step:812/1575 train_time:36290ms step_avg:44.69ms step:813/1575 train_time:36352ms step_avg:44.71ms step:814/1575 train_time:36412ms step_avg:44.73ms step:815/1575 train_time:36475ms step_avg:44.75ms step:816/1575 train_time:36534ms step_avg:44.77ms step:817/1575 train_time:36598ms step_avg:44.80ms step:818/1575 train_time:36658ms step_avg:44.81ms step:819/1575 train_time:36721ms step_avg:44.84ms step:820/1575 train_time:36780ms step_avg:44.85ms step:821/1575 train_time:36844ms step_avg:44.88ms step:822/1575 train_time:36903ms step_avg:44.89ms step:823/1575 train_time:36966ms step_avg:44.92ms step:824/1575 train_time:37026ms step_avg:44.93ms step:825/1575 train_time:37089ms step_avg:44.96ms step:826/1575 train_time:37148ms step_avg:44.97ms step:827/1575 train_time:37211ms step_avg:44.99ms step:828/1575 train_time:37270ms step_avg:45.01ms step:829/1575 train_time:37333ms step_avg:45.03ms step:830/1575 train_time:37392ms step_avg:45.05ms step:831/1575 train_time:37455ms step_avg:45.07ms step:832/1575 train_time:37514ms step_avg:45.09ms step:833/1575 train_time:37578ms step_avg:45.11ms step:834/1575 train_time:37637ms step_avg:45.13ms step:835/1575 train_time:37701ms step_avg:45.15ms step:836/1575 train_time:37760ms step_avg:45.17ms step:837/1575 train_time:37825ms step_avg:45.19ms step:838/1575 train_time:37884ms step_avg:45.21ms step:839/1575 train_time:37948ms step_avg:45.23ms step:840/1575 train_time:38007ms step_avg:45.25ms step:841/1575 train_time:38070ms step_avg:45.27ms step:842/1575 train_time:38130ms step_avg:45.28ms step:843/1575 train_time:38193ms step_avg:45.31ms step:844/1575 train_time:38252ms step_avg:45.32ms step:845/1575 train_time:38318ms step_avg:45.35ms step:846/1575 train_time:38376ms step_avg:45.36ms step:847/1575 train_time:38439ms step_avg:45.38ms step:848/1575 train_time:38497ms step_avg:45.40ms step:849/1575 train_time:38562ms step_avg:45.42ms step:850/1575 train_time:38622ms step_avg:45.44ms step:851/1575 train_time:38684ms step_avg:45.46ms step:852/1575 train_time:38745ms step_avg:45.48ms step:853/1575 train_time:38808ms step_avg:45.50ms step:854/1575 train_time:38866ms step_avg:45.51ms step:855/1575 train_time:38929ms step_avg:45.53ms step:856/1575 train_time:38988ms step_avg:45.55ms step:857/1575 train_time:39051ms step_avg:45.57ms step:858/1575 train_time:39111ms step_avg:45.58ms step:859/1575 train_time:39174ms step_avg:45.60ms step:860/1575 train_time:39233ms step_avg:45.62ms step:861/1575 train_time:39296ms step_avg:45.64ms step:862/1575 train_time:39355ms step_avg:45.66ms step:863/1575 train_time:39418ms step_avg:45.68ms step:864/1575 train_time:39477ms step_avg:45.69ms step:865/1575 train_time:39540ms step_avg:45.71ms step:866/1575 train_time:39600ms step_avg:45.73ms step:867/1575 train_time:39663ms step_avg:45.75ms step:868/1575 train_time:39722ms step_avg:45.76ms step:869/1575 train_time:39786ms step_avg:45.78ms step:870/1575 train_time:39846ms step_avg:45.80ms step:871/1575 train_time:39910ms step_avg:45.82ms step:872/1575 train_time:39969ms step_avg:45.84ms step:873/1575 train_time:40033ms step_avg:45.86ms step:874/1575 train_time:40092ms step_avg:45.87ms step:875/1575 train_time:40155ms step_avg:45.89ms step:876/1575 train_time:40214ms step_avg:45.91ms step:877/1575 train_time:40278ms step_avg:45.93ms step:878/1575 train_time:40337ms step_avg:45.94ms step:879/1575 train_time:40401ms step_avg:45.96ms step:880/1575 train_time:40459ms step_avg:45.98ms step:881/1575 train_time:40522ms step_avg:46.00ms step:882/1575 train_time:40582ms step_avg:46.01ms step:883/1575 train_time:40645ms step_avg:46.03ms step:884/1575 train_time:40704ms step_avg:46.05ms step:885/1575 train_time:40768ms step_avg:46.07ms step:886/1575 train_time:40832ms step_avg:46.09ms step:887/1575 train_time:40892ms step_avg:46.10ms step:888/1575 train_time:40952ms step_avg:46.12ms step:889/1575 train_time:41014ms step_avg:46.13ms step:890/1575 train_time:41072ms step_avg:46.15ms step:891/1575 train_time:41136ms step_avg:46.17ms step:892/1575 train_time:41195ms step_avg:46.18ms step:893/1575 train_time:41259ms step_avg:46.20ms step:894/1575 train_time:41319ms step_avg:46.22ms step:895/1575 train_time:41381ms step_avg:46.24ms step:896/1575 train_time:41440ms step_avg:46.25ms step:897/1575 train_time:41503ms step_avg:46.27ms step:898/1575 train_time:41563ms step_avg:46.28ms step:899/1575 train_time:41626ms step_avg:46.30ms step:900/1575 train_time:41685ms step_avg:46.32ms step:901/1575 train_time:41748ms step_avg:46.34ms step:902/1575 train_time:41808ms step_avg:46.35ms step:903/1575 train_time:41871ms step_avg:46.37ms step:904/1575 train_time:41931ms step_avg:46.38ms step:905/1575 train_time:41994ms step_avg:46.40ms step:906/1575 train_time:42053ms step_avg:46.42ms step:907/1575 train_time:42116ms step_avg:46.43ms step:908/1575 train_time:42175ms step_avg:46.45ms step:909/1575 train_time:42239ms step_avg:46.47ms step:910/1575 train_time:42298ms step_avg:46.48ms step:911/1575 train_time:42361ms step_avg:46.50ms step:912/1575 train_time:42421ms step_avg:46.51ms step:913/1575 train_time:42483ms step_avg:46.53ms step:914/1575 train_time:42543ms step_avg:46.55ms step:915/1575 train_time:42607ms step_avg:46.57ms step:916/1575 train_time:42666ms step_avg:46.58ms step:917/1575 train_time:42730ms step_avg:46.60ms step:918/1575 train_time:42790ms step_avg:46.61ms step:919/1575 train_time:42853ms step_avg:46.63ms step:920/1575 train_time:42912ms step_avg:46.64ms step:921/1575 train_time:42976ms step_avg:46.66ms step:922/1575 train_time:43035ms step_avg:46.68ms step:923/1575 train_time:43098ms step_avg:46.69ms step:924/1575 train_time:43157ms step_avg:46.71ms step:925/1575 train_time:43220ms step_avg:46.72ms step:926/1575 train_time:43280ms step_avg:46.74ms step:927/1575 train_time:43343ms step_avg:46.76ms step:928/1575 train_time:43403ms step_avg:46.77ms step:929/1575 train_time:43465ms step_avg:46.79ms step:930/1575 train_time:43524ms step_avg:46.80ms step:931/1575 train_time:43589ms step_avg:46.82ms step:932/1575 train_time:43648ms step_avg:46.83ms step:933/1575 train_time:43712ms step_avg:46.85ms step:934/1575 train_time:43772ms step_avg:46.86ms step:935/1575 train_time:43833ms step_avg:46.88ms step:936/1575 train_time:43892ms step_avg:46.89ms step:937/1575 train_time:43957ms step_avg:46.91ms step:938/1575 train_time:44016ms step_avg:46.93ms step:939/1575 train_time:44079ms step_avg:46.94ms step:940/1575 train_time:44139ms step_avg:46.96ms step:941/1575 train_time:44201ms step_avg:46.97ms step:942/1575 train_time:44260ms step_avg:46.99ms step:943/1575 train_time:44324ms step_avg:47.00ms step:944/1575 train_time:44383ms step_avg:47.02ms step:945/1575 train_time:44446ms step_avg:47.03ms step:946/1575 train_time:44505ms step_avg:47.05ms step:947/1575 train_time:44569ms step_avg:47.06ms step:948/1575 train_time:44629ms step_avg:47.08ms step:949/1575 train_time:44691ms step_avg:47.09ms step:950/1575 train_time:44750ms step_avg:47.11ms step:951/1575 train_time:44814ms step_avg:47.12ms step:952/1575 train_time:44872ms step_avg:47.13ms step:953/1575 train_time:44935ms step_avg:47.15ms step:954/1575 train_time:44994ms step_avg:47.16ms step:955/1575 train_time:45058ms step_avg:47.18ms step:956/1575 train_time:45117ms step_avg:47.19ms step:957/1575 train_time:45182ms step_avg:47.21ms step:958/1575 train_time:45241ms step_avg:47.22ms step:959/1575 train_time:45305ms step_avg:47.24ms step:960/1575 train_time:45364ms step_avg:47.25ms step:961/1575 train_time:45428ms step_avg:47.27ms step:962/1575 train_time:45487ms step_avg:47.28ms step:963/1575 train_time:45550ms step_avg:47.30ms step:964/1575 train_time:45609ms step_avg:47.31ms step:965/1575 train_time:45672ms step_avg:47.33ms step:966/1575 train_time:45732ms step_avg:47.34ms step:967/1575 train_time:45795ms step_avg:47.36ms step:968/1575 train_time:45854ms step_avg:47.37ms step:969/1575 train_time:45917ms step_avg:47.39ms step:970/1575 train_time:45976ms step_avg:47.40ms step:971/1575 train_time:46040ms step_avg:47.41ms step:972/1575 train_time:46099ms step_avg:47.43ms step:973/1575 train_time:46162ms step_avg:47.44ms step:974/1575 train_time:46222ms step_avg:47.46ms step:975/1575 train_time:46286ms step_avg:47.47ms step:976/1575 train_time:46346ms step_avg:47.49ms step:977/1575 train_time:46409ms step_avg:47.50ms step:978/1575 train_time:46468ms step_avg:47.51ms step:979/1575 train_time:46531ms step_avg:47.53ms step:980/1575 train_time:46590ms step_avg:47.54ms step:981/1575 train_time:46653ms step_avg:47.56ms step:982/1575 train_time:46712ms step_avg:47.57ms step:983/1575 train_time:46775ms step_avg:47.58ms step:984/1575 train_time:46834ms step_avg:47.60ms step:985/1575 train_time:46898ms step_avg:47.61ms step:986/1575 train_time:46957ms step_avg:47.62ms step:987/1575 train_time:47020ms step_avg:47.64ms step:988/1575 train_time:47080ms step_avg:47.65ms step:989/1575 train_time:47144ms step_avg:47.67ms step:990/1575 train_time:47203ms step_avg:47.68ms step:991/1575 train_time:47266ms step_avg:47.70ms step:992/1575 train_time:47327ms step_avg:47.71ms step:993/1575 train_time:47389ms step_avg:47.72ms step:994/1575 train_time:47449ms step_avg:47.74ms step:995/1575 train_time:47512ms step_avg:47.75ms step:996/1575 train_time:47571ms step_avg:47.76ms step:997/1575 train_time:47634ms step_avg:47.78ms step:998/1575 train_time:47693ms step_avg:47.79ms step:999/1575 train_time:47756ms step_avg:47.80ms step:1000/1575 train_time:47815ms step_avg:47.81ms step:1000/1575 val_loss:3.5816 train_time:47861ms step_avg:47.86ms step:1001/1575 train_time:47882ms step_avg:47.83ms step:1002/1575 train_time:47943ms step_avg:47.85ms step:1003/1575 train_time:48008ms step_avg:47.86ms step:1004/1575 train_time:48070ms step_avg:47.88ms step:1005/1575 train_time:48133ms step_avg:47.89ms step:1006/1575 train_time:48194ms step_avg:47.91ms step:1007/1575 train_time:48257ms step_avg:47.92ms step:1008/1575 train_time:48315ms step_avg:47.93ms step:1009/1575 train_time:48377ms step_avg:47.95ms step:1010/1575 train_time:48436ms step_avg:47.96ms step:1011/1575 train_time:48499ms step_avg:47.97ms step:1012/1575 train_time:48559ms step_avg:47.98ms step:1013/1575 train_time:48622ms step_avg:48.00ms step:1014/1575 train_time:48680ms step_avg:48.01ms step:1015/1575 train_time:48743ms step_avg:48.02ms step:1016/1575 train_time:48802ms step_avg:48.03ms step:1017/1575 train_time:48866ms step_avg:48.05ms step:1018/1575 train_time:48925ms step_avg:48.06ms step:1019/1575 train_time:48990ms step_avg:48.08ms step:1020/1575 train_time:49051ms step_avg:48.09ms step:1021/1575 train_time:49115ms step_avg:48.10ms step:1022/1575 train_time:49174ms step_avg:48.12ms step:1023/1575 train_time:49237ms step_avg:48.13ms step:1024/1575 train_time:49296ms step_avg:48.14ms step:1025/1575 train_time:49366ms step_avg:48.16ms step:1026/1575 train_time:49451ms step_avg:48.20ms step:1027/1575 train_time:49541ms step_avg:48.24ms step:1028/1575 train_time:49625ms step_avg:48.27ms step:1029/1575 train_time:49714ms step_avg:48.31ms step:1030/1575 train_time:49800ms step_avg:48.35ms step:1031/1575 train_time:49889ms step_avg:48.39ms step:1032/1575 train_time:49975ms step_avg:48.43ms step:1033/1575 train_time:50066ms step_avg:48.47ms step:1034/1575 train_time:50152ms step_avg:48.50ms step:1035/1575 train_time:50243ms step_avg:48.54ms step:1036/1575 train_time:50328ms step_avg:48.58ms step:1037/1575 train_time:50417ms step_avg:48.62ms step:1038/1575 train_time:50502ms step_avg:48.65ms step:1039/1575 train_time:50591ms step_avg:48.69ms step:1040/1575 train_time:50677ms step_avg:48.73ms step:1041/1575 train_time:50766ms step_avg:48.77ms step:1042/1575 train_time:50851ms step_avg:48.80ms step:1043/1575 train_time:50941ms step_avg:48.84ms step:1044/1575 train_time:51026ms step_avg:48.88ms step:1045/1575 train_time:51118ms step_avg:48.92ms step:1046/1575 train_time:51203ms step_avg:48.95ms step:1047/1575 train_time:51293ms step_avg:48.99ms step:1048/1575 train_time:51379ms step_avg:49.03ms step:1049/1575 train_time:51468ms step_avg:49.06ms step:1050/1575 train_time:51552ms step_avg:49.10ms step:1051/1575 train_time:51642ms step_avg:49.14ms step:1052/1575 train_time:51727ms step_avg:49.17ms step:1053/1575 train_time:51817ms step_avg:49.21ms step:1054/1575 train_time:51902ms step_avg:49.24ms step:1055/1575 train_time:51991ms step_avg:49.28ms step:1056/1575 train_time:52078ms step_avg:49.32ms step:1057/1575 train_time:52168ms step_avg:49.35ms step:1058/1575 train_time:52254ms step_avg:49.39ms step:1059/1575 train_time:52344ms step_avg:49.43ms step:1060/1575 train_time:52430ms step_avg:49.46ms step:1061/1575 train_time:52519ms step_avg:49.50ms step:1062/1575 train_time:52604ms step_avg:49.53ms step:1063/1575 train_time:52692ms step_avg:49.57ms step:1064/1575 train_time:52779ms step_avg:49.60ms step:1065/1575 train_time:52869ms step_avg:49.64ms step:1066/1575 train_time:52955ms step_avg:49.68ms step:1067/1575 train_time:53046ms step_avg:49.72ms step:1068/1575 train_time:53131ms step_avg:49.75ms step:1069/1575 train_time:53221ms step_avg:49.79ms step:1070/1575 train_time:53306ms step_avg:49.82ms step:1071/1575 train_time:53395ms step_avg:49.86ms step:1072/1575 train_time:53481ms step_avg:49.89ms step:1073/1575 train_time:53570ms step_avg:49.93ms step:1074/1575 train_time:53656ms step_avg:49.96ms step:1075/1575 train_time:53746ms step_avg:50.00ms step:1076/1575 train_time:53831ms step_avg:50.03ms step:1077/1575 train_time:53920ms step_avg:50.06ms step:1078/1575 train_time:54005ms step_avg:50.10ms step:1079/1575 train_time:54094ms step_avg:50.13ms step:1080/1575 train_time:54182ms step_avg:50.17ms step:1081/1575 train_time:54271ms step_avg:50.20ms step:1082/1575 train_time:54356ms step_avg:50.24ms step:1083/1575 train_time:54446ms step_avg:50.27ms step:1084/1575 train_time:54532ms step_avg:50.31ms step:1085/1575 train_time:54621ms step_avg:50.34ms step:1086/1575 train_time:54707ms step_avg:50.37ms step:1087/1575 train_time:54796ms step_avg:50.41ms step:1088/1575 train_time:54882ms step_avg:50.44ms step:1089/1575 train_time:54972ms step_avg:50.48ms step:1090/1575 train_time:55057ms step_avg:50.51ms step:1091/1575 train_time:55146ms step_avg:50.55ms step:1092/1575 train_time:55231ms step_avg:50.58ms step:1093/1575 train_time:55321ms step_avg:50.61ms step:1094/1575 train_time:55406ms step_avg:50.65ms step:1095/1575 train_time:55496ms step_avg:50.68ms step:1096/1575 train_time:55583ms step_avg:50.71ms step:1097/1575 train_time:55671ms step_avg:50.75ms step:1098/1575 train_time:55757ms step_avg:50.78ms step:1099/1575 train_time:55847ms step_avg:50.82ms step:1100/1575 train_time:55933ms step_avg:50.85ms step:1101/1575 train_time:56023ms step_avg:50.88ms step:1102/1575 train_time:56108ms step_avg:50.92ms step:1103/1575 train_time:56198ms step_avg:50.95ms step:1104/1575 train_time:56283ms step_avg:50.98ms step:1105/1575 train_time:56374ms step_avg:51.02ms step:1106/1575 train_time:56459ms step_avg:51.05ms step:1107/1575 train_time:56550ms step_avg:51.08ms step:1108/1575 train_time:56635ms step_avg:51.11ms step:1109/1575 train_time:56730ms step_avg:51.15ms step:1110/1575 train_time:56821ms step_avg:51.19ms step:1111/1575 train_time:56903ms step_avg:51.22ms step:1112/1575 train_time:56988ms step_avg:51.25ms step:1113/1575 train_time:57077ms step_avg:51.28ms step:1114/1575 train_time:57162ms step_avg:51.31ms step:1115/1575 train_time:57252ms step_avg:51.35ms step:1116/1575 train_time:57338ms step_avg:51.38ms step:1117/1575 train_time:57427ms step_avg:51.41ms step:1118/1575 train_time:57514ms step_avg:51.44ms step:1119/1575 train_time:57604ms step_avg:51.48ms step:1120/1575 train_time:57692ms step_avg:51.51ms step:1121/1575 train_time:57780ms step_avg:51.54ms step:1122/1575 train_time:57865ms step_avg:51.57ms step:1123/1575 train_time:57954ms step_avg:51.61ms step:1124/1575 train_time:58039ms step_avg:51.64ms step:1125/1575 train_time:58128ms step_avg:51.67ms step:1126/1575 train_time:58216ms step_avg:51.70ms step:1127/1575 train_time:58305ms step_avg:51.73ms step:1128/1575 train_time:58389ms step_avg:51.76ms step:1129/1575 train_time:58480ms step_avg:51.80ms step:1130/1575 train_time:58565ms step_avg:51.83ms step:1131/1575 train_time:58654ms step_avg:51.86ms step:1132/1575 train_time:58740ms step_avg:51.89ms step:1133/1575 train_time:58829ms step_avg:51.92ms step:1134/1575 train_time:58915ms step_avg:51.95ms step:1135/1575 train_time:59004ms step_avg:51.99ms step:1136/1575 train_time:59090ms step_avg:52.02ms step:1137/1575 train_time:59179ms step_avg:52.05ms step:1138/1575 train_time:59264ms step_avg:52.08ms step:1139/1575 train_time:59353ms step_avg:52.11ms step:1140/1575 train_time:59439ms step_avg:52.14ms step:1141/1575 train_time:59528ms step_avg:52.17ms step:1142/1575 train_time:59614ms step_avg:52.20ms step:1143/1575 train_time:59704ms step_avg:52.23ms step:1144/1575 train_time:59792ms step_avg:52.27ms step:1145/1575 train_time:59880ms step_avg:52.30ms step:1146/1575 train_time:59969ms step_avg:52.33ms step:1147/1575 train_time:60056ms step_avg:52.36ms step:1148/1575 train_time:60142ms step_avg:52.39ms step:1149/1575 train_time:60230ms step_avg:52.42ms step:1150/1575 train_time:60316ms step_avg:52.45ms step:1151/1575 train_time:60406ms step_avg:52.48ms step:1152/1575 train_time:60492ms step_avg:52.51ms step:1153/1575 train_time:60580ms step_avg:52.54ms step:1154/1575 train_time:60666ms step_avg:52.57ms step:1155/1575 train_time:60755ms step_avg:52.60ms step:1156/1575 train_time:60842ms step_avg:52.63ms step:1157/1575 train_time:60930ms step_avg:52.66ms step:1158/1575 train_time:61017ms step_avg:52.69ms step:1159/1575 train_time:61106ms step_avg:52.72ms step:1160/1575 train_time:61192ms step_avg:52.75ms step:1161/1575 train_time:61282ms step_avg:52.78ms step:1162/1575 train_time:61367ms step_avg:52.81ms step:1163/1575 train_time:61457ms step_avg:52.84ms step:1164/1575 train_time:61543ms step_avg:52.87ms step:1165/1575 train_time:61633ms step_avg:52.90ms step:1166/1575 train_time:61717ms step_avg:52.93ms step:1167/1575 train_time:61807ms step_avg:52.96ms step:1168/1575 train_time:61894ms step_avg:52.99ms step:1169/1575 train_time:61983ms step_avg:53.02ms step:1170/1575 train_time:62067ms step_avg:53.05ms step:1171/1575 train_time:62157ms step_avg:53.08ms step:1172/1575 train_time:62243ms step_avg:53.11ms step:1173/1575 train_time:62331ms step_avg:53.14ms step:1174/1575 train_time:62417ms step_avg:53.17ms step:1175/1575 train_time:62507ms step_avg:53.20ms step:1176/1575 train_time:62592ms step_avg:53.22ms step:1177/1575 train_time:62681ms step_avg:53.25ms step:1178/1575 train_time:62767ms step_avg:53.28ms step:1179/1575 train_time:62856ms step_avg:53.31ms step:1180/1575 train_time:62942ms step_avg:53.34ms step:1181/1575 train_time:63031ms step_avg:53.37ms step:1182/1575 train_time:63116ms step_avg:53.40ms step:1183/1575 train_time:63206ms step_avg:53.43ms step:1184/1575 train_time:63293ms step_avg:53.46ms step:1185/1575 train_time:63381ms step_avg:53.49ms step:1186/1575 train_time:63466ms step_avg:53.51ms step:1187/1575 train_time:63555ms step_avg:53.54ms step:1188/1575 train_time:63641ms step_avg:53.57ms step:1189/1575 train_time:63731ms step_avg:53.60ms step:1190/1575 train_time:63816ms step_avg:53.63ms step:1191/1575 train_time:63907ms step_avg:53.66ms step:1192/1575 train_time:63992ms step_avg:53.68ms step:1193/1575 train_time:64081ms step_avg:53.71ms step:1194/1575 train_time:64166ms step_avg:53.74ms step:1195/1575 train_time:64256ms step_avg:53.77ms step:1196/1575 train_time:64341ms step_avg:53.80ms step:1197/1575 train_time:64429ms step_avg:53.83ms step:1198/1575 train_time:64516ms step_avg:53.85ms step:1199/1575 train_time:64606ms step_avg:53.88ms step:1200/1575 train_time:64691ms step_avg:53.91ms step:1201/1575 train_time:64782ms step_avg:53.94ms step:1202/1575 train_time:64866ms step_avg:53.97ms step:1203/1575 train_time:64956ms step_avg:53.99ms step:1204/1575 train_time:65042ms step_avg:54.02ms step:1205/1575 train_time:65132ms step_avg:54.05ms step:1206/1575 train_time:65218ms step_avg:54.08ms step:1207/1575 train_time:65308ms step_avg:54.11ms step:1208/1575 train_time:65393ms step_avg:54.13ms step:1209/1575 train_time:65484ms step_avg:54.16ms step:1210/1575 train_time:65568ms step_avg:54.19ms step:1211/1575 train_time:65658ms step_avg:54.22ms step:1212/1575 train_time:65743ms step_avg:54.24ms step:1213/1575 train_time:65833ms step_avg:54.27ms step:1214/1575 train_time:65920ms step_avg:54.30ms step:1215/1575 train_time:66009ms step_avg:54.33ms step:1216/1575 train_time:66094ms step_avg:54.35ms step:1217/1575 train_time:66184ms step_avg:54.38ms step:1218/1575 train_time:66269ms step_avg:54.41ms step:1219/1575 train_time:66358ms step_avg:54.44ms step:1220/1575 train_time:66444ms step_avg:54.46ms step:1221/1575 train_time:66532ms step_avg:54.49ms step:1222/1575 train_time:66619ms step_avg:54.52ms step:1223/1575 train_time:66708ms step_avg:54.54ms step:1224/1575 train_time:66794ms step_avg:54.57ms step:1225/1575 train_time:66885ms step_avg:54.60ms step:1226/1575 train_time:66970ms step_avg:54.62ms step:1227/1575 train_time:67059ms step_avg:54.65ms step:1228/1575 train_time:67145ms step_avg:54.68ms step:1229/1575 train_time:67234ms step_avg:54.71ms step:1230/1575 train_time:67320ms step_avg:54.73ms step:1231/1575 train_time:67409ms step_avg:54.76ms step:1232/1575 train_time:67495ms step_avg:54.78ms step:1233/1575 train_time:67585ms step_avg:54.81ms step:1234/1575 train_time:67671ms step_avg:54.84ms step:1235/1575 train_time:67761ms step_avg:54.87ms step:1236/1575 train_time:67848ms step_avg:54.89ms step:1237/1575 train_time:67938ms step_avg:54.92ms step:1238/1575 train_time:68022ms step_avg:54.94ms step:1239/1575 train_time:68111ms step_avg:54.97ms step:1240/1575 train_time:68196ms step_avg:55.00ms step:1241/1575 train_time:68286ms step_avg:55.02ms step:1242/1575 train_time:68371ms step_avg:55.05ms step:1243/1575 train_time:68462ms step_avg:55.08ms step:1244/1575 train_time:68546ms step_avg:55.10ms step:1245/1575 train_time:68636ms step_avg:55.13ms step:1246/1575 train_time:68722ms step_avg:55.15ms step:1247/1575 train_time:68811ms step_avg:55.18ms step:1248/1575 train_time:68897ms step_avg:55.21ms step:1249/1575 train_time:68988ms step_avg:55.23ms step:1250/1575 train_time:69074ms step_avg:55.26ms step:1250/1575 val_loss:3.4069 train_time:69146ms step_avg:55.32ms step:1251/1575 train_time:69167ms step_avg:55.29ms step:1252/1575 train_time:69253ms step_avg:55.31ms step:1253/1575 train_time:69348ms step_avg:55.35ms step:1254/1575 train_time:69434ms step_avg:55.37ms step:1255/1575 train_time:69524ms step_avg:55.40ms step:1256/1575 train_time:69609ms step_avg:55.42ms step:1257/1575 train_time:69697ms step_avg:55.45ms step:1258/1575 train_time:69781ms step_avg:55.47ms step:1259/1575 train_time:69869ms step_avg:55.50ms step:1260/1575 train_time:69954ms step_avg:55.52ms step:1261/1575 train_time:70043ms step_avg:55.55ms step:1262/1575 train_time:70129ms step_avg:55.57ms step:1263/1575 train_time:70222ms step_avg:55.60ms step:1264/1575 train_time:70310ms step_avg:55.62ms step:1265/1575 train_time:70401ms step_avg:55.65ms step:1266/1575 train_time:70486ms step_avg:55.68ms step:1267/1575 train_time:70575ms step_avg:55.70ms step:1268/1575 train_time:70660ms step_avg:55.73ms step:1269/1575 train_time:70748ms step_avg:55.75ms step:1270/1575 train_time:70834ms step_avg:55.77ms step:1271/1575 train_time:70922ms step_avg:55.80ms step:1272/1575 train_time:71007ms step_avg:55.82ms step:1273/1575 train_time:71096ms step_avg:55.85ms step:1274/1575 train_time:71183ms step_avg:55.87ms step:1275/1575 train_time:71274ms step_avg:55.90ms step:1276/1575 train_time:71360ms step_avg:55.92ms step:1277/1575 train_time:71451ms step_avg:55.95ms step:1278/1575 train_time:71537ms step_avg:55.98ms step:1279/1575 train_time:71625ms step_avg:56.00ms step:1280/1575 train_time:71712ms step_avg:56.02ms step:1281/1575 train_time:71802ms step_avg:56.05ms step:1282/1575 train_time:71887ms step_avg:56.07ms step:1283/1575 train_time:71975ms step_avg:56.10ms step:1284/1575 train_time:72060ms step_avg:56.12ms step:1285/1575 train_time:72149ms step_avg:56.15ms step:1286/1575 train_time:72236ms step_avg:56.17ms step:1287/1575 train_time:72327ms step_avg:56.20ms step:1288/1575 train_time:72414ms step_avg:56.22ms step:1289/1575 train_time:72504ms step_avg:56.25ms step:1290/1575 train_time:72590ms step_avg:56.27ms step:1291/1575 train_time:72682ms step_avg:56.30ms step:1292/1575 train_time:72765ms step_avg:56.32ms step:1293/1575 train_time:72854ms step_avg:56.35ms step:1294/1575 train_time:72940ms step_avg:56.37ms step:1295/1575 train_time:73028ms step_avg:56.39ms step:1296/1575 train_time:73116ms step_avg:56.42ms step:1297/1575 train_time:73204ms step_avg:56.44ms step:1298/1575 train_time:73291ms step_avg:56.46ms step:1299/1575 train_time:73381ms step_avg:56.49ms step:1300/1575 train_time:73467ms step_avg:56.51ms step:1301/1575 train_time:73557ms step_avg:56.54ms step:1302/1575 train_time:73643ms step_avg:56.56ms step:1303/1575 train_time:73732ms step_avg:56.59ms step:1304/1575 train_time:73818ms step_avg:56.61ms step:1305/1575 train_time:73906ms step_avg:56.63ms step:1306/1575 train_time:73992ms step_avg:56.66ms step:1307/1575 train_time:74081ms step_avg:56.68ms step:1308/1575 train_time:74166ms step_avg:56.70ms step:1309/1575 train_time:74256ms step_avg:56.73ms step:1310/1575 train_time:74342ms step_avg:56.75ms step:1311/1575 train_time:74432ms step_avg:56.78ms step:1312/1575 train_time:74518ms step_avg:56.80ms step:1313/1575 train_time:74608ms step_avg:56.82ms step:1314/1575 train_time:74693ms step_avg:56.84ms step:1315/1575 train_time:74782ms step_avg:56.87ms step:1316/1575 train_time:74868ms step_avg:56.89ms step:1317/1575 train_time:74957ms step_avg:56.91ms step:1318/1575 train_time:75042ms step_avg:56.94ms step:1319/1575 train_time:75132ms step_avg:56.96ms step:1320/1575 train_time:75218ms step_avg:56.98ms step:1321/1575 train_time:75307ms step_avg:57.01ms step:1322/1575 train_time:75394ms step_avg:57.03ms step:1323/1575 train_time:75483ms step_avg:57.05ms step:1324/1575 train_time:75569ms step_avg:57.08ms step:1325/1575 train_time:75659ms step_avg:57.10ms step:1326/1575 train_time:75745ms step_avg:57.12ms step:1327/1575 train_time:75834ms step_avg:57.15ms step:1328/1575 train_time:75919ms step_avg:57.17ms step:1329/1575 train_time:76008ms step_avg:57.19ms step:1330/1575 train_time:76095ms step_avg:57.21ms step:1331/1575 train_time:76184ms step_avg:57.24ms step:1332/1575 train_time:76270ms step_avg:57.26ms step:1333/1575 train_time:76360ms step_avg:57.28ms step:1334/1575 train_time:76445ms step_avg:57.31ms step:1335/1575 train_time:76534ms step_avg:57.33ms step:1336/1575 train_time:76621ms step_avg:57.35ms step:1337/1575 train_time:76710ms step_avg:57.37ms step:1338/1575 train_time:76795ms step_avg:57.40ms step:1339/1575 train_time:76885ms step_avg:57.42ms step:1340/1575 train_time:76971ms step_avg:57.44ms step:1341/1575 train_time:77061ms step_avg:57.47ms step:1342/1575 train_time:77146ms step_avg:57.49ms step:1343/1575 train_time:77236ms step_avg:57.51ms step:1344/1575 train_time:77321ms step_avg:57.53ms step:1345/1575 train_time:77411ms step_avg:57.55ms step:1346/1575 train_time:77497ms step_avg:57.58ms step:1347/1575 train_time:77587ms step_avg:57.60ms step:1348/1575 train_time:77673ms step_avg:57.62ms step:1349/1575 train_time:77763ms step_avg:57.65ms step:1350/1575 train_time:77848ms step_avg:57.67ms step:1351/1575 train_time:77937ms step_avg:57.69ms step:1352/1575 train_time:78022ms step_avg:57.71ms step:1353/1575 train_time:78111ms step_avg:57.73ms step:1354/1575 train_time:78198ms step_avg:57.75ms step:1355/1575 train_time:78286ms step_avg:57.78ms step:1356/1575 train_time:78372ms step_avg:57.80ms step:1357/1575 train_time:78462ms step_avg:57.82ms step:1358/1575 train_time:78552ms step_avg:57.84ms step:1359/1575 train_time:78639ms step_avg:57.87ms step:1360/1575 train_time:78725ms step_avg:57.89ms step:1361/1575 train_time:78813ms step_avg:57.91ms step:1362/1575 train_time:78900ms step_avg:57.93ms step:1363/1575 train_time:78988ms step_avg:57.95ms step:1364/1575 train_time:79074ms step_avg:57.97ms step:1365/1575 train_time:79163ms step_avg:58.00ms step:1366/1575 train_time:79249ms step_avg:58.02ms step:1367/1575 train_time:79339ms step_avg:58.04ms step:1368/1575 train_time:79425ms step_avg:58.06ms step:1369/1575 train_time:79514ms step_avg:58.08ms step:1370/1575 train_time:79600ms step_avg:58.10ms step:1371/1575 train_time:79689ms step_avg:58.13ms step:1372/1575 train_time:79775ms step_avg:58.15ms step:1373/1575 train_time:79865ms step_avg:58.17ms step:1374/1575 train_time:79950ms step_avg:58.19ms step:1375/1575 train_time:80040ms step_avg:58.21ms step:1376/1575 train_time:80127ms step_avg:58.23ms step:1377/1575 train_time:80215ms step_avg:58.25ms step:1378/1575 train_time:80300ms step_avg:58.27ms step:1379/1575 train_time:80393ms step_avg:58.30ms step:1380/1575 train_time:80478ms step_avg:58.32ms step:1381/1575 train_time:80566ms step_avg:58.34ms step:1382/1575 train_time:80652ms step_avg:58.36ms step:1383/1575 train_time:80741ms step_avg:58.38ms step:1384/1575 train_time:80826ms step_avg:58.40ms step:1385/1575 train_time:80916ms step_avg:58.42ms step:1386/1575 train_time:81002ms step_avg:58.44ms step:1387/1575 train_time:81092ms step_avg:58.47ms step:1388/1575 train_time:81177ms step_avg:58.49ms step:1389/1575 train_time:81266ms step_avg:58.51ms step:1390/1575 train_time:81352ms step_avg:58.53ms step:1391/1575 train_time:81441ms step_avg:58.55ms step:1392/1575 train_time:81527ms step_avg:58.57ms step:1393/1575 train_time:81617ms step_avg:58.59ms step:1394/1575 train_time:81702ms step_avg:58.61ms step:1395/1575 train_time:81791ms step_avg:58.63ms step:1396/1575 train_time:81877ms step_avg:58.65ms step:1397/1575 train_time:81968ms step_avg:58.67ms step:1398/1575 train_time:82053ms step_avg:58.69ms step:1399/1575 train_time:82143ms step_avg:58.72ms step:1400/1575 train_time:82229ms step_avg:58.74ms step:1401/1575 train_time:82319ms step_avg:58.76ms step:1402/1575 train_time:82404ms step_avg:58.78ms step:1403/1575 train_time:82494ms step_avg:58.80ms step:1404/1575 train_time:82579ms step_avg:58.82ms step:1405/1575 train_time:82668ms step_avg:58.84ms step:1406/1575 train_time:82755ms step_avg:58.86ms step:1407/1575 train_time:82845ms step_avg:58.88ms step:1408/1575 train_time:82933ms step_avg:58.90ms step:1409/1575 train_time:83019ms step_avg:58.92ms step:1410/1575 train_time:83104ms step_avg:58.94ms step:1411/1575 train_time:83194ms step_avg:58.96ms step:1412/1575 train_time:83280ms step_avg:58.98ms step:1413/1575 train_time:83369ms step_avg:59.00ms step:1414/1575 train_time:83456ms step_avg:59.02ms step:1415/1575 train_time:83543ms step_avg:59.04ms step:1416/1575 train_time:83630ms step_avg:59.06ms step:1417/1575 train_time:83720ms step_avg:59.08ms step:1418/1575 train_time:83805ms step_avg:59.10ms step:1419/1575 train_time:83895ms step_avg:59.12ms step:1420/1575 train_time:83982ms step_avg:59.14ms step:1421/1575 train_time:84071ms step_avg:59.16ms step:1422/1575 train_time:84157ms step_avg:59.18ms step:1423/1575 train_time:84246ms step_avg:59.20ms step:1424/1575 train_time:84333ms step_avg:59.22ms step:1425/1575 train_time:84422ms step_avg:59.24ms step:1426/1575 train_time:84507ms step_avg:59.26ms step:1427/1575 train_time:84597ms step_avg:59.28ms step:1428/1575 train_time:84682ms step_avg:59.30ms step:1429/1575 train_time:84772ms step_avg:59.32ms step:1430/1575 train_time:84858ms step_avg:59.34ms step:1431/1575 train_time:84947ms step_avg:59.36ms step:1432/1575 train_time:85034ms step_avg:59.38ms step:1433/1575 train_time:85124ms step_avg:59.40ms step:1434/1575 train_time:85209ms step_avg:59.42ms step:1435/1575 train_time:85301ms step_avg:59.44ms step:1436/1575 train_time:85386ms step_avg:59.46ms step:1437/1575 train_time:85475ms step_avg:59.48ms step:1438/1575 train_time:85561ms step_avg:59.50ms step:1439/1575 train_time:85650ms step_avg:59.52ms step:1440/1575 train_time:85736ms step_avg:59.54ms step:1441/1575 train_time:85829ms step_avg:59.56ms step:1442/1575 train_time:85912ms step_avg:59.58ms step:1443/1575 train_time:86000ms step_avg:59.60ms step:1444/1575 train_time:86086ms step_avg:59.62ms step:1445/1575 train_time:86175ms step_avg:59.64ms step:1446/1575 train_time:86261ms step_avg:59.65ms step:1447/1575 train_time:86351ms step_avg:59.68ms step:1448/1575 train_time:86436ms step_avg:59.69ms step:1449/1575 train_time:86525ms step_avg:59.71ms step:1450/1575 train_time:86611ms step_avg:59.73ms step:1451/1575 train_time:86700ms step_avg:59.75ms step:1452/1575 train_time:86786ms step_avg:59.77ms step:1453/1575 train_time:86878ms step_avg:59.79ms step:1454/1575 train_time:86963ms step_avg:59.81ms step:1455/1575 train_time:87051ms step_avg:59.83ms step:1456/1575 train_time:87138ms step_avg:59.85ms step:1457/1575 train_time:87228ms step_avg:59.87ms step:1458/1575 train_time:87314ms step_avg:59.89ms step:1459/1575 train_time:87404ms step_avg:59.91ms step:1460/1575 train_time:87491ms step_avg:59.93ms step:1461/1575 train_time:87579ms step_avg:59.94ms step:1462/1575 train_time:87665ms step_avg:59.96ms step:1463/1575 train_time:87753ms step_avg:59.98ms step:1464/1575 train_time:87839ms step_avg:60.00ms step:1465/1575 train_time:87928ms step_avg:60.02ms step:1466/1575 train_time:88016ms step_avg:60.04ms step:1467/1575 train_time:88104ms step_avg:60.06ms step:1468/1575 train_time:88191ms step_avg:60.08ms step:1469/1575 train_time:88280ms step_avg:60.10ms step:1470/1575 train_time:88365ms step_avg:60.11ms step:1471/1575 train_time:88455ms step_avg:60.13ms step:1472/1575 train_time:88541ms step_avg:60.15ms step:1473/1575 train_time:88631ms step_avg:60.17ms step:1474/1575 train_time:88717ms step_avg:60.19ms step:1475/1575 train_time:88806ms step_avg:60.21ms step:1476/1575 train_time:88892ms step_avg:60.22ms step:1477/1575 train_time:88982ms step_avg:60.25ms step:1478/1575 train_time:89068ms step_avg:60.26ms step:1479/1575 train_time:89158ms step_avg:60.28ms step:1480/1575 train_time:89244ms step_avg:60.30ms step:1481/1575 train_time:89333ms step_avg:60.32ms step:1482/1575 train_time:89419ms step_avg:60.34ms step:1483/1575 train_time:89508ms step_avg:60.36ms step:1484/1575 train_time:89594ms step_avg:60.37ms step:1485/1575 train_time:89684ms step_avg:60.39ms step:1486/1575 train_time:89770ms step_avg:60.41ms step:1487/1575 train_time:89860ms step_avg:60.43ms step:1488/1575 train_time:89945ms step_avg:60.45ms step:1489/1575 train_time:90036ms step_avg:60.47ms step:1490/1575 train_time:90121ms step_avg:60.48ms step:1491/1575 train_time:90212ms step_avg:60.50ms step:1492/1575 train_time:90297ms step_avg:60.52ms step:1493/1575 train_time:90385ms step_avg:60.54ms step:1494/1575 train_time:90472ms step_avg:60.56ms step:1495/1575 train_time:90562ms step_avg:60.58ms step:1496/1575 train_time:90648ms step_avg:60.59ms step:1497/1575 train_time:90737ms step_avg:60.61ms step:1498/1575 train_time:90823ms step_avg:60.63ms step:1499/1575 train_time:90912ms step_avg:60.65ms step:1500/1575 train_time:90998ms step_avg:60.67ms step:1500/1575 val_loss:3.3014 train_time:91070ms step_avg:60.71ms step:1501/1575 train_time:91091ms step_avg:60.69ms step:1502/1575 train_time:91179ms step_avg:60.70ms step:1503/1575 train_time:91273ms step_avg:60.73ms step:1504/1575 train_time:91359ms step_avg:60.74ms step:1505/1575 train_time:91448ms step_avg:60.76ms step:1506/1575 train_time:91533ms step_avg:60.78ms step:1507/1575 train_time:91622ms step_avg:60.80ms step:1508/1575 train_time:91706ms step_avg:60.81ms step:1509/1575 train_time:91794ms step_avg:60.83ms step:1510/1575 train_time:91879ms step_avg:60.85ms step:1511/1575 train_time:91967ms step_avg:60.86ms step:1512/1575 train_time:92053ms step_avg:60.88ms step:1513/1575 train_time:92145ms step_avg:60.90ms step:1514/1575 train_time:92235ms step_avg:60.92ms step:1515/1575 train_time:92325ms step_avg:60.94ms step:1516/1575 train_time:92411ms step_avg:60.96ms step:1517/1575 train_time:92500ms step_avg:60.98ms step:1518/1575 train_time:92585ms step_avg:60.99ms step:1519/1575 train_time:92674ms step_avg:61.01ms step:1520/1575 train_time:92760ms step_avg:61.03ms step:1521/1575 train_time:92848ms step_avg:61.04ms step:1522/1575 train_time:92933ms step_avg:61.06ms step:1523/1575 train_time:93023ms step_avg:61.08ms step:1524/1575 train_time:93109ms step_avg:61.09ms step:1525/1575 train_time:93200ms step_avg:61.11ms step:1526/1575 train_time:93286ms step_avg:61.13ms step:1527/1575 train_time:93377ms step_avg:61.15ms step:1528/1575 train_time:93462ms step_avg:61.17ms step:1529/1575 train_time:93551ms step_avg:61.18ms step:1530/1575 train_time:93636ms step_avg:61.20ms step:1531/1575 train_time:93727ms step_avg:61.22ms step:1532/1575 train_time:93811ms step_avg:61.23ms step:1533/1575 train_time:93899ms step_avg:61.25ms step:1534/1575 train_time:93984ms step_avg:61.27ms step:1535/1575 train_time:94074ms step_avg:61.29ms step:1536/1575 train_time:94167ms step_avg:61.31ms step:1537/1575 train_time:94257ms step_avg:61.33ms step:1538/1575 train_time:94343ms step_avg:61.34ms step:1539/1575 train_time:94433ms step_avg:61.36ms step:1540/1575 train_time:94519ms step_avg:61.38ms step:1541/1575 train_time:94610ms step_avg:61.40ms step:1542/1575 train_time:94694ms step_avg:61.41ms step:1543/1575 train_time:94784ms step_avg:61.43ms step:1544/1575 train_time:94869ms step_avg:61.44ms step:1545/1575 train_time:94958ms step_avg:61.46ms step:1546/1575 train_time:95043ms step_avg:61.48ms step:1547/1575 train_time:95134ms step_avg:61.50ms step:1548/1575 train_time:95221ms step_avg:61.51ms step:1549/1575 train_time:95311ms step_avg:61.53ms step:1550/1575 train_time:95398ms step_avg:61.55ms step:1551/1575 train_time:95487ms step_avg:61.57ms step:1552/1575 train_time:95573ms step_avg:61.58ms step:1553/1575 train_time:95663ms step_avg:61.60ms step:1554/1575 train_time:95751ms step_avg:61.62ms step:1555/1575 train_time:95838ms step_avg:61.63ms step:1556/1575 train_time:95923ms step_avg:61.65ms step:1557/1575 train_time:96013ms step_avg:61.67ms step:1558/1575 train_time:96102ms step_avg:61.68ms step:1559/1575 train_time:96190ms step_avg:61.70ms step:1560/1575 train_time:96276ms step_avg:61.72ms step:1561/1575 train_time:96366ms step_avg:61.73ms step:1562/1575 train_time:96453ms step_avg:61.75ms step:1563/1575 train_time:96544ms step_avg:61.77ms step:1564/1575 train_time:96629ms step_avg:61.78ms step:1565/1575 train_time:96719ms step_avg:61.80ms step:1566/1575 train_time:96805ms step_avg:61.82ms step:1567/1575 train_time:96895ms step_avg:61.83ms step:1568/1575 train_time:96980ms step_avg:61.85ms step:1569/1575 train_time:97073ms step_avg:61.87ms step:1570/1575 train_time:97158ms step_avg:61.88ms step:1571/1575 train_time:97247ms step_avg:61.90ms step:1572/1575 train_time:97333ms step_avg:61.92ms step:1573/1575 train_time:97424ms step_avg:61.94ms step:1574/1575 train_time:97509ms step_avg:61.95ms step:1575/1575 train_time:97599ms step_avg:61.97ms step:1575/1575 val_loss:3.2799 train_time:97666ms step_avg:62.01ms peak memory allocated: 31016 MiB reserved: 46998 MiB