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:36:55 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 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:62:00.0 Off | 0 | | N/A 38C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 40C P0 120W / 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 42C P0 132W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 40C P0 124W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 36C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 257483 C /usr/bin/python3 1510MiB | | 1 N/A N/A 257484 C /usr/bin/python3 1510MiB | | 2 N/A N/A 257485 C /usr/bin/python3 1510MiB | | 3 N/A N/A 257486 C /usr/bin/python3 1510MiB | | 4 N/A N/A 257487 C /usr/bin/python3 1510MiB | | 5 N/A N/A 257488 C /usr/bin/python3 1510MiB | | 6 N/A N/A 257489 C /usr/bin/python3 1510MiB | | 7 N/A N/A 257490 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.8298 train_time:0ms step_avg:0.03ms step:1/1575 train_time:76ms step_avg:75.52ms step:2/1575 train_time:96ms step_avg:48.21ms step:3/1575 train_time:114ms step_avg:38.15ms step:4/1575 train_time:157ms step_avg:39.27ms step:5/1575 train_time:188ms step_avg:37.50ms step:6/1575 train_time:269ms step_avg:44.78ms step:7/1575 train_time:285ms step_avg:40.71ms step:8/1575 train_time:325ms step_avg:40.62ms step:9/1575 train_time:356ms step_avg:39.51ms step:10/1575 train_time:394ms step_avg:39.42ms step:11/1575 train_time:425ms step_avg:38.62ms step:12/1575 train_time:463ms step_avg:38.61ms step:13/1575 train_time:495ms step_avg:38.08ms step:14/1575 train_time:533ms step_avg:38.10ms step:15/1575 train_time:565ms step_avg:37.65ms step:16/1575 train_time:603ms step_avg:37.71ms step:17/1575 train_time:635ms step_avg:37.34ms step:18/1575 train_time:673ms step_avg:37.39ms step:19/1575 train_time:704ms step_avg:37.07ms step:20/1575 train_time:743ms step_avg:37.16ms step:21/1575 train_time:774ms step_avg:36.87ms step:22/1575 train_time:813ms step_avg:36.93ms step:23/1575 train_time:844ms step_avg:36.68ms step:24/1575 train_time:882ms step_avg:36.74ms step:25/1575 train_time:913ms step_avg:36.53ms step:26/1575 train_time:952ms step_avg:36.60ms step:27/1575 train_time:983ms step_avg:36.41ms step:28/1575 train_time:1022ms step_avg:36.48ms step:29/1575 train_time:1053ms step_avg:36.31ms step:30/1575 train_time:1091ms step_avg:36.38ms step:31/1575 train_time:1122ms step_avg:36.21ms step:32/1575 train_time:1161ms step_avg:36.27ms step:33/1575 train_time:1192ms step_avg:36.13ms step:34/1575 train_time:1231ms step_avg:36.20ms step:35/1575 train_time:1262ms step_avg:36.06ms step:36/1575 train_time:1301ms step_avg:36.14ms step:37/1575 train_time:1332ms step_avg:36.01ms step:38/1575 train_time:1371ms step_avg:36.07ms step:39/1575 train_time:1402ms step_avg:35.94ms step:40/1575 train_time:1440ms step_avg:36.00ms step:41/1575 train_time:1471ms step_avg:35.89ms step:42/1575 train_time:1510ms step_avg:35.96ms step:43/1575 train_time:1541ms step_avg:35.84ms step:44/1575 train_time:1580ms step_avg:35.91ms step:45/1575 train_time:1611ms step_avg:35.81ms step:46/1575 train_time:1649ms step_avg:35.86ms step:47/1575 train_time:1681ms step_avg:35.76ms step:48/1575 train_time:1719ms step_avg:35.82ms step:49/1575 train_time:1751ms step_avg:35.74ms step:50/1575 train_time:1789ms step_avg:35.79ms step:51/1575 train_time:1820ms step_avg:35.69ms step:52/1575 train_time:1858ms step_avg:35.74ms step:53/1575 train_time:1890ms step_avg:35.66ms step:54/1575 train_time:1928ms step_avg:35.70ms step:55/1575 train_time:1960ms step_avg:35.63ms step:56/1575 train_time:1998ms step_avg:35.68ms step:57/1575 train_time:2030ms step_avg:35.61ms step:58/1575 train_time:2068ms step_avg:35.66ms step:59/1575 train_time:2099ms step_avg:35.58ms step:60/1575 train_time:2138ms step_avg:35.63ms step:61/1575 train_time:2169ms step_avg:35.55ms step:62/1575 train_time:2207ms step_avg:35.60ms step:63/1575 train_time:2239ms step_avg:35.53ms step:64/1575 train_time:2277ms step_avg:35.57ms step:65/1575 train_time:2308ms step_avg:35.51ms step:66/1575 train_time:2346ms step_avg:35.55ms step:67/1575 train_time:2378ms step_avg:35.49ms step:68/1575 train_time:2416ms step_avg:35.53ms step:69/1575 train_time:2447ms step_avg:35.47ms step:70/1575 train_time:2486ms step_avg:35.51ms step:71/1575 train_time:2517ms step_avg:35.46ms step:72/1575 train_time:2556ms step_avg:35.50ms step:73/1575 train_time:2588ms step_avg:35.45ms step:74/1575 train_time:2626ms step_avg:35.48ms step:75/1575 train_time:2657ms step_avg:35.43ms step:76/1575 train_time:2696ms step_avg:35.47ms step:77/1575 train_time:2727ms step_avg:35.42ms step:78/1575 train_time:2766ms step_avg:35.46ms step:79/1575 train_time:2797ms step_avg:35.41ms step:80/1575 train_time:2836ms step_avg:35.45ms step:81/1575 train_time:2867ms step_avg:35.39ms step:82/1575 train_time:2905ms step_avg:35.43ms step:83/1575 train_time:2936ms step_avg:35.38ms step:84/1575 train_time:2975ms step_avg:35.41ms step:85/1575 train_time:3006ms step_avg:35.37ms step:86/1575 train_time:3045ms step_avg:35.41ms step:87/1575 train_time:3076ms step_avg:35.36ms step:88/1575 train_time:3114ms step_avg:35.39ms step:89/1575 train_time:3146ms step_avg:35.34ms step:90/1575 train_time:3184ms step_avg:35.38ms step:91/1575 train_time:3215ms step_avg:35.33ms step:92/1575 train_time:3254ms step_avg:35.37ms step:93/1575 train_time:3285ms step_avg:35.32ms step:94/1575 train_time:3323ms step_avg:35.36ms step:95/1575 train_time:3355ms step_avg:35.32ms step:96/1575 train_time:3394ms step_avg:35.35ms step:97/1575 train_time:3424ms step_avg:35.30ms step:98/1575 train_time:3463ms step_avg:35.33ms step:99/1575 train_time:3494ms step_avg:35.29ms step:100/1575 train_time:3532ms step_avg:35.32ms step:101/1575 train_time:3563ms step_avg:35.28ms step:102/1575 train_time:3602ms step_avg:35.31ms step:103/1575 train_time:3634ms step_avg:35.28ms step:104/1575 train_time:3672ms step_avg:35.31ms step:105/1575 train_time:3703ms step_avg:35.27ms step:106/1575 train_time:3741ms step_avg:35.29ms step:107/1575 train_time:3773ms step_avg:35.26ms step:108/1575 train_time:3811ms step_avg:35.29ms step:109/1575 train_time:3842ms step_avg:35.25ms step:110/1575 train_time:3880ms step_avg:35.27ms step:111/1575 train_time:3912ms step_avg:35.24ms step:112/1575 train_time:3950ms step_avg:35.27ms step:113/1575 train_time:3982ms step_avg:35.24ms step:114/1575 train_time:4020ms step_avg:35.26ms step:115/1575 train_time:4051ms step_avg:35.23ms step:116/1575 train_time:4090ms step_avg:35.26ms step:117/1575 train_time:4121ms step_avg:35.22ms step:118/1575 train_time:4161ms step_avg:35.26ms step:119/1575 train_time:4191ms step_avg:35.22ms step:120/1575 train_time:4229ms step_avg:35.24ms step:121/1575 train_time:4260ms step_avg:35.21ms step:122/1575 train_time:4299ms step_avg:35.24ms step:123/1575 train_time:4331ms step_avg:35.21ms step:124/1575 train_time:4369ms step_avg:35.24ms step:125/1575 train_time:4400ms step_avg:35.20ms step:126/1575 train_time:4439ms step_avg:35.23ms step:127/1575 train_time:4470ms step_avg:35.20ms step:128/1575 train_time:4508ms step_avg:35.22ms step:129/1575 train_time:4540ms step_avg:35.19ms step:130/1575 train_time:4579ms step_avg:35.22ms step:131/1575 train_time:4611ms step_avg:35.20ms step:132/1575 train_time:4649ms step_avg:35.22ms step:133/1575 train_time:4680ms step_avg:35.18ms step:134/1575 train_time:4718ms step_avg:35.21ms step:135/1575 train_time:4749ms step_avg:35.18ms step:136/1575 train_time:4787ms step_avg:35.20ms step:137/1575 train_time:4819ms step_avg:35.17ms step:138/1575 train_time:4857ms step_avg:35.20ms step:139/1575 train_time:4888ms step_avg:35.17ms step:140/1575 train_time:4927ms step_avg:35.19ms step:141/1575 train_time:4958ms step_avg:35.16ms step:142/1575 train_time:4996ms step_avg:35.19ms step:143/1575 train_time:5028ms step_avg:35.16ms step:144/1575 train_time:5067ms step_avg:35.18ms step:145/1575 train_time:5098ms step_avg:35.16ms step:146/1575 train_time:5136ms step_avg:35.18ms step:147/1575 train_time:5167ms step_avg:35.15ms step:148/1575 train_time:5206ms step_avg:35.18ms step:149/1575 train_time:5237ms step_avg:35.15ms step:150/1575 train_time:5275ms step_avg:35.17ms step:151/1575 train_time:5307ms step_avg:35.14ms step:152/1575 train_time:5345ms step_avg:35.17ms step:153/1575 train_time:5377ms step_avg:35.14ms step:154/1575 train_time:5415ms step_avg:35.16ms step:155/1575 train_time:5446ms step_avg:35.14ms step:156/1575 train_time:5484ms step_avg:35.16ms step:157/1575 train_time:5515ms step_avg:35.13ms step:158/1575 train_time:5554ms step_avg:35.15ms step:159/1575 train_time:5585ms step_avg:35.13ms step:160/1575 train_time:5624ms step_avg:35.15ms step:161/1575 train_time:5654ms step_avg:35.12ms step:162/1575 train_time:5693ms step_avg:35.14ms step:163/1575 train_time:5724ms step_avg:35.12ms step:164/1575 train_time:5762ms step_avg:35.14ms step:165/1575 train_time:5793ms step_avg:35.11ms step:166/1575 train_time:5832ms step_avg:35.13ms step:167/1575 train_time:5863ms step_avg:35.11ms step:168/1575 train_time:5902ms step_avg:35.13ms step:169/1575 train_time:5933ms step_avg:35.11ms step:170/1575 train_time:5971ms step_avg:35.12ms step:171/1575 train_time:6002ms step_avg:35.10ms step:172/1575 train_time:6040ms step_avg:35.12ms step:173/1575 train_time:6072ms step_avg:35.10ms step:174/1575 train_time:6110ms step_avg:35.11ms step:175/1575 train_time:6141ms step_avg:35.09ms step:176/1575 train_time:6179ms step_avg:35.11ms step:177/1575 train_time:6210ms step_avg:35.09ms step:178/1575 train_time:6248ms step_avg:35.10ms step:179/1575 train_time:6280ms step_avg:35.08ms step:180/1575 train_time:6318ms step_avg:35.10ms step:181/1575 train_time:6349ms step_avg:35.08ms step:182/1575 train_time:6387ms step_avg:35.10ms step:183/1575 train_time:6418ms step_avg:35.07ms step:184/1575 train_time:6457ms step_avg:35.09ms step:185/1575 train_time:6488ms step_avg:35.07ms step:186/1575 train_time:6526ms step_avg:35.09ms step:187/1575 train_time:6557ms step_avg:35.07ms step:188/1575 train_time:6596ms step_avg:35.08ms step:189/1575 train_time:6627ms step_avg:35.06ms step:190/1575 train_time:6665ms step_avg:35.08ms step:191/1575 train_time:6697ms step_avg:35.06ms step:192/1575 train_time:6735ms step_avg:35.08ms step:193/1575 train_time:6766ms step_avg:35.06ms step:194/1575 train_time:6805ms step_avg:35.08ms step:195/1575 train_time:6836ms step_avg:35.06ms step:196/1575 train_time:6874ms step_avg:35.07ms step:197/1575 train_time:6906ms step_avg:35.05ms step:198/1575 train_time:6944ms step_avg:35.07ms step:199/1575 train_time:6975ms step_avg:35.05ms step:200/1575 train_time:7013ms step_avg:35.07ms step:201/1575 train_time:7044ms step_avg:35.05ms step:202/1575 train_time:7083ms step_avg:35.06ms step:203/1575 train_time:7114ms step_avg:35.04ms step:204/1575 train_time:7152ms step_avg:35.06ms step:205/1575 train_time:7183ms step_avg:35.04ms step:206/1575 train_time:7222ms step_avg:35.06ms step:207/1575 train_time:7253ms step_avg:35.04ms step:208/1575 train_time:7291ms step_avg:35.05ms step:209/1575 train_time:7322ms step_avg:35.03ms step:210/1575 train_time:7361ms step_avg:35.05ms step:211/1575 train_time:7392ms step_avg:35.03ms step:212/1575 train_time:7430ms step_avg:35.05ms step:213/1575 train_time:7461ms step_avg:35.03ms step:214/1575 train_time:7499ms step_avg:35.04ms step:215/1575 train_time:7531ms step_avg:35.03ms step:216/1575 train_time:7569ms step_avg:35.04ms step:217/1575 train_time:7601ms step_avg:35.03ms step:218/1575 train_time:7639ms step_avg:35.04ms step:219/1575 train_time:7671ms step_avg:35.03ms step:220/1575 train_time:7709ms step_avg:35.04ms step:221/1575 train_time:7740ms step_avg:35.02ms step:222/1575 train_time:7778ms step_avg:35.04ms step:223/1575 train_time:7810ms step_avg:35.02ms step:224/1575 train_time:7848ms step_avg:35.04ms step:225/1575 train_time:7879ms step_avg:35.02ms step:226/1575 train_time:7918ms step_avg:35.04ms step:227/1575 train_time:7949ms step_avg:35.02ms step:228/1575 train_time:7988ms step_avg:35.03ms step:229/1575 train_time:8019ms step_avg:35.02ms step:230/1575 train_time:8057ms step_avg:35.03ms step:231/1575 train_time:8089ms step_avg:35.02ms step:232/1575 train_time:8128ms step_avg:35.03ms step:233/1575 train_time:8159ms step_avg:35.02ms step:234/1575 train_time:8197ms step_avg:35.03ms step:235/1575 train_time:8228ms step_avg:35.01ms step:236/1575 train_time:8267ms step_avg:35.03ms step:237/1575 train_time:8298ms step_avg:35.01ms step:238/1575 train_time:8337ms step_avg:35.03ms step:239/1575 train_time:8368ms step_avg:35.01ms step:240/1575 train_time:8406ms step_avg:35.03ms step:241/1575 train_time:8437ms step_avg:35.01ms step:242/1575 train_time:8476ms step_avg:35.02ms step:243/1575 train_time:8507ms step_avg:35.01ms step:244/1575 train_time:8545ms step_avg:35.02ms step:245/1575 train_time:8576ms step_avg:35.00ms step:246/1575 train_time:8614ms step_avg:35.02ms step:247/1575 train_time:8645ms step_avg:35.00ms step:248/1575 train_time:8684ms step_avg:35.01ms step:249/1575 train_time:8715ms step_avg:35.00ms step:250/1575 train_time:8753ms step_avg:35.01ms step:250/1575 val_loss:4.5756 train_time:8802ms step_avg:35.21ms step:251/1575 train_time:8820ms step_avg:35.14ms step:252/1575 train_time:8838ms step_avg:35.07ms step:253/1575 train_time:8858ms step_avg:35.01ms step:254/1575 train_time:8897ms step_avg:35.03ms step:255/1575 train_time:8929ms step_avg:35.02ms step:256/1575 train_time:8969ms step_avg:35.03ms step:257/1575 train_time:9001ms step_avg:35.02ms step:258/1575 train_time:9040ms step_avg:35.04ms step:259/1575 train_time:9071ms step_avg:35.02ms step:260/1575 train_time:9110ms step_avg:35.04ms step:261/1575 train_time:9141ms step_avg:35.02ms step:262/1575 train_time:9180ms step_avg:35.04ms step:263/1575 train_time:9211ms step_avg:35.02ms step:264/1575 train_time:9249ms step_avg:35.03ms step:265/1575 train_time:9280ms step_avg:35.02ms step:266/1575 train_time:9318ms step_avg:35.03ms step:267/1575 train_time:9349ms step_avg:35.02ms step:268/1575 train_time:9387ms step_avg:35.03ms step:269/1575 train_time:9419ms step_avg:35.01ms step:270/1575 train_time:9458ms step_avg:35.03ms step:271/1575 train_time:9488ms step_avg:35.01ms step:272/1575 train_time:9526ms step_avg:35.02ms step:273/1575 train_time:9557ms step_avg:35.01ms step:274/1575 train_time:9596ms step_avg:35.02ms step:275/1575 train_time:9627ms step_avg:35.01ms step:276/1575 train_time:9665ms step_avg:35.02ms step:277/1575 train_time:9697ms step_avg:35.01ms step:278/1575 train_time:9735ms step_avg:35.02ms step:279/1575 train_time:9766ms step_avg:35.00ms step:280/1575 train_time:9805ms step_avg:35.02ms step:281/1575 train_time:9836ms step_avg:35.00ms step:282/1575 train_time:9874ms step_avg:35.01ms step:283/1575 train_time:9905ms step_avg:35.00ms step:284/1575 train_time:9943ms step_avg:35.01ms step:285/1575 train_time:9974ms step_avg:35.00ms step:286/1575 train_time:10012ms step_avg:35.01ms step:287/1575 train_time:10044ms step_avg:35.00ms step:288/1575 train_time:10083ms step_avg:35.01ms step:289/1575 train_time:10114ms step_avg:35.00ms step:290/1575 train_time:10153ms step_avg:35.01ms step:291/1575 train_time:10184ms step_avg:35.00ms step:292/1575 train_time:10222ms step_avg:35.01ms step:293/1575 train_time:10253ms step_avg:34.99ms step:294/1575 train_time:10292ms step_avg:35.01ms step:295/1575 train_time:10322ms step_avg:34.99ms step:296/1575 train_time:10361ms step_avg:35.00ms step:297/1575 train_time:10392ms step_avg:34.99ms step:298/1575 train_time:10430ms step_avg:35.00ms step:299/1575 train_time:10461ms step_avg:34.99ms step:300/1575 train_time:10499ms step_avg:35.00ms step:301/1575 train_time:10530ms step_avg:34.98ms step:302/1575 train_time:10569ms step_avg:35.00ms step:303/1575 train_time:10600ms step_avg:34.98ms step:304/1575 train_time:10638ms step_avg:34.99ms step:305/1575 train_time:10669ms step_avg:34.98ms step:306/1575 train_time:10707ms step_avg:34.99ms step:307/1575 train_time:10738ms step_avg:34.98ms step:308/1575 train_time:10777ms step_avg:34.99ms step:309/1575 train_time:10808ms step_avg:34.98ms step:310/1575 train_time:10846ms step_avg:34.99ms step:311/1575 train_time:10877ms step_avg:34.97ms step:312/1575 train_time:10915ms step_avg:34.98ms step:313/1575 train_time:10946ms step_avg:34.97ms step:314/1575 train_time:10984ms step_avg:34.98ms step:315/1575 train_time:11015ms step_avg:34.97ms step:316/1575 train_time:11054ms step_avg:34.98ms step:317/1575 train_time:11085ms step_avg:34.97ms step:318/1575 train_time:11123ms step_avg:34.98ms step:319/1575 train_time:11154ms step_avg:34.97ms step:320/1575 train_time:11193ms step_avg:34.98ms step:321/1575 train_time:11223ms step_avg:34.96ms step:322/1575 train_time:11262ms step_avg:34.98ms step:323/1575 train_time:11293ms step_avg:34.96ms step:324/1575 train_time:11331ms step_avg:34.97ms step:325/1575 train_time:11362ms step_avg:34.96ms step:326/1575 train_time:11400ms step_avg:34.97ms step:327/1575 train_time:11431ms step_avg:34.96ms step:328/1575 train_time:11470ms step_avg:34.97ms step:329/1575 train_time:11501ms step_avg:34.96ms step:330/1575 train_time:11539ms step_avg:34.97ms step:331/1575 train_time:11570ms step_avg:34.96ms step:332/1575 train_time:11608ms step_avg:34.96ms step:333/1575 train_time:11639ms step_avg:34.95ms step:334/1575 train_time:11677ms step_avg:34.96ms step:335/1575 train_time:11708ms step_avg:34.95ms step:336/1575 train_time:11747ms step_avg:34.96ms step:337/1575 train_time:11778ms step_avg:34.95ms step:338/1575 train_time:11816ms step_avg:34.96ms step:339/1575 train_time:11847ms step_avg:34.95ms step:340/1575 train_time:11886ms step_avg:34.96ms step:341/1575 train_time:11917ms step_avg:34.95ms step:342/1575 train_time:11955ms step_avg:34.96ms step:343/1575 train_time:11986ms step_avg:34.94ms step:344/1575 train_time:12024ms step_avg:34.95ms step:345/1575 train_time:12056ms step_avg:34.94ms step:346/1575 train_time:12094ms step_avg:34.95ms step:347/1575 train_time:12125ms step_avg:34.94ms step:348/1575 train_time:12164ms step_avg:34.95ms step:349/1575 train_time:12195ms step_avg:34.94ms step:350/1575 train_time:12233ms step_avg:34.95ms step:351/1575 train_time:12264ms step_avg:34.94ms step:352/1575 train_time:12302ms step_avg:34.95ms step:353/1575 train_time:12333ms step_avg:34.94ms step:354/1575 train_time:12372ms step_avg:34.95ms step:355/1575 train_time:12402ms step_avg:34.94ms step:356/1575 train_time:12441ms step_avg:34.95ms step:357/1575 train_time:12472ms step_avg:34.94ms step:358/1575 train_time:12510ms step_avg:34.94ms step:359/1575 train_time:12541ms step_avg:34.93ms step:360/1575 train_time:12579ms step_avg:34.94ms step:361/1575 train_time:12610ms step_avg:34.93ms step:362/1575 train_time:12649ms step_avg:34.94ms step:363/1575 train_time:12680ms step_avg:34.93ms step:364/1575 train_time:12718ms step_avg:34.94ms step:365/1575 train_time:12749ms step_avg:34.93ms step:366/1575 train_time:12788ms step_avg:34.94ms step:367/1575 train_time:12819ms step_avg:34.93ms step:368/1575 train_time:12857ms step_avg:34.94ms step:369/1575 train_time:12888ms step_avg:34.93ms step:370/1575 train_time:12926ms step_avg:34.94ms step:371/1575 train_time:12958ms step_avg:34.93ms step:372/1575 train_time:12996ms step_avg:34.93ms step:373/1575 train_time:13027ms step_avg:34.92ms step:374/1575 train_time:13065ms step_avg:34.93ms step:375/1575 train_time:13096ms step_avg:34.92ms step:376/1575 train_time:13135ms step_avg:34.93ms step:377/1575 train_time:13166ms step_avg:34.92ms step:378/1575 train_time:13204ms step_avg:34.93ms step:379/1575 train_time:13236ms step_avg:34.92ms step:380/1575 train_time:13274ms step_avg:34.93ms step:381/1575 train_time:13305ms step_avg:34.92ms step:382/1575 train_time:13343ms step_avg:34.93ms step:383/1575 train_time:13375ms step_avg:34.92ms step:384/1575 train_time:13413ms step_avg:34.93ms step:385/1575 train_time:13444ms step_avg:34.92ms step:386/1575 train_time:13482ms step_avg:34.93ms step:387/1575 train_time:13513ms step_avg:34.92ms step:388/1575 train_time:13551ms step_avg:34.93ms step:389/1575 train_time:13582ms step_avg:34.92ms step:390/1575 train_time:13621ms step_avg:34.93ms step:391/1575 train_time:13652ms step_avg:34.92ms step:392/1575 train_time:13690ms step_avg:34.92ms step:393/1575 train_time:13721ms step_avg:34.91ms step:394/1575 train_time:13760ms step_avg:34.92ms step:395/1575 train_time:13791ms step_avg:34.91ms step:396/1575 train_time:13829ms step_avg:34.92ms step:397/1575 train_time:13860ms step_avg:34.91ms step:398/1575 train_time:13898ms step_avg:34.92ms step:399/1575 train_time:13929ms step_avg:34.91ms step:400/1575 train_time:13968ms step_avg:34.92ms step:401/1575 train_time:13999ms step_avg:34.91ms step:402/1575 train_time:14038ms step_avg:34.92ms step:403/1575 train_time:14069ms step_avg:34.91ms step:404/1575 train_time:14107ms step_avg:34.92ms step:405/1575 train_time:14138ms step_avg:34.91ms step:406/1575 train_time:14176ms step_avg:34.92ms step:407/1575 train_time:14207ms step_avg:34.91ms step:408/1575 train_time:14246ms step_avg:34.92ms step:409/1575 train_time:14277ms step_avg:34.91ms step:410/1575 train_time:14315ms step_avg:34.91ms step:411/1575 train_time:14346ms step_avg:34.91ms step:412/1575 train_time:14385ms step_avg:34.92ms step:413/1575 train_time:14416ms step_avg:34.91ms step:414/1575 train_time:14454ms step_avg:34.91ms step:415/1575 train_time:14485ms step_avg:34.90ms step:416/1575 train_time:14524ms step_avg:34.91ms step:417/1575 train_time:14555ms step_avg:34.90ms step:418/1575 train_time:14594ms step_avg:34.91ms step:419/1575 train_time:14624ms step_avg:34.90ms step:420/1575 train_time:14663ms step_avg:34.91ms step:421/1575 train_time:14694ms step_avg:34.90ms step:422/1575 train_time:14732ms step_avg:34.91ms step:423/1575 train_time:14763ms step_avg:34.90ms step:424/1575 train_time:14802ms step_avg:34.91ms step:425/1575 train_time:14833ms step_avg:34.90ms step:426/1575 train_time:14871ms step_avg:34.91ms step:427/1575 train_time:14902ms step_avg:34.90ms step:428/1575 train_time:14940ms step_avg:34.91ms step:429/1575 train_time:14971ms step_avg:34.90ms step:430/1575 train_time:15010ms step_avg:34.91ms step:431/1575 train_time:15041ms step_avg:34.90ms step:432/1575 train_time:15080ms step_avg:34.91ms step:433/1575 train_time:15110ms step_avg:34.90ms step:434/1575 train_time:15149ms step_avg:34.90ms step:435/1575 train_time:15180ms step_avg:34.90ms step:436/1575 train_time:15218ms step_avg:34.90ms step:437/1575 train_time:15249ms step_avg:34.90ms step:438/1575 train_time:15287ms step_avg:34.90ms step:439/1575 train_time:15318ms step_avg:34.89ms step:440/1575 train_time:15357ms step_avg:34.90ms step:441/1575 train_time:15388ms step_avg:34.89ms step:442/1575 train_time:15426ms step_avg:34.90ms step:443/1575 train_time:15457ms step_avg:34.89ms step:444/1575 train_time:15495ms step_avg:34.90ms step:445/1575 train_time:15526ms step_avg:34.89ms step:446/1575 train_time:15564ms step_avg:34.90ms step:447/1575 train_time:15596ms step_avg:34.89ms step:448/1575 train_time:15634ms step_avg:34.90ms step:449/1575 train_time:15665ms step_avg:34.89ms step:450/1575 train_time:15703ms step_avg:34.90ms step:451/1575 train_time:15735ms step_avg:34.89ms step:452/1575 train_time:15773ms step_avg:34.90ms step:453/1575 train_time:15804ms step_avg:34.89ms step:454/1575 train_time:15842ms step_avg:34.90ms step:455/1575 train_time:15873ms step_avg:34.89ms step:456/1575 train_time:15912ms step_avg:34.89ms step:457/1575 train_time:15943ms step_avg:34.89ms step:458/1575 train_time:15981ms step_avg:34.89ms step:459/1575 train_time:16012ms step_avg:34.88ms step:460/1575 train_time:16050ms step_avg:34.89ms step:461/1575 train_time:16081ms step_avg:34.88ms step:462/1575 train_time:16119ms step_avg:34.89ms step:463/1575 train_time:16150ms step_avg:34.88ms step:464/1575 train_time:16189ms step_avg:34.89ms step:465/1575 train_time:16220ms step_avg:34.88ms step:466/1575 train_time:16258ms step_avg:34.89ms step:467/1575 train_time:16290ms step_avg:34.88ms step:468/1575 train_time:16328ms step_avg:34.89ms step:469/1575 train_time:16359ms step_avg:34.88ms step:470/1575 train_time:16398ms step_avg:34.89ms step:471/1575 train_time:16429ms step_avg:34.88ms step:472/1575 train_time:16468ms step_avg:34.89ms step:473/1575 train_time:16499ms step_avg:34.88ms step:474/1575 train_time:16537ms step_avg:34.89ms step:475/1575 train_time:16568ms step_avg:34.88ms step:476/1575 train_time:16606ms step_avg:34.89ms step:477/1575 train_time:16638ms step_avg:34.88ms step:478/1575 train_time:16675ms step_avg:34.89ms step:479/1575 train_time:16707ms step_avg:34.88ms step:480/1575 train_time:16745ms step_avg:34.89ms step:481/1575 train_time:16776ms step_avg:34.88ms step:482/1575 train_time:16815ms step_avg:34.88ms step:483/1575 train_time:16846ms step_avg:34.88ms step:484/1575 train_time:16885ms step_avg:34.89ms step:485/1575 train_time:16916ms step_avg:34.88ms step:486/1575 train_time:16955ms step_avg:34.89ms step:487/1575 train_time:16986ms step_avg:34.88ms step:488/1575 train_time:17024ms step_avg:34.89ms step:489/1575 train_time:17055ms step_avg:34.88ms step:490/1575 train_time:17093ms step_avg:34.88ms step:491/1575 train_time:17124ms step_avg:34.88ms step:492/1575 train_time:17162ms step_avg:34.88ms step:493/1575 train_time:17194ms step_avg:34.88ms step:494/1575 train_time:17232ms step_avg:34.88ms step:495/1575 train_time:17264ms step_avg:34.88ms step:496/1575 train_time:17302ms step_avg:34.88ms step:497/1575 train_time:17333ms step_avg:34.88ms step:498/1575 train_time:17371ms step_avg:34.88ms step:499/1575 train_time:17402ms step_avg:34.87ms step:500/1575 train_time:17441ms step_avg:34.88ms step:500/1575 val_loss:4.2375 train_time:17489ms step_avg:34.98ms step:501/1575 train_time:17507ms step_avg:34.94ms step:502/1575 train_time:17525ms step_avg:34.91ms step:503/1575 train_time:17544ms step_avg:34.88ms step:504/1575 train_time:17584ms step_avg:34.89ms step:505/1575 train_time:17617ms step_avg:34.89ms step:506/1575 train_time:17656ms step_avg:34.89ms step:507/1575 train_time:17688ms step_avg:34.89ms step:508/1575 train_time:17726ms step_avg:34.89ms step:509/1575 train_time:17757ms step_avg:34.89ms step:510/1575 train_time:17796ms step_avg:34.89ms step:511/1575 train_time:17827ms step_avg:34.89ms step:512/1575 train_time:17865ms step_avg:34.89ms step:513/1575 train_time:17935ms step_avg:34.96ms step:514/1575 train_time:17992ms step_avg:35.00ms step:515/1575 train_time:18055ms step_avg:35.06ms step:516/1575 train_time:18113ms step_avg:35.10ms step:517/1575 train_time:18176ms step_avg:35.16ms step:518/1575 train_time:18236ms step_avg:35.20ms step:519/1575 train_time:18299ms step_avg:35.26ms step:520/1575 train_time:18357ms step_avg:35.30ms step:521/1575 train_time:18420ms step_avg:35.36ms step:522/1575 train_time:18480ms step_avg:35.40ms step:523/1575 train_time:18545ms step_avg:35.46ms step:524/1575 train_time:18605ms step_avg:35.51ms step:525/1575 train_time:18669ms step_avg:35.56ms step:526/1575 train_time:18729ms step_avg:35.61ms step:527/1575 train_time:18795ms step_avg:35.66ms step:528/1575 train_time:18854ms step_avg:35.71ms step:529/1575 train_time:18918ms step_avg:35.76ms step:530/1575 train_time:18977ms step_avg:35.81ms step:531/1575 train_time:19041ms step_avg:35.86ms step:532/1575 train_time:19099ms step_avg:35.90ms step:533/1575 train_time:19162ms step_avg:35.95ms step:534/1575 train_time:19222ms step_avg:36.00ms step:535/1575 train_time:19285ms step_avg:36.05ms step:536/1575 train_time:19344ms step_avg:36.09ms step:537/1575 train_time:19406ms step_avg:36.14ms step:538/1575 train_time:19465ms step_avg:36.18ms step:539/1575 train_time:19529ms step_avg:36.23ms step:540/1575 train_time:19588ms step_avg:36.27ms step:541/1575 train_time:19651ms step_avg:36.32ms step:542/1575 train_time:19711ms step_avg:36.37ms step:543/1575 train_time:19775ms step_avg:36.42ms step:544/1575 train_time:19835ms step_avg:36.46ms step:545/1575 train_time:19899ms step_avg:36.51ms step:546/1575 train_time:19958ms step_avg:36.55ms step:547/1575 train_time:20022ms step_avg:36.60ms step:548/1575 train_time:20080ms step_avg:36.64ms step:549/1575 train_time:20144ms step_avg:36.69ms step:550/1575 train_time:20203ms step_avg:36.73ms step:551/1575 train_time:20265ms step_avg:36.78ms step:552/1575 train_time:20324ms step_avg:36.82ms step:553/1575 train_time:20387ms step_avg:36.87ms step:554/1575 train_time:20447ms step_avg:36.91ms step:555/1575 train_time:20510ms step_avg:36.95ms step:556/1575 train_time:20569ms step_avg:36.99ms step:557/1575 train_time:20632ms step_avg:37.04ms step:558/1575 train_time:20691ms step_avg:37.08ms step:559/1575 train_time:20755ms step_avg:37.13ms step:560/1575 train_time:20815ms step_avg:37.17ms step:561/1575 train_time:20880ms step_avg:37.22ms step:562/1575 train_time:20939ms step_avg:37.26ms step:563/1575 train_time:21002ms step_avg:37.30ms step:564/1575 train_time:21062ms step_avg:37.34ms step:565/1575 train_time:21124ms step_avg:37.39ms step:566/1575 train_time:21183ms step_avg:37.43ms step:567/1575 train_time:21246ms step_avg:37.47ms step:568/1575 train_time:21305ms step_avg:37.51ms step:569/1575 train_time:21373ms step_avg:37.56ms step:570/1575 train_time:21430ms step_avg:37.60ms step:571/1575 train_time:21492ms step_avg:37.64ms step:572/1575 train_time:21551ms step_avg:37.68ms step:573/1575 train_time:21614ms step_avg:37.72ms step:574/1575 train_time:21673ms step_avg:37.76ms step:575/1575 train_time:21737ms step_avg:37.80ms step:576/1575 train_time:21796ms step_avg:37.84ms step:577/1575 train_time:21859ms step_avg:37.88ms step:578/1575 train_time:21919ms step_avg:37.92ms step:579/1575 train_time:21982ms step_avg:37.97ms step:580/1575 train_time:22042ms step_avg:38.00ms step:581/1575 train_time:22105ms step_avg:38.05ms step:582/1575 train_time:22164ms step_avg:38.08ms step:583/1575 train_time:22228ms step_avg:38.13ms step:584/1575 train_time:22287ms step_avg:38.16ms step:585/1575 train_time:22350ms step_avg:38.20ms step:586/1575 train_time:22409ms step_avg:38.24ms step:587/1575 train_time:22472ms step_avg:38.28ms step:588/1575 train_time:22531ms step_avg:38.32ms step:589/1575 train_time:22594ms step_avg:38.36ms step:590/1575 train_time:22654ms step_avg:38.40ms step:591/1575 train_time:22717ms step_avg:38.44ms step:592/1575 train_time:22776ms step_avg:38.47ms step:593/1575 train_time:22840ms step_avg:38.52ms step:594/1575 train_time:22899ms step_avg:38.55ms step:595/1575 train_time:22962ms step_avg:38.59ms step:596/1575 train_time:23022ms step_avg:38.63ms step:597/1575 train_time:23085ms step_avg:38.67ms step:598/1575 train_time:23147ms step_avg:38.71ms step:599/1575 train_time:23209ms step_avg:38.75ms step:600/1575 train_time:23268ms step_avg:38.78ms step:601/1575 train_time:23331ms step_avg:38.82ms step:602/1575 train_time:23389ms step_avg:38.85ms step:603/1575 train_time:23453ms step_avg:38.89ms step:604/1575 train_time:23513ms step_avg:38.93ms step:605/1575 train_time:23575ms step_avg:38.97ms step:606/1575 train_time:23635ms step_avg:39.00ms step:607/1575 train_time:23698ms step_avg:39.04ms step:608/1575 train_time:23758ms step_avg:39.07ms step:609/1575 train_time:23821ms step_avg:39.11ms step:610/1575 train_time:23880ms step_avg:39.15ms step:611/1575 train_time:23944ms step_avg:39.19ms step:612/1575 train_time:24004ms step_avg:39.22ms step:613/1575 train_time:24067ms step_avg:39.26ms step:614/1575 train_time:24127ms step_avg:39.29ms step:615/1575 train_time:24190ms step_avg:39.33ms step:616/1575 train_time:24250ms step_avg:39.37ms step:617/1575 train_time:24313ms step_avg:39.41ms step:618/1575 train_time:24373ms step_avg:39.44ms step:619/1575 train_time:24436ms step_avg:39.48ms step:620/1575 train_time:24495ms step_avg:39.51ms step:621/1575 train_time:24558ms step_avg:39.55ms step:622/1575 train_time:24617ms step_avg:39.58ms step:623/1575 train_time:24680ms step_avg:39.61ms step:624/1575 train_time:24739ms step_avg:39.65ms step:625/1575 train_time:24803ms step_avg:39.68ms step:626/1575 train_time:24862ms step_avg:39.72ms step:627/1575 train_time:24925ms step_avg:39.75ms step:628/1575 train_time:24984ms step_avg:39.78ms step:629/1575 train_time:25048ms step_avg:39.82ms step:630/1575 train_time:25107ms step_avg:39.85ms step:631/1575 train_time:25170ms step_avg:39.89ms step:632/1575 train_time:25230ms step_avg:39.92ms step:633/1575 train_time:25294ms step_avg:39.96ms step:634/1575 train_time:25353ms step_avg:39.99ms step:635/1575 train_time:25416ms step_avg:40.03ms step:636/1575 train_time:25475ms step_avg:40.06ms step:637/1575 train_time:25539ms step_avg:40.09ms step:638/1575 train_time:25598ms step_avg:40.12ms step:639/1575 train_time:25662ms step_avg:40.16ms step:640/1575 train_time:25722ms step_avg:40.19ms step:641/1575 train_time:25785ms step_avg:40.23ms step:642/1575 train_time:25845ms step_avg:40.26ms step:643/1575 train_time:25908ms step_avg:40.29ms step:644/1575 train_time:25967ms step_avg:40.32ms step:645/1575 train_time:26031ms step_avg:40.36ms step:646/1575 train_time:26090ms step_avg:40.39ms step:647/1575 train_time:26154ms step_avg:40.42ms step:648/1575 train_time:26215ms step_avg:40.46ms step:649/1575 train_time:26278ms step_avg:40.49ms step:650/1575 train_time:26339ms step_avg:40.52ms step:651/1575 train_time:26401ms step_avg:40.55ms step:652/1575 train_time:26460ms step_avg:40.58ms step:653/1575 train_time:26523ms step_avg:40.62ms step:654/1575 train_time:26582ms step_avg:40.65ms step:655/1575 train_time:26646ms step_avg:40.68ms step:656/1575 train_time:26705ms step_avg:40.71ms step:657/1575 train_time:26768ms step_avg:40.74ms step:658/1575 train_time:26827ms step_avg:40.77ms step:659/1575 train_time:26891ms step_avg:40.81ms step:660/1575 train_time:26950ms step_avg:40.83ms step:661/1575 train_time:27013ms step_avg:40.87ms step:662/1575 train_time:27072ms step_avg:40.89ms step:663/1575 train_time:27135ms step_avg:40.93ms step:664/1575 train_time:27195ms step_avg:40.96ms step:665/1575 train_time:27259ms step_avg:40.99ms step:666/1575 train_time:27318ms step_avg:41.02ms step:667/1575 train_time:27382ms step_avg:41.05ms step:668/1575 train_time:27442ms step_avg:41.08ms step:669/1575 train_time:27504ms step_avg:41.11ms step:670/1575 train_time:27564ms step_avg:41.14ms step:671/1575 train_time:27627ms step_avg:41.17ms step:672/1575 train_time:27692ms step_avg:41.21ms step:673/1575 train_time:27750ms step_avg:41.23ms step:674/1575 train_time:27809ms step_avg:41.26ms step:675/1575 train_time:27872ms step_avg:41.29ms step:676/1575 train_time:27931ms step_avg:41.32ms step:677/1575 train_time:27994ms step_avg:41.35ms step:678/1575 train_time:28054ms step_avg:41.38ms step:679/1575 train_time:28118ms step_avg:41.41ms step:680/1575 train_time:28177ms step_avg:41.44ms step:681/1575 train_time:28241ms step_avg:41.47ms step:682/1575 train_time:28301ms step_avg:41.50ms step:683/1575 train_time:28364ms step_avg:41.53ms step:684/1575 train_time:28423ms step_avg:41.55ms step:685/1575 train_time:28486ms step_avg:41.59ms step:686/1575 train_time:28546ms step_avg:41.61ms step:687/1575 train_time:28609ms step_avg:41.64ms step:688/1575 train_time:28670ms step_avg:41.67ms step:689/1575 train_time:28732ms step_avg:41.70ms step:690/1575 train_time:28791ms step_avg:41.73ms step:691/1575 train_time:28854ms step_avg:41.76ms step:692/1575 train_time:28914ms step_avg:41.78ms step:693/1575 train_time:28977ms step_avg:41.81ms step:694/1575 train_time:29036ms step_avg:41.84ms step:695/1575 train_time:29101ms step_avg:41.87ms step:696/1575 train_time:29160ms step_avg:41.90ms step:697/1575 train_time:29223ms step_avg:41.93ms step:698/1575 train_time:29282ms step_avg:41.95ms step:699/1575 train_time:29346ms step_avg:41.98ms step:700/1575 train_time:29405ms step_avg:42.01ms step:701/1575 train_time:29468ms step_avg:42.04ms step:702/1575 train_time:29527ms step_avg:42.06ms step:703/1575 train_time:29590ms step_avg:42.09ms step:704/1575 train_time:29649ms step_avg:42.11ms step:705/1575 train_time:29712ms step_avg:42.14ms step:706/1575 train_time:29771ms step_avg:42.17ms step:707/1575 train_time:29834ms step_avg:42.20ms step:708/1575 train_time:29894ms step_avg:42.22ms step:709/1575 train_time:29957ms step_avg:42.25ms step:710/1575 train_time:30017ms step_avg:42.28ms step:711/1575 train_time:30081ms step_avg:42.31ms step:712/1575 train_time:30141ms step_avg:42.33ms step:713/1575 train_time:30204ms step_avg:42.36ms step:714/1575 train_time:30263ms step_avg:42.39ms step:715/1575 train_time:30327ms step_avg:42.42ms step:716/1575 train_time:30387ms step_avg:42.44ms step:717/1575 train_time:30450ms step_avg:42.47ms step:718/1575 train_time:30509ms step_avg:42.49ms step:719/1575 train_time:30572ms step_avg:42.52ms step:720/1575 train_time:30631ms step_avg:42.54ms step:721/1575 train_time:30694ms step_avg:42.57ms step:722/1575 train_time:30754ms step_avg:42.60ms step:723/1575 train_time:30818ms step_avg:42.62ms step:724/1575 train_time:30877ms step_avg:42.65ms step:725/1575 train_time:30940ms step_avg:42.68ms step:726/1575 train_time:30999ms step_avg:42.70ms step:727/1575 train_time:31064ms step_avg:42.73ms step:728/1575 train_time:31125ms step_avg:42.75ms step:729/1575 train_time:31186ms step_avg:42.78ms step:730/1575 train_time:31245ms step_avg:42.80ms step:731/1575 train_time:31308ms step_avg:42.83ms step:732/1575 train_time:31368ms step_avg:42.85ms step:733/1575 train_time:31431ms step_avg:42.88ms step:734/1575 train_time:31490ms step_avg:42.90ms step:735/1575 train_time:31554ms step_avg:42.93ms step:736/1575 train_time:31613ms step_avg:42.95ms step:737/1575 train_time:31677ms step_avg:42.98ms step:738/1575 train_time:31736ms step_avg:43.00ms step:739/1575 train_time:31799ms step_avg:43.03ms step:740/1575 train_time:31859ms step_avg:43.05ms step:741/1575 train_time:31922ms step_avg:43.08ms step:742/1575 train_time:31982ms step_avg:43.10ms step:743/1575 train_time:32044ms step_avg:43.13ms step:744/1575 train_time:32103ms step_avg:43.15ms step:745/1575 train_time:32167ms step_avg:43.18ms step:746/1575 train_time:32226ms step_avg:43.20ms step:747/1575 train_time:32289ms step_avg:43.23ms step:748/1575 train_time:32350ms step_avg:43.25ms step:749/1575 train_time:32412ms step_avg:43.27ms step:750/1575 train_time:32471ms step_avg:43.29ms step:750/1575 val_loss:3.8927 train_time:32519ms step_avg:43.36ms step:751/1575 train_time:32537ms step_avg:43.33ms step:752/1575 train_time:32597ms step_avg:43.35ms step:753/1575 train_time:32664ms step_avg:43.38ms step:754/1575 train_time:32725ms step_avg:43.40ms step:755/1575 train_time:32789ms step_avg:43.43ms step:756/1575 train_time:32848ms step_avg:43.45ms step:757/1575 train_time:32910ms step_avg:43.47ms step:758/1575 train_time:32969ms step_avg:43.49ms step:759/1575 train_time:33032ms step_avg:43.52ms step:760/1575 train_time:33091ms step_avg:43.54ms step:761/1575 train_time:33154ms step_avg:43.57ms step:762/1575 train_time:33213ms step_avg:43.59ms step:763/1575 train_time:33275ms step_avg:43.61ms step:764/1575 train_time:33334ms step_avg:43.63ms step:765/1575 train_time:33398ms step_avg:43.66ms step:766/1575 train_time:33457ms step_avg:43.68ms step:767/1575 train_time:33521ms step_avg:43.70ms step:768/1575 train_time:33582ms step_avg:43.73ms step:769/1575 train_time:33647ms step_avg:43.75ms step:770/1575 train_time:33706ms step_avg:43.77ms step:771/1575 train_time:33770ms step_avg:43.80ms step:772/1575 train_time:33830ms step_avg:43.82ms step:773/1575 train_time:33893ms step_avg:43.85ms step:774/1575 train_time:33952ms step_avg:43.87ms step:775/1575 train_time:34015ms step_avg:43.89ms step:776/1575 train_time:34074ms step_avg:43.91ms step:777/1575 train_time:34137ms step_avg:43.93ms step:778/1575 train_time:34196ms step_avg:43.95ms step:779/1575 train_time:34259ms step_avg:43.98ms step:780/1575 train_time:34318ms step_avg:44.00ms step:781/1575 train_time:34381ms step_avg:44.02ms step:782/1575 train_time:34441ms step_avg:44.04ms step:783/1575 train_time:34504ms step_avg:44.07ms step:784/1575 train_time:34564ms step_avg:44.09ms step:785/1575 train_time:34628ms step_avg:44.11ms step:786/1575 train_time:34688ms step_avg:44.13ms step:787/1575 train_time:34751ms step_avg:44.16ms step:788/1575 train_time:34810ms step_avg:44.18ms step:789/1575 train_time:34873ms step_avg:44.20ms step:790/1575 train_time:34932ms step_avg:44.22ms step:791/1575 train_time:34995ms step_avg:44.24ms step:792/1575 train_time:35054ms step_avg:44.26ms step:793/1575 train_time:35118ms step_avg:44.28ms step:794/1575 train_time:35176ms step_avg:44.30ms step:795/1575 train_time:35239ms step_avg:44.33ms step:796/1575 train_time:35298ms step_avg:44.34ms step:797/1575 train_time:35362ms step_avg:44.37ms step:798/1575 train_time:35421ms step_avg:44.39ms step:799/1575 train_time:35486ms step_avg:44.41ms step:800/1575 train_time:35546ms step_avg:44.43ms step:801/1575 train_time:35609ms step_avg:44.46ms step:802/1575 train_time:35669ms step_avg:44.47ms step:803/1575 train_time:35732ms step_avg:44.50ms step:804/1575 train_time:35791ms step_avg:44.52ms step:805/1575 train_time:35855ms step_avg:44.54ms step:806/1575 train_time:35915ms step_avg:44.56ms step:807/1575 train_time:35979ms step_avg:44.58ms step:808/1575 train_time:36037ms step_avg:44.60ms step:809/1575 train_time:36100ms step_avg:44.62ms step:810/1575 train_time:36160ms step_avg:44.64ms step:811/1575 train_time:36224ms step_avg:44.67ms step:812/1575 train_time:36282ms step_avg:44.68ms step:813/1575 train_time:36345ms step_avg:44.71ms step:814/1575 train_time:36404ms step_avg:44.72ms step:815/1575 train_time:36468ms step_avg:44.75ms step:816/1575 train_time:36528ms step_avg:44.76ms step:817/1575 train_time:36590ms step_avg:44.79ms step:818/1575 train_time:36650ms step_avg:44.80ms step:819/1575 train_time:36713ms step_avg:44.83ms step:820/1575 train_time:36772ms step_avg:44.84ms step:821/1575 train_time:36835ms step_avg:44.87ms step:822/1575 train_time:36894ms step_avg:44.88ms step:823/1575 train_time:36958ms step_avg:44.91ms step:824/1575 train_time:37016ms step_avg:44.92ms step:825/1575 train_time:37080ms step_avg:44.95ms step:826/1575 train_time:37139ms step_avg:44.96ms step:827/1575 train_time:37202ms step_avg:44.98ms step:828/1575 train_time:37262ms step_avg:45.00ms step:829/1575 train_time:37326ms step_avg:45.02ms step:830/1575 train_time:37385ms step_avg:45.04ms step:831/1575 train_time:37448ms step_avg:45.06ms step:832/1575 train_time:37507ms step_avg:45.08ms step:833/1575 train_time:37571ms step_avg:45.10ms step:834/1575 train_time:37630ms step_avg:45.12ms step:835/1575 train_time:37694ms step_avg:45.14ms step:836/1575 train_time:37754ms step_avg:45.16ms step:837/1575 train_time:37816ms step_avg:45.18ms step:838/1575 train_time:37875ms step_avg:45.20ms step:839/1575 train_time:37938ms step_avg:45.22ms step:840/1575 train_time:37997ms step_avg:45.23ms step:841/1575 train_time:38061ms step_avg:45.26ms step:842/1575 train_time:38120ms step_avg:45.27ms step:843/1575 train_time:38184ms step_avg:45.30ms step:844/1575 train_time:38243ms step_avg:45.31ms step:845/1575 train_time:38307ms step_avg:45.33ms step:846/1575 train_time:38366ms step_avg:45.35ms step:847/1575 train_time:38429ms step_avg:45.37ms step:848/1575 train_time:38489ms step_avg:45.39ms step:849/1575 train_time:38552ms step_avg:45.41ms step:850/1575 train_time:38611ms step_avg:45.43ms step:851/1575 train_time:38674ms step_avg:45.45ms step:852/1575 train_time:38735ms step_avg:45.46ms step:853/1575 train_time:38798ms step_avg:45.48ms step:854/1575 train_time:38855ms step_avg:45.50ms step:855/1575 train_time:38918ms step_avg:45.52ms step:856/1575 train_time:38978ms step_avg:45.54ms step:857/1575 train_time:39041ms step_avg:45.56ms step:858/1575 train_time:39101ms step_avg:45.57ms step:859/1575 train_time:39165ms step_avg:45.59ms step:860/1575 train_time:39224ms step_avg:45.61ms step:861/1575 train_time:39287ms step_avg:45.63ms step:862/1575 train_time:39347ms step_avg:45.65ms step:863/1575 train_time:39411ms step_avg:45.67ms step:864/1575 train_time:39470ms step_avg:45.68ms step:865/1575 train_time:39533ms step_avg:45.70ms step:866/1575 train_time:39593ms step_avg:45.72ms step:867/1575 train_time:39656ms step_avg:45.74ms step:868/1575 train_time:39716ms step_avg:45.76ms step:869/1575 train_time:39779ms step_avg:45.78ms step:870/1575 train_time:39838ms step_avg:45.79ms step:871/1575 train_time:39902ms step_avg:45.81ms step:872/1575 train_time:39962ms step_avg:45.83ms step:873/1575 train_time:40025ms step_avg:45.85ms step:874/1575 train_time:40084ms step_avg:45.86ms step:875/1575 train_time:40147ms step_avg:45.88ms step:876/1575 train_time:40206ms step_avg:45.90ms step:877/1575 train_time:40270ms step_avg:45.92ms step:878/1575 train_time:40329ms step_avg:45.93ms step:879/1575 train_time:40395ms step_avg:45.96ms step:880/1575 train_time:40454ms step_avg:45.97ms step:881/1575 train_time:40515ms step_avg:45.99ms step:882/1575 train_time:40575ms step_avg:46.00ms step:883/1575 train_time:40638ms step_avg:46.02ms step:884/1575 train_time:40697ms step_avg:46.04ms step:885/1575 train_time:40761ms step_avg:46.06ms step:886/1575 train_time:40825ms step_avg:46.08ms step:887/1575 train_time:40886ms step_avg:46.09ms step:888/1575 train_time:40945ms step_avg:46.11ms step:889/1575 train_time:41008ms step_avg:46.13ms step:890/1575 train_time:41067ms step_avg:46.14ms step:891/1575 train_time:41130ms step_avg:46.16ms step:892/1575 train_time:41188ms step_avg:46.18ms step:893/1575 train_time:41251ms step_avg:46.19ms step:894/1575 train_time:41311ms step_avg:46.21ms step:895/1575 train_time:41374ms step_avg:46.23ms step:896/1575 train_time:41433ms step_avg:46.24ms step:897/1575 train_time:41496ms step_avg:46.26ms step:898/1575 train_time:41555ms step_avg:46.27ms step:899/1575 train_time:41618ms step_avg:46.29ms step:900/1575 train_time:41678ms step_avg:46.31ms step:901/1575 train_time:41741ms step_avg:46.33ms step:902/1575 train_time:41800ms step_avg:46.34ms step:903/1575 train_time:41864ms step_avg:46.36ms step:904/1575 train_time:41923ms step_avg:46.37ms step:905/1575 train_time:41986ms step_avg:46.39ms step:906/1575 train_time:42046ms step_avg:46.41ms step:907/1575 train_time:42109ms step_avg:46.43ms step:908/1575 train_time:42169ms step_avg:46.44ms step:909/1575 train_time:42232ms step_avg:46.46ms step:910/1575 train_time:42291ms step_avg:46.47ms step:911/1575 train_time:42354ms step_avg:46.49ms step:912/1575 train_time:42413ms step_avg:46.51ms step:913/1575 train_time:42477ms step_avg:46.52ms step:914/1575 train_time:42536ms step_avg:46.54ms step:915/1575 train_time:42600ms step_avg:46.56ms step:916/1575 train_time:42659ms step_avg:46.57ms step:917/1575 train_time:42722ms step_avg:46.59ms step:918/1575 train_time:42781ms step_avg:46.60ms step:919/1575 train_time:42845ms step_avg:46.62ms step:920/1575 train_time:42904ms step_avg:46.63ms step:921/1575 train_time:42967ms step_avg:46.65ms step:922/1575 train_time:43026ms step_avg:46.67ms step:923/1575 train_time:43089ms step_avg:46.68ms step:924/1575 train_time:43149ms step_avg:46.70ms step:925/1575 train_time:43212ms step_avg:46.72ms step:926/1575 train_time:43271ms step_avg:46.73ms step:927/1575 train_time:43335ms step_avg:46.75ms step:928/1575 train_time:43394ms step_avg:46.76ms step:929/1575 train_time:43458ms step_avg:46.78ms step:930/1575 train_time:43517ms step_avg:46.79ms step:931/1575 train_time:43580ms step_avg:46.81ms step:932/1575 train_time:43640ms step_avg:46.82ms step:933/1575 train_time:43703ms step_avg:46.84ms step:934/1575 train_time:43763ms step_avg:46.86ms step:935/1575 train_time:43827ms step_avg:46.87ms step:936/1575 train_time:43885ms step_avg:46.89ms step:937/1575 train_time:43949ms step_avg:46.90ms step:938/1575 train_time:44008ms step_avg:46.92ms step:939/1575 train_time:44072ms step_avg:46.93ms step:940/1575 train_time:44131ms step_avg:46.95ms step:941/1575 train_time:44194ms step_avg:46.97ms step:942/1575 train_time:44253ms step_avg:46.98ms step:943/1575 train_time:44317ms step_avg:47.00ms step:944/1575 train_time:44376ms step_avg:47.01ms step:945/1575 train_time:44439ms step_avg:47.03ms step:946/1575 train_time:44498ms step_avg:47.04ms step:947/1575 train_time:44562ms step_avg:47.06ms step:948/1575 train_time:44621ms step_avg:47.07ms step:949/1575 train_time:44685ms step_avg:47.09ms step:950/1575 train_time:44744ms step_avg:47.10ms step:951/1575 train_time:44808ms step_avg:47.12ms step:952/1575 train_time:44867ms step_avg:47.13ms step:953/1575 train_time:44930ms step_avg:47.15ms step:954/1575 train_time:44989ms step_avg:47.16ms step:955/1575 train_time:45052ms step_avg:47.17ms step:956/1575 train_time:45111ms step_avg:47.19ms step:957/1575 train_time:45176ms step_avg:47.21ms step:958/1575 train_time:45236ms step_avg:47.22ms step:959/1575 train_time:45299ms step_avg:47.24ms step:960/1575 train_time:45358ms step_avg:47.25ms step:961/1575 train_time:45420ms step_avg:47.26ms step:962/1575 train_time:45479ms step_avg:47.28ms step:963/1575 train_time:45543ms step_avg:47.29ms step:964/1575 train_time:45602ms step_avg:47.30ms step:965/1575 train_time:45666ms step_avg:47.32ms step:966/1575 train_time:45725ms step_avg:47.33ms step:967/1575 train_time:45788ms step_avg:47.35ms step:968/1575 train_time:45847ms step_avg:47.36ms step:969/1575 train_time:45911ms step_avg:47.38ms step:970/1575 train_time:45970ms step_avg:47.39ms step:971/1575 train_time:46033ms step_avg:47.41ms step:972/1575 train_time:46093ms step_avg:47.42ms step:973/1575 train_time:46157ms step_avg:47.44ms step:974/1575 train_time:46216ms step_avg:47.45ms step:975/1575 train_time:46279ms step_avg:47.47ms step:976/1575 train_time:46338ms step_avg:47.48ms step:977/1575 train_time:46402ms step_avg:47.49ms step:978/1575 train_time:46461ms step_avg:47.51ms step:979/1575 train_time:46526ms step_avg:47.52ms step:980/1575 train_time:46585ms step_avg:47.54ms step:981/1575 train_time:46648ms step_avg:47.55ms step:982/1575 train_time:46707ms step_avg:47.56ms step:983/1575 train_time:46771ms step_avg:47.58ms step:984/1575 train_time:46831ms step_avg:47.59ms step:985/1575 train_time:46894ms step_avg:47.61ms step:986/1575 train_time:46953ms step_avg:47.62ms step:987/1575 train_time:47017ms step_avg:47.64ms step:988/1575 train_time:47076ms step_avg:47.65ms step:989/1575 train_time:47140ms step_avg:47.66ms step:990/1575 train_time:47199ms step_avg:47.68ms step:991/1575 train_time:47263ms step_avg:47.69ms step:992/1575 train_time:47322ms step_avg:47.70ms step:993/1575 train_time:47386ms step_avg:47.72ms step:994/1575 train_time:47445ms step_avg:47.73ms step:995/1575 train_time:47508ms step_avg:47.75ms step:996/1575 train_time:47567ms step_avg:47.76ms step:997/1575 train_time:47630ms step_avg:47.77ms step:998/1575 train_time:47690ms step_avg:47.79ms step:999/1575 train_time:47753ms step_avg:47.80ms step:1000/1575 train_time:47812ms step_avg:47.81ms step:1000/1575 val_loss:3.5812 train_time:47861ms step_avg:47.86ms step:1001/1575 train_time:47879ms step_avg:47.83ms step:1002/1575 train_time:47937ms step_avg:47.84ms step:1003/1575 train_time:48003ms step_avg:47.86ms step:1004/1575 train_time:48064ms step_avg:47.87ms step:1005/1575 train_time:48127ms step_avg:47.89ms step:1006/1575 train_time:48186ms step_avg:47.90ms step:1007/1575 train_time:48250ms step_avg:47.91ms step:1008/1575 train_time:48309ms step_avg:47.93ms step:1009/1575 train_time:48372ms step_avg:47.94ms step:1010/1575 train_time:48431ms step_avg:47.95ms step:1011/1575 train_time:48494ms step_avg:47.97ms step:1012/1575 train_time:48553ms step_avg:47.98ms step:1013/1575 train_time:48615ms step_avg:47.99ms step:1014/1575 train_time:48674ms step_avg:48.00ms step:1015/1575 train_time:48736ms step_avg:48.02ms step:1016/1575 train_time:48796ms step_avg:48.03ms step:1017/1575 train_time:48859ms step_avg:48.04ms step:1018/1575 train_time:48919ms step_avg:48.05ms step:1019/1575 train_time:48984ms step_avg:48.07ms step:1020/1575 train_time:49043ms step_avg:48.08ms step:1021/1575 train_time:49108ms step_avg:48.10ms step:1022/1575 train_time:49167ms step_avg:48.11ms step:1023/1575 train_time:49230ms step_avg:48.12ms step:1024/1575 train_time:49289ms step_avg:48.13ms step:1025/1575 train_time:49361ms step_avg:48.16ms step:1026/1575 train_time:49444ms step_avg:48.19ms step:1027/1575 train_time:49535ms step_avg:48.23ms step:1028/1575 train_time:49621ms step_avg:48.27ms step:1029/1575 train_time:49711ms step_avg:48.31ms step:1030/1575 train_time:49796ms step_avg:48.35ms step:1031/1575 train_time:49888ms step_avg:48.39ms step:1032/1575 train_time:49974ms step_avg:48.42ms step:1033/1575 train_time:50064ms step_avg:48.46ms step:1034/1575 train_time:50150ms step_avg:48.50ms step:1035/1575 train_time:50239ms step_avg:48.54ms step:1036/1575 train_time:50325ms step_avg:48.58ms step:1037/1575 train_time:50414ms step_avg:48.62ms step:1038/1575 train_time:50499ms step_avg:48.65ms step:1039/1575 train_time:50589ms step_avg:48.69ms step:1040/1575 train_time:50675ms step_avg:48.73ms step:1041/1575 train_time:50765ms 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:51116ms step_avg:48.92ms step:1046/1575 train_time:51202ms step_avg:48.95ms step:1047/1575 train_time:51292ms step_avg:48.99ms step:1048/1575 train_time:51378ms step_avg:49.02ms step:1049/1575 train_time:51467ms 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:51903ms step_avg:49.24ms step:1055/1575 train_time:51993ms step_avg:49.28ms step:1056/1575 train_time:52079ms step_avg:49.32ms step:1057/1575 train_time:52168ms step_avg:49.36ms step:1058/1575 train_time:52255ms step_avg:49.39ms step:1059/1575 train_time:52345ms step_avg:49.43ms step:1060/1575 train_time:52430ms step_avg:49.46ms step:1061/1575 train_time:52520ms step_avg:49.50ms step:1062/1575 train_time:52605ms step_avg:49.53ms step:1063/1575 train_time:52694ms step_avg:49.57ms step:1064/1575 train_time:52780ms step_avg:49.61ms 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:53045ms step_avg:49.71ms step:1068/1575 train_time:53130ms step_avg:49.75ms step:1069/1575 train_time:53220ms step_avg:49.78ms 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.92ms 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:53832ms step_avg:50.03ms step:1077/1575 train_time:53922ms step_avg:50.07ms step:1078/1575 train_time:54008ms step_avg:50.10ms step:1079/1575 train_time:54098ms step_avg:50.14ms step:1080/1575 train_time:54184ms step_avg:50.17ms step:1081/1575 train_time:54273ms step_avg:50.21ms step:1082/1575 train_time:54359ms step_avg:50.24ms step:1083/1575 train_time:54448ms step_avg:50.28ms step:1084/1575 train_time:54538ms step_avg:50.31ms step:1085/1575 train_time:54626ms step_avg:50.35ms step:1086/1575 train_time:54711ms step_avg:50.38ms step:1087/1575 train_time:54801ms step_avg:50.41ms step:1088/1575 train_time:54886ms step_avg:50.45ms step:1089/1575 train_time:54975ms step_avg:50.48ms step:1090/1575 train_time:55061ms step_avg:50.51ms step:1091/1575 train_time:55151ms step_avg:50.55ms step:1092/1575 train_time:55237ms step_avg:50.58ms step:1093/1575 train_time:55327ms step_avg:50.62ms step:1094/1575 train_time:55412ms step_avg:50.65ms step:1095/1575 train_time:55502ms step_avg:50.69ms step:1096/1575 train_time:55587ms step_avg:50.72ms step:1097/1575 train_time:55677ms step_avg:50.75ms step:1098/1575 train_time:55763ms step_avg:50.79ms step:1099/1575 train_time:55852ms step_avg:50.82ms step:1100/1575 train_time:55938ms step_avg:50.85ms step:1101/1575 train_time:56028ms step_avg:50.89ms step:1102/1575 train_time:56113ms step_avg:50.92ms step:1103/1575 train_time:56203ms step_avg:50.95ms step:1104/1575 train_time:56289ms step_avg:50.99ms step:1105/1575 train_time:56378ms step_avg:51.02ms step:1106/1575 train_time:56464ms step_avg:51.05ms step:1107/1575 train_time:56554ms step_avg:51.09ms step:1108/1575 train_time:56640ms step_avg:51.12ms step:1109/1575 train_time:56730ms step_avg:51.15ms step:1110/1575 train_time:56816ms step_avg:51.19ms step:1111/1575 train_time:56906ms step_avg:51.22ms step:1112/1575 train_time:56991ms step_avg:51.25ms step:1113/1575 train_time:57081ms step_avg:51.29ms step:1114/1575 train_time:57167ms step_avg:51.32ms step:1115/1575 train_time:57256ms step_avg:51.35ms step:1116/1575 train_time:57342ms step_avg:51.38ms step:1117/1575 train_time:57431ms step_avg:51.42ms step:1118/1575 train_time:57517ms step_avg:51.45ms step:1119/1575 train_time:57606ms step_avg:51.48ms step:1120/1575 train_time:57693ms step_avg:51.51ms step:1121/1575 train_time:57782ms step_avg:51.55ms step:1122/1575 train_time:57868ms step_avg:51.58ms step:1123/1575 train_time:57957ms step_avg:51.61ms step:1124/1575 train_time:58043ms step_avg:51.64ms step:1125/1575 train_time:58132ms step_avg:51.67ms step:1126/1575 train_time:58221ms step_avg:51.71ms step:1127/1575 train_time:58309ms step_avg:51.74ms step:1128/1575 train_time:58400ms step_avg:51.77ms step:1129/1575 train_time:58489ms step_avg:51.81ms step:1130/1575 train_time:58570ms step_avg:51.83ms step:1131/1575 train_time:58659ms step_avg:51.86ms step:1132/1575 train_time:58744ms step_avg:51.89ms step:1133/1575 train_time:58833ms step_avg:51.93ms step:1134/1575 train_time:58919ms step_avg:51.96ms step:1135/1575 train_time:59008ms step_avg:51.99ms step:1136/1575 train_time:59093ms step_avg:52.02ms step:1137/1575 train_time:59183ms step_avg:52.05ms step:1138/1575 train_time:59269ms step_avg:52.08ms step:1139/1575 train_time:59361ms step_avg:52.12ms step:1140/1575 train_time:59445ms step_avg:52.14ms step:1141/1575 train_time:59536ms step_avg:52.18ms step:1142/1575 train_time:59620ms step_avg:52.21ms step:1143/1575 train_time:59709ms step_avg:52.24ms step:1144/1575 train_time:59794ms step_avg:52.27ms step:1145/1575 train_time:59884ms step_avg:52.30ms step:1146/1575 train_time:59975ms step_avg:52.33ms step:1147/1575 train_time:60063ms step_avg:52.37ms step:1148/1575 train_time:60148ms step_avg:52.39ms step:1149/1575 train_time:60237ms step_avg:52.43ms step:1150/1575 train_time:60322ms step_avg:52.45ms step:1151/1575 train_time:60411ms step_avg:52.49ms step:1152/1575 train_time:60497ms step_avg:52.52ms step:1153/1575 train_time:60587ms step_avg:52.55ms step:1154/1575 train_time:60672ms step_avg:52.58ms step:1155/1575 train_time:60761ms step_avg:52.61ms step:1156/1575 train_time:60847ms step_avg:52.64ms step:1157/1575 train_time:60937ms step_avg:52.67ms step:1158/1575 train_time:61023ms step_avg:52.70ms step:1159/1575 train_time:61113ms step_avg:52.73ms step:1160/1575 train_time:61198ms step_avg:52.76ms step:1161/1575 train_time:61288ms step_avg:52.79ms step:1162/1575 train_time:61374ms step_avg:52.82ms step:1163/1575 train_time:61463ms step_avg:52.85ms step:1164/1575 train_time:61549ms step_avg:52.88ms step:1165/1575 train_time:61638ms step_avg:52.91ms step:1166/1575 train_time:61724ms step_avg:52.94ms step:1167/1575 train_time:61813ms step_avg:52.97ms step:1168/1575 train_time:61899ms step_avg:53.00ms step:1169/1575 train_time:61989ms step_avg:53.03ms step:1170/1575 train_time:62075ms step_avg:53.06ms step:1171/1575 train_time:62164ms step_avg:53.09ms step:1172/1575 train_time:62249ms step_avg:53.11ms step:1173/1575 train_time:62339ms step_avg:53.15ms step:1174/1575 train_time:62425ms step_avg:53.17ms step:1175/1575 train_time:62514ms step_avg:53.20ms step:1176/1575 train_time:62600ms step_avg:53.23ms step:1177/1575 train_time:62688ms step_avg:53.26ms step:1178/1575 train_time:62774ms step_avg:53.29ms step:1179/1575 train_time:62864ms step_avg:53.32ms step:1180/1575 train_time:62950ms step_avg:53.35ms step:1181/1575 train_time:63039ms step_avg:53.38ms step:1182/1575 train_time:63124ms step_avg:53.40ms step:1183/1575 train_time:63214ms step_avg:53.44ms step:1184/1575 train_time:63301ms step_avg:53.46ms step:1185/1575 train_time:63389ms step_avg:53.49ms step:1186/1575 train_time:63475ms step_avg:53.52ms step:1187/1575 train_time:63565ms step_avg:53.55ms step:1188/1575 train_time:63651ms step_avg:53.58ms step:1189/1575 train_time:63742ms step_avg:53.61ms step:1190/1575 train_time:63827ms step_avg:53.64ms step:1191/1575 train_time:63916ms step_avg:53.67ms step:1192/1575 train_time:64002ms step_avg:53.69ms step:1193/1575 train_time:64090ms step_avg:53.72ms step:1194/1575 train_time:64177ms step_avg:53.75ms step:1195/1575 train_time:64266ms step_avg:53.78ms step:1196/1575 train_time:64352ms step_avg:53.81ms step:1197/1575 train_time:64443ms step_avg:53.84ms step:1198/1575 train_time:64528ms step_avg:53.86ms step:1199/1575 train_time:64618ms step_avg:53.89ms step:1200/1575 train_time:64703ms step_avg:53.92ms step:1201/1575 train_time:64792ms step_avg:53.95ms step:1202/1575 train_time:64879ms step_avg:53.98ms step:1203/1575 train_time:64968ms step_avg:54.01ms step:1204/1575 train_time:65054ms step_avg:54.03ms step:1205/1575 train_time:65143ms step_avg:54.06ms step:1206/1575 train_time:65230ms step_avg:54.09ms step:1207/1575 train_time:65320ms step_avg:54.12ms step:1208/1575 train_time:65405ms step_avg:54.14ms step:1209/1575 train_time:65496ms step_avg:54.17ms step:1210/1575 train_time:65582ms step_avg:54.20ms step:1211/1575 train_time:65671ms step_avg:54.23ms step:1212/1575 train_time:65758ms step_avg:54.26ms step:1213/1575 train_time:65847ms step_avg:54.28ms step:1214/1575 train_time:65932ms step_avg:54.31ms step:1215/1575 train_time:66023ms step_avg:54.34ms step:1216/1575 train_time:66108ms step_avg:54.37ms step:1217/1575 train_time:66197ms step_avg:54.39ms step:1218/1575 train_time:66284ms step_avg:54.42ms step:1219/1575 train_time:66373ms step_avg:54.45ms step:1220/1575 train_time:66459ms step_avg:54.47ms step:1221/1575 train_time:66548ms step_avg:54.50ms step:1222/1575 train_time:66634ms step_avg:54.53ms step:1223/1575 train_time:66724ms step_avg:54.56ms step:1224/1575 train_time:66809ms step_avg:54.58ms step:1225/1575 train_time:66900ms step_avg:54.61ms step:1226/1575 train_time:66985ms step_avg:54.64ms step:1227/1575 train_time:67074ms step_avg:54.66ms step:1228/1575 train_time:67161ms step_avg:54.69ms step:1229/1575 train_time:67250ms step_avg:54.72ms step:1230/1575 train_time:67335ms step_avg:54.74ms step:1231/1575 train_time:67425ms step_avg:54.77ms step:1232/1575 train_time:67510ms step_avg:54.80ms step:1233/1575 train_time:67599ms step_avg:54.83ms step:1234/1575 train_time:67685ms step_avg:54.85ms step:1235/1575 train_time:67775ms step_avg:54.88ms step:1236/1575 train_time:67861ms step_avg:54.90ms step:1237/1575 train_time:67950ms step_avg:54.93ms step:1238/1575 train_time:68036ms step_avg:54.96ms step:1239/1575 train_time:68126ms step_avg:54.99ms step:1240/1575 train_time:68212ms step_avg:55.01ms step:1241/1575 train_time:68301ms step_avg:55.04ms step:1242/1575 train_time:68387ms step_avg:55.06ms step:1243/1575 train_time:68477ms step_avg:55.09ms step:1244/1575 train_time:68562ms step_avg:55.11ms step:1245/1575 train_time:68652ms step_avg:55.14ms step:1246/1575 train_time:68738ms step_avg:55.17ms step:1247/1575 train_time:68828ms step_avg:55.19ms step:1248/1575 train_time:68914ms step_avg:55.22ms step:1249/1575 train_time:69004ms step_avg:55.25ms step:1250/1575 train_time:69089ms step_avg:55.27ms step:1250/1575 val_loss:3.4062 train_time:69165ms step_avg:55.33ms step:1251/1575 train_time:69184ms step_avg:55.30ms step:1252/1575 train_time:69269ms step_avg:55.33ms step:1253/1575 train_time:69363ms step_avg:55.36ms step:1254/1575 train_time:69449ms step_avg:55.38ms step:1255/1575 train_time:69538ms step_avg:55.41ms step:1256/1575 train_time:69623ms step_avg:55.43ms step:1257/1575 train_time:69711ms step_avg:55.46ms step:1258/1575 train_time:69796ms step_avg:55.48ms step:1259/1575 train_time:69884ms step_avg:55.51ms step:1260/1575 train_time:69970ms step_avg:55.53ms step:1261/1575 train_time:70059ms step_avg:55.56ms step:1262/1575 train_time:70145ms step_avg:55.58ms step:1263/1575 train_time:70240ms step_avg:55.61ms step:1264/1575 train_time:70327ms step_avg:55.64ms step:1265/1575 train_time:70417ms step_avg:55.67ms step:1266/1575 train_time:70503ms step_avg:55.69ms step:1267/1575 train_time:70591ms step_avg:55.72ms step:1268/1575 train_time:70676ms step_avg:55.74ms step:1269/1575 train_time:70765ms step_avg:55.76ms step:1270/1575 train_time:70849ms step_avg:55.79ms step:1271/1575 train_time:70938ms step_avg:55.81ms step:1272/1575 train_time:71023ms step_avg:55.84ms step:1273/1575 train_time:71113ms step_avg:55.86ms step:1274/1575 train_time:71200ms step_avg:55.89ms step:1275/1575 train_time:71291ms step_avg:55.91ms step:1276/1575 train_time:71378ms step_avg:55.94ms step:1277/1575 train_time:71468ms step_avg:55.97ms step:1278/1575 train_time:71553ms step_avg:55.99ms step:1279/1575 train_time:71643ms step_avg:56.01ms step:1280/1575 train_time:71728ms step_avg:56.04ms step:1281/1575 train_time:71816ms step_avg:56.06ms step:1282/1575 train_time:71901ms step_avg:56.08ms step:1283/1575 train_time:71990ms step_avg:56.11ms step:1284/1575 train_time:72076ms step_avg:56.13ms step:1285/1575 train_time:72166ms step_avg:56.16ms step:1286/1575 train_time:72253ms step_avg:56.18ms step:1287/1575 train_time:72343ms step_avg:56.21ms step:1288/1575 train_time:72430ms step_avg:56.23ms step:1289/1575 train_time:72520ms step_avg:56.26ms step:1290/1575 train_time:72606ms step_avg:56.28ms step:1291/1575 train_time:72695ms step_avg:56.31ms step:1292/1575 train_time:72780ms step_avg:56.33ms step:1293/1575 train_time:72868ms step_avg:56.36ms step:1294/1575 train_time:72954ms step_avg:56.38ms step:1295/1575 train_time:73043ms step_avg:56.40ms step:1296/1575 train_time:73129ms step_avg:56.43ms step:1297/1575 train_time:73220ms step_avg:56.45ms step:1298/1575 train_time:73306ms step_avg:56.48ms step:1299/1575 train_time:73397ms step_avg:56.50ms step:1300/1575 train_time:73484ms step_avg:56.53ms step:1301/1575 train_time:73574ms step_avg:56.55ms step:1302/1575 train_time:73658ms step_avg:56.57ms step:1303/1575 train_time:73748ms step_avg:56.60ms step:1304/1575 train_time:73832ms step_avg:56.62ms step:1305/1575 train_time:73921ms step_avg:56.64ms step:1306/1575 train_time:74006ms step_avg:56.67ms step:1307/1575 train_time:74095ms step_avg:56.69ms step:1308/1575 train_time:74181ms step_avg:56.71ms step:1309/1575 train_time:74271ms step_avg:56.74ms step:1310/1575 train_time:74357ms step_avg:56.76ms step:1311/1575 train_time:74447ms step_avg:56.79ms step:1312/1575 train_time:74533ms step_avg:56.81ms step:1313/1575 train_time:74623ms step_avg:56.83ms step:1314/1575 train_time:74709ms step_avg:56.86ms step:1315/1575 train_time:74799ms step_avg:56.88ms step:1316/1575 train_time:74884ms step_avg:56.90ms step:1317/1575 train_time:74972ms step_avg:56.93ms step:1318/1575 train_time:75058ms step_avg:56.95ms step:1319/1575 train_time:75148ms step_avg:56.97ms step:1320/1575 train_time:75234ms step_avg:57.00ms step:1321/1575 train_time:75323ms step_avg:57.02ms step:1322/1575 train_time:75410ms step_avg:57.04ms step:1323/1575 train_time:75500ms step_avg:57.07ms step:1324/1575 train_time:75586ms step_avg:57.09ms step:1325/1575 train_time:75676ms step_avg:57.11ms step:1326/1575 train_time:75761ms step_avg:57.14ms step:1327/1575 train_time:75850ms step_avg:57.16ms step:1328/1575 train_time:75936ms step_avg:57.18ms step:1329/1575 train_time:76025ms step_avg:57.20ms step:1330/1575 train_time:76110ms step_avg:57.23ms step:1331/1575 train_time:76200ms step_avg:57.25ms step:1332/1575 train_time:76286ms step_avg:57.27ms step:1333/1575 train_time:76377ms step_avg:57.30ms step:1334/1575 train_time:76462ms step_avg:57.32ms step:1335/1575 train_time:76552ms step_avg:57.34ms step:1336/1575 train_time:76638ms step_avg:57.36ms step:1337/1575 train_time:76728ms step_avg:57.39ms step:1338/1575 train_time:76813ms step_avg:57.41ms step:1339/1575 train_time:76903ms step_avg:57.43ms step:1340/1575 train_time:76988ms step_avg:57.45ms step:1341/1575 train_time:77078ms step_avg:57.48ms step:1342/1575 train_time:77163ms step_avg:57.50ms step:1343/1575 train_time:77252ms step_avg:57.52ms step:1344/1575 train_time:77338ms step_avg:57.54ms step:1345/1575 train_time:77428ms step_avg:57.57ms step:1346/1575 train_time:77514ms step_avg:57.59ms step:1347/1575 train_time:77604ms step_avg:57.61ms step:1348/1575 train_time:77690ms step_avg:57.63ms step:1349/1575 train_time:77779ms step_avg:57.66ms step:1350/1575 train_time:77864ms step_avg:57.68ms step:1351/1575 train_time:77953ms step_avg:57.70ms step:1352/1575 train_time:78039ms step_avg:57.72ms step:1353/1575 train_time:78128ms step_avg:57.74ms step:1354/1575 train_time:78214ms step_avg:57.77ms step:1355/1575 train_time:78304ms step_avg:57.79ms step:1356/1575 train_time:78389ms step_avg:57.81ms step:1357/1575 train_time:78480ms step_avg:57.83ms step:1358/1575 train_time:78570ms step_avg:57.86ms step:1359/1575 train_time:78658ms step_avg:57.88ms step:1360/1575 train_time:78744ms step_avg:57.90ms step:1361/1575 train_time:78833ms step_avg:57.92ms step:1362/1575 train_time:78918ms step_avg:57.94ms step:1363/1575 train_time:79007ms step_avg:57.97ms step:1364/1575 train_time:79093ms step_avg:57.99ms step:1365/1575 train_time:79182ms step_avg:58.01ms step:1366/1575 train_time:79267ms step_avg:58.03ms step:1367/1575 train_time:79358ms step_avg:58.05ms step:1368/1575 train_time:79443ms step_avg:58.07ms step:1369/1575 train_time:79533ms step_avg:58.10ms step:1370/1575 train_time:79618ms step_avg:58.12ms step:1371/1575 train_time:79708ms step_avg:58.14ms step:1372/1575 train_time:79794ms step_avg:58.16ms step:1373/1575 train_time:79884ms step_avg:58.18ms step:1374/1575 train_time:79970ms step_avg:58.20ms step:1375/1575 train_time:80060ms step_avg:58.23ms step:1376/1575 train_time:80145ms step_avg:58.25ms step:1377/1575 train_time:80235ms step_avg:58.27ms step:1378/1575 train_time:80320ms step_avg:58.29ms step:1379/1575 train_time:80410ms step_avg:58.31ms step:1380/1575 train_time:80496ms step_avg:58.33ms step:1381/1575 train_time:80586ms step_avg:58.35ms step:1382/1575 train_time:80671ms step_avg:58.37ms step:1383/1575 train_time:80761ms step_avg:58.40ms step:1384/1575 train_time:80847ms step_avg:58.42ms step:1385/1575 train_time:80937ms step_avg:58.44ms step:1386/1575 train_time:81023ms step_avg:58.46ms step:1387/1575 train_time:81113ms step_avg:58.48ms step:1388/1575 train_time:81198ms step_avg:58.50ms step:1389/1575 train_time:81288ms step_avg:58.52ms step:1390/1575 train_time:81374ms step_avg:58.54ms step:1391/1575 train_time:81463ms step_avg:58.56ms step:1392/1575 train_time:81549ms step_avg:58.58ms step:1393/1575 train_time:81640ms step_avg:58.61ms step:1394/1575 train_time:81725ms step_avg:58.63ms step:1395/1575 train_time:81815ms step_avg:58.65ms step:1396/1575 train_time:81900ms step_avg:58.67ms step:1397/1575 train_time:81990ms step_avg:58.69ms step:1398/1575 train_time:82076ms step_avg:58.71ms step:1399/1575 train_time:82165ms step_avg:58.73ms step:1400/1575 train_time:82250ms step_avg:58.75ms step:1401/1575 train_time:82340ms step_avg:58.77ms step:1402/1575 train_time:82426ms step_avg:58.79ms step:1403/1575 train_time:82515ms step_avg:58.81ms step:1404/1575 train_time:82601ms step_avg:58.83ms step:1405/1575 train_time:82691ms step_avg:58.85ms step:1406/1575 train_time:82776ms step_avg:58.87ms step:1407/1575 train_time:82866ms step_avg:58.90ms step:1408/1575 train_time:82951ms step_avg:58.91ms step:1409/1575 train_time:83041ms step_avg:58.94ms step:1410/1575 train_time:83127ms step_avg:58.96ms step:1411/1575 train_time:83217ms step_avg:58.98ms step:1412/1575 train_time:83302ms step_avg:59.00ms step:1413/1575 train_time:83391ms step_avg:59.02ms step:1414/1575 train_time:83477ms step_avg:59.04ms step:1415/1575 train_time:83567ms step_avg:59.06ms step:1416/1575 train_time:83655ms step_avg:59.08ms step:1417/1575 train_time:83742ms step_avg:59.10ms step:1418/1575 train_time:83828ms step_avg:59.12ms step:1419/1575 train_time:83917ms step_avg:59.14ms step:1420/1575 train_time:84003ms step_avg:59.16ms step:1421/1575 train_time:84093ms step_avg:59.18ms step:1422/1575 train_time:84179ms step_avg:59.20ms step:1423/1575 train_time:84268ms step_avg:59.22ms step:1424/1575 train_time:84354ms step_avg:59.24ms step:1425/1575 train_time:84443ms step_avg:59.26ms step:1426/1575 train_time:84529ms step_avg:59.28ms step:1427/1575 train_time:84620ms step_avg:59.30ms step:1428/1575 train_time:84705ms step_avg:59.32ms step:1429/1575 train_time:84795ms step_avg:59.34ms step:1430/1575 train_time:84881ms step_avg:59.36ms step:1431/1575 train_time:84970ms step_avg:59.38ms step:1432/1575 train_time:85057ms step_avg:59.40ms step:1433/1575 train_time:85146ms step_avg:59.42ms step:1434/1575 train_time:85232ms step_avg:59.44ms step:1435/1575 train_time:85322ms step_avg:59.46ms step:1436/1575 train_time:85407ms step_avg:59.48ms step:1437/1575 train_time:85497ms step_avg:59.50ms step:1438/1575 train_time:85583ms step_avg:59.52ms step:1439/1575 train_time:85674ms step_avg:59.54ms step:1440/1575 train_time:85759ms step_avg:59.55ms step:1441/1575 train_time:85849ms step_avg:59.58ms step:1442/1575 train_time:85935ms step_avg:59.59ms step:1443/1575 train_time:86025ms step_avg:59.62ms step:1444/1575 train_time:86110ms step_avg:59.63ms step:1445/1575 train_time:86200ms step_avg:59.65ms step:1446/1575 train_time:86285ms step_avg:59.67ms step:1447/1575 train_time:86375ms step_avg:59.69ms step:1448/1575 train_time:86460ms step_avg:59.71ms step:1449/1575 train_time:86550ms step_avg:59.73ms step:1450/1575 train_time:86636ms step_avg:59.75ms step:1451/1575 train_time:86725ms step_avg:59.77ms step:1452/1575 train_time:86812ms step_avg:59.79ms step:1453/1575 train_time:86902ms step_avg:59.81ms step:1454/1575 train_time:86988ms step_avg:59.83ms step:1455/1575 train_time:87077ms step_avg:59.85ms step:1456/1575 train_time:87162ms step_avg:59.86ms step:1457/1575 train_time:87251ms step_avg:59.88ms step:1458/1575 train_time:87337ms step_avg:59.90ms step:1459/1575 train_time:87426ms step_avg:59.92ms step:1460/1575 train_time:87512ms step_avg:59.94ms step:1461/1575 train_time:87602ms step_avg:59.96ms step:1462/1575 train_time:87687ms step_avg:59.98ms step:1463/1575 train_time:87778ms step_avg:60.00ms step:1464/1575 train_time:87864ms step_avg:60.02ms step:1465/1575 train_time:87953ms step_avg:60.04ms step:1466/1575 train_time:88039ms step_avg:60.05ms step:1467/1575 train_time:88129ms step_avg:60.07ms step:1468/1575 train_time:88214ms step_avg:60.09ms step:1469/1575 train_time:88303ms step_avg:60.11ms step:1470/1575 train_time:88388ms step_avg:60.13ms step:1471/1575 train_time:88479ms step_avg:60.15ms step:1472/1575 train_time:88565ms step_avg:60.17ms step:1473/1575 train_time:88655ms step_avg:60.19ms step:1474/1575 train_time:88740ms step_avg:60.20ms step:1475/1575 train_time:88831ms step_avg:60.22ms step:1476/1575 train_time:88916ms step_avg:60.24ms step:1477/1575 train_time:89007ms step_avg:60.26ms step:1478/1575 train_time:89091ms step_avg:60.28ms step:1479/1575 train_time:89180ms step_avg:60.30ms step:1480/1575 train_time:89265ms step_avg:60.31ms step:1481/1575 train_time:89354ms step_avg:60.33ms step:1482/1575 train_time:89439ms step_avg:60.35ms step:1483/1575 train_time:89529ms step_avg:60.37ms step:1484/1575 train_time:89615ms step_avg:60.39ms step:1485/1575 train_time:89704ms step_avg:60.41ms step:1486/1575 train_time:89791ms step_avg:60.42ms step:1487/1575 train_time:89882ms step_avg:60.44ms step:1488/1575 train_time:89967ms step_avg:60.46ms step:1489/1575 train_time:90057ms step_avg:60.48ms step:1490/1575 train_time:90143ms step_avg:60.50ms step:1491/1575 train_time:90232ms step_avg:60.52ms step:1492/1575 train_time:90318ms step_avg:60.53ms step:1493/1575 train_time:90407ms step_avg:60.55ms step:1494/1575 train_time:90493ms step_avg:60.57ms step:1495/1575 train_time:90582ms step_avg:60.59ms step:1496/1575 train_time:90668ms step_avg:60.61ms step:1497/1575 train_time:90758ms step_avg:60.63ms step:1498/1575 train_time:90844ms step_avg:60.64ms step:1499/1575 train_time:90935ms step_avg:60.66ms step:1500/1575 train_time:91020ms step_avg:60.68ms step:1500/1575 val_loss:3.2992 train_time:91094ms step_avg:60.73ms step:1501/1575 train_time:91113ms step_avg:60.70ms step:1502/1575 train_time:91201ms step_avg:60.72ms step:1503/1575 train_time:91293ms step_avg:60.74ms step:1504/1575 train_time:91380ms step_avg:60.76ms step:1505/1575 train_time:91469ms step_avg:60.78ms step:1506/1575 train_time:91554ms step_avg:60.79ms step:1507/1575 train_time:91642ms step_avg:60.81ms step:1508/1575 train_time:91727ms step_avg:60.83ms step:1509/1575 train_time:91815ms step_avg:60.85ms step:1510/1575 train_time:91900ms step_avg:60.86ms step:1511/1575 train_time:91988ms step_avg:60.88ms step:1512/1575 train_time:92074ms step_avg:60.90ms step:1513/1575 train_time:92166ms step_avg:60.92ms step:1514/1575 train_time:92255ms step_avg:60.93ms step:1515/1575 train_time:92344ms step_avg:60.95ms step:1516/1575 train_time:92430ms step_avg:60.97ms step:1517/1575 train_time:92520ms step_avg:60.99ms step:1518/1575 train_time:92605ms step_avg:61.00ms step:1519/1575 train_time:92694ms step_avg:61.02ms step:1520/1575 train_time:92778ms step_avg:61.04ms step:1521/1575 train_time:92867ms step_avg:61.06ms step:1522/1575 train_time:92952ms step_avg:61.07ms step:1523/1575 train_time:93042ms step_avg:61.09ms step:1524/1575 train_time:93129ms step_avg:61.11ms step:1525/1575 train_time:93222ms step_avg:61.13ms step:1526/1575 train_time:93308ms step_avg:61.15ms step:1527/1575 train_time:93399ms step_avg:61.16ms step:1528/1575 train_time:93484ms step_avg:61.18ms step:1529/1575 train_time:93574ms step_avg:61.20ms step:1530/1575 train_time:93659ms step_avg:61.21ms step:1531/1575 train_time:93748ms step_avg:61.23ms step:1532/1575 train_time:93832ms step_avg:61.25ms step:1533/1575 train_time:93921ms step_avg:61.27ms step:1534/1575 train_time:94007ms step_avg:61.28ms step:1535/1575 train_time:94097ms step_avg:61.30ms step:1536/1575 train_time:94191ms step_avg:61.32ms step:1537/1575 train_time:94280ms step_avg:61.34ms step:1538/1575 train_time:94366ms step_avg:61.36ms step:1539/1575 train_time:94456ms step_avg:61.37ms step:1540/1575 train_time:94542ms step_avg:61.39ms step:1541/1575 train_time:94632ms step_avg:61.41ms step:1542/1575 train_time:94718ms step_avg:61.43ms step:1543/1575 train_time:94807ms step_avg:61.44ms step:1544/1575 train_time:94892ms step_avg:61.46ms step:1545/1575 train_time:94982ms step_avg:61.48ms step:1546/1575 train_time:95067ms step_avg:61.49ms step:1547/1575 train_time:95158ms step_avg:61.51ms step:1548/1575 train_time:95245ms step_avg:61.53ms step:1549/1575 train_time:95336ms step_avg:61.55ms step:1550/1575 train_time:95421ms step_avg:61.56ms step:1551/1575 train_time:95513ms step_avg:61.58ms step:1552/1575 train_time:95598ms step_avg:61.60ms step:1553/1575 train_time:95689ms step_avg:61.62ms step:1554/1575 train_time:95775ms step_avg:61.63ms step:1555/1575 train_time:95863ms step_avg:61.65ms step:1556/1575 train_time:95949ms step_avg:61.66ms step:1557/1575 train_time:96039ms step_avg:61.68ms step:1558/1575 train_time:96125ms step_avg:61.70ms step:1559/1575 train_time:96215ms step_avg:61.72ms step:1560/1575 train_time:96302ms step_avg:61.73ms step:1561/1575 train_time:96392ms step_avg:61.75ms step:1562/1575 train_time:96478ms step_avg:61.77ms step:1563/1575 train_time:96569ms step_avg:61.78ms step:1564/1575 train_time:96655ms step_avg:61.80ms step:1565/1575 train_time:96745ms step_avg:61.82ms step:1566/1575 train_time:96831ms step_avg:61.83ms step:1567/1575 train_time:96920ms step_avg:61.85ms step:1568/1575 train_time:97007ms step_avg:61.87ms step:1569/1575 train_time:97101ms step_avg:61.89ms step:1570/1575 train_time:97185ms step_avg:61.90ms step:1571/1575 train_time:97275ms step_avg:61.92ms step:1572/1575 train_time:97361ms step_avg:61.93ms step:1573/1575 train_time:97451ms step_avg:61.95ms step:1574/1575 train_time:97537ms step_avg:61.97ms step:1575/1575 train_time:97627ms step_avg:61.99ms step:1575/1575 val_loss:3.2775 train_time:97697ms step_avg:62.03ms peak memory allocated: 31016 MiB reserved: 46998 MiB