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:30:12 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 34C P0 121W / 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 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 35C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:6A:00.0 Off | 0 | | N/A 35C P0 121W / 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 126W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 36C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 250711 C /usr/bin/python3 1510MiB | | 1 N/A N/A 250712 C /usr/bin/python3 1510MiB | | 2 N/A N/A 250713 C /usr/bin/python3 1510MiB | | 3 N/A N/A 250714 C /usr/bin/python3 1510MiB | | 4 N/A N/A 250715 C /usr/bin/python3 1510MiB | | 5 N/A N/A 250716 C /usr/bin/python3 1510MiB | | 6 N/A N/A 250717 C /usr/bin/python3 1510MiB | | 7 N/A N/A 250718 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.8328 train_time:0ms step_avg:0.04ms step:1/1575 train_time:88ms step_avg:87.92ms step:2/1575 train_time:111ms step_avg:55.44ms step:3/1575 train_time:132ms step_avg:43.93ms step:4/1575 train_time:166ms step_avg:41.50ms step:5/1575 train_time:197ms step_avg:39.35ms step:6/1575 train_time:296ms step_avg:49.29ms step:7/1575 train_time:314ms step_avg:44.87ms step:8/1575 train_time:337ms step_avg:42.14ms step:9/1575 train_time:368ms step_avg:40.86ms step:10/1575 train_time:406ms step_avg:40.63ms step:11/1575 train_time:437ms step_avg:39.72ms step:12/1575 train_time:476ms step_avg:39.68ms step:13/1575 train_time:507ms step_avg:38.99ms step:14/1575 train_time:546ms step_avg:38.99ms step:15/1575 train_time:576ms step_avg:38.43ms step:16/1575 train_time:615ms step_avg:38.44ms step:17/1575 train_time:646ms step_avg:38.00ms step:18/1575 train_time:685ms step_avg:38.04ms step:19/1575 train_time:716ms step_avg:37.67ms step:20/1575 train_time:755ms step_avg:37.73ms step:21/1575 train_time:786ms step_avg:37.43ms step:22/1575 train_time:825ms step_avg:37.49ms step:23/1575 train_time:855ms step_avg:37.19ms step:24/1575 train_time:894ms step_avg:37.26ms step:25/1575 train_time:925ms step_avg:37.01ms step:26/1575 train_time:964ms step_avg:37.08ms step:27/1575 train_time:995ms step_avg:36.84ms step:28/1575 train_time:1034ms step_avg:36.92ms step:29/1575 train_time:1065ms step_avg:36.71ms step:30/1575 train_time:1103ms step_avg:36.78ms step:31/1575 train_time:1134ms step_avg:36.59ms step:32/1575 train_time:1173ms step_avg:36.66ms step:33/1575 train_time:1204ms step_avg:36.49ms step:34/1575 train_time:1243ms step_avg:36.56ms step:35/1575 train_time:1274ms step_avg:36.40ms step:36/1575 train_time:1313ms step_avg:36.46ms step:37/1575 train_time:1344ms step_avg:36.31ms step:38/1575 train_time:1382ms step_avg:36.37ms step:39/1575 train_time:1413ms step_avg:36.24ms step:40/1575 train_time:1452ms step_avg:36.30ms step:41/1575 train_time:1483ms step_avg:36.17ms step:42/1575 train_time:1522ms step_avg:36.25ms step:43/1575 train_time:1552ms step_avg:36.10ms step:44/1575 train_time:1592ms step_avg:36.17ms step:45/1575 train_time:1622ms step_avg:36.05ms step:46/1575 train_time:1661ms step_avg:36.11ms step:47/1575 train_time:1692ms step_avg:36.00ms step:48/1575 train_time:1731ms step_avg:36.06ms step:49/1575 train_time:1762ms step_avg:35.96ms step:50/1575 train_time:1801ms step_avg:36.02ms step:51/1575 train_time:1832ms step_avg:35.92ms step:52/1575 train_time:1871ms step_avg:35.98ms step:53/1575 train_time:1902ms step_avg:35.88ms step:54/1575 train_time:1941ms step_avg:35.94ms step:55/1575 train_time:1972ms step_avg:35.85ms step:56/1575 train_time:2011ms step_avg:35.90ms step:57/1575 train_time:2041ms step_avg:35.81ms step:58/1575 train_time:2080ms step_avg:35.86ms step:59/1575 train_time:2111ms step_avg:35.78ms step:60/1575 train_time:2150ms step_avg:35.83ms step:61/1575 train_time:2181ms step_avg:35.75ms step:62/1575 train_time:2220ms step_avg:35.80ms step:63/1575 train_time:2251ms step_avg:35.73ms step:64/1575 train_time:2290ms step_avg:35.77ms step:65/1575 train_time:2320ms step_avg:35.70ms step:66/1575 train_time:2359ms step_avg:35.75ms step:67/1575 train_time:2390ms step_avg:35.67ms step:68/1575 train_time:2429ms step_avg:35.71ms step:69/1575 train_time:2459ms step_avg:35.64ms step:70/1575 train_time:2498ms step_avg:35.69ms step:71/1575 train_time:2529ms step_avg:35.62ms step:72/1575 train_time:2568ms step_avg:35.67ms step:73/1575 train_time:2599ms step_avg:35.60ms step:74/1575 train_time:2638ms step_avg:35.64ms step:75/1575 train_time:2669ms step_avg:35.58ms step:76/1575 train_time:2707ms step_avg:35.62ms step:77/1575 train_time:2738ms step_avg:35.56ms step:78/1575 train_time:2777ms step_avg:35.60ms step:79/1575 train_time:2808ms step_avg:35.54ms step:80/1575 train_time:2847ms step_avg:35.58ms step:81/1575 train_time:2878ms step_avg:35.53ms step:82/1575 train_time:2916ms step_avg:35.57ms step:83/1575 train_time:2947ms step_avg:35.51ms step:84/1575 train_time:2986ms step_avg:35.55ms step:85/1575 train_time:3017ms step_avg:35.50ms step:86/1575 train_time:3056ms step_avg:35.53ms step:87/1575 train_time:3086ms step_avg:35.48ms step:88/1575 train_time:3125ms step_avg:35.51ms step:89/1575 train_time:3156ms step_avg:35.46ms step:90/1575 train_time:3194ms step_avg:35.49ms step:91/1575 train_time:3226ms step_avg:35.45ms step:92/1575 train_time:3265ms step_avg:35.48ms step:93/1575 train_time:3295ms step_avg:35.43ms step:94/1575 train_time:3334ms step_avg:35.47ms step:95/1575 train_time:3365ms step_avg:35.42ms step:96/1575 train_time:3404ms step_avg:35.46ms step:97/1575 train_time:3435ms step_avg:35.41ms step:98/1575 train_time:3473ms step_avg:35.44ms step:99/1575 train_time:3504ms step_avg:35.40ms step:100/1575 train_time:3543ms step_avg:35.43ms step:101/1575 train_time:3573ms step_avg:35.38ms step:102/1575 train_time:3612ms step_avg:35.42ms step:103/1575 train_time:3643ms step_avg:35.37ms step:104/1575 train_time:3682ms step_avg:35.41ms step:105/1575 train_time:3713ms step_avg:35.36ms step:106/1575 train_time:3752ms step_avg:35.39ms step:107/1575 train_time:3783ms step_avg:35.35ms step:108/1575 train_time:3821ms step_avg:35.38ms step:109/1575 train_time:3852ms step_avg:35.34ms step:110/1575 train_time:3891ms step_avg:35.37ms step:111/1575 train_time:3922ms step_avg:35.33ms step:112/1575 train_time:3961ms step_avg:35.36ms step:113/1575 train_time:3992ms step_avg:35.32ms step:114/1575 train_time:4031ms step_avg:35.36ms step:115/1575 train_time:4061ms step_avg:35.31ms step:116/1575 train_time:4100ms step_avg:35.34ms step:117/1575 train_time:4131ms step_avg:35.31ms step:118/1575 train_time:4169ms step_avg:35.33ms step:119/1575 train_time:4200ms step_avg:35.30ms step:120/1575 train_time:4239ms step_avg:35.32ms step:121/1575 train_time:4270ms step_avg:35.29ms step:122/1575 train_time:4308ms step_avg:35.31ms step:123/1575 train_time:4339ms step_avg:35.28ms step:124/1575 train_time:4378ms step_avg:35.31ms step:125/1575 train_time:4409ms step_avg:35.27ms step:126/1575 train_time:4447ms step_avg:35.30ms step:127/1575 train_time:4478ms step_avg:35.26ms step:128/1575 train_time:4517ms step_avg:35.29ms step:129/1575 train_time:4548ms step_avg:35.25ms step:130/1575 train_time:4586ms step_avg:35.28ms step:131/1575 train_time:4617ms step_avg:35.24ms step:132/1575 train_time:4656ms step_avg:35.27ms step:133/1575 train_time:4687ms step_avg:35.24ms step:134/1575 train_time:4725ms step_avg:35.26ms step:135/1575 train_time:4756ms step_avg:35.23ms step:136/1575 train_time:4795ms step_avg:35.25ms step:137/1575 train_time:4825ms step_avg:35.22ms step:138/1575 train_time:4864ms step_avg:35.25ms step:139/1575 train_time:4895ms step_avg:35.21ms step:140/1575 train_time:4934ms step_avg:35.24ms step:141/1575 train_time:4964ms step_avg:35.21ms step:142/1575 train_time:5004ms step_avg:35.24ms step:143/1575 train_time:5034ms step_avg:35.21ms step:144/1575 train_time:5073ms step_avg:35.23ms step:145/1575 train_time:5104ms step_avg:35.20ms step:146/1575 train_time:5143ms step_avg:35.22ms step:147/1575 train_time:5173ms step_avg:35.19ms step:148/1575 train_time:5212ms step_avg:35.22ms step:149/1575 train_time:5243ms step_avg:35.19ms step:150/1575 train_time:5281ms step_avg:35.21ms step:151/1575 train_time:5312ms step_avg:35.18ms step:152/1575 train_time:5351ms step_avg:35.21ms step:153/1575 train_time:5382ms step_avg:35.18ms step:154/1575 train_time:5421ms step_avg:35.20ms step:155/1575 train_time:5452ms step_avg:35.17ms step:156/1575 train_time:5491ms step_avg:35.20ms step:157/1575 train_time:5522ms step_avg:35.17ms step:158/1575 train_time:5560ms step_avg:35.19ms step:159/1575 train_time:5591ms step_avg:35.16ms step:160/1575 train_time:5630ms step_avg:35.19ms step:161/1575 train_time:5661ms step_avg:35.16ms step:162/1575 train_time:5700ms step_avg:35.18ms step:163/1575 train_time:5731ms step_avg:35.16ms step:164/1575 train_time:5769ms step_avg:35.18ms step:165/1575 train_time:5800ms step_avg:35.15ms step:166/1575 train_time:5840ms step_avg:35.18ms step:167/1575 train_time:5870ms step_avg:35.15ms step:168/1575 train_time:5908ms step_avg:35.17ms step:169/1575 train_time:5939ms step_avg:35.14ms step:170/1575 train_time:5978ms step_avg:35.16ms step:171/1575 train_time:6009ms step_avg:35.14ms step:172/1575 train_time:6047ms step_avg:35.16ms step:173/1575 train_time:6078ms step_avg:35.13ms step:174/1575 train_time:6117ms step_avg:35.16ms step:175/1575 train_time:6148ms step_avg:35.13ms step:176/1575 train_time:6186ms step_avg:35.15ms step:177/1575 train_time:6217ms step_avg:35.12ms step:178/1575 train_time:6256ms step_avg:35.14ms step:179/1575 train_time:6287ms step_avg:35.12ms step:180/1575 train_time:6325ms step_avg:35.14ms step:181/1575 train_time:6356ms step_avg:35.11ms step:182/1575 train_time:6394ms step_avg:35.13ms step:183/1575 train_time:6425ms step_avg:35.11ms step:184/1575 train_time:6464ms step_avg:35.13ms step:185/1575 train_time:6495ms step_avg:35.11ms step:186/1575 train_time:6534ms step_avg:35.13ms step:187/1575 train_time:6564ms step_avg:35.10ms step:188/1575 train_time:6603ms step_avg:35.12ms step:189/1575 train_time:6633ms step_avg:35.10ms step:190/1575 train_time:6673ms step_avg:35.12ms step:191/1575 train_time:6703ms step_avg:35.10ms step:192/1575 train_time:6742ms step_avg:35.12ms step:193/1575 train_time:6773ms step_avg:35.09ms step:194/1575 train_time:6812ms step_avg:35.11ms step:195/1575 train_time:6844ms step_avg:35.10ms step:196/1575 train_time:6883ms step_avg:35.12ms step:197/1575 train_time:6914ms step_avg:35.09ms step:198/1575 train_time:6952ms step_avg:35.11ms step:199/1575 train_time:6983ms step_avg:35.09ms step:200/1575 train_time:7022ms step_avg:35.11ms step:201/1575 train_time:7053ms step_avg:35.09ms step:202/1575 train_time:7092ms step_avg:35.11ms step:203/1575 train_time:7123ms step_avg:35.09ms step:204/1575 train_time:7161ms step_avg:35.10ms step:205/1575 train_time:7193ms step_avg:35.09ms step:206/1575 train_time:7231ms step_avg:35.10ms step:207/1575 train_time:7262ms step_avg:35.08ms step:208/1575 train_time:7301ms step_avg:35.10ms step:209/1575 train_time:7332ms step_avg:35.08ms step:210/1575 train_time:7371ms step_avg:35.10ms step:211/1575 train_time:7401ms step_avg:35.08ms step:212/1575 train_time:7440ms step_avg:35.09ms step:213/1575 train_time:7471ms step_avg:35.07ms step:214/1575 train_time:7510ms step_avg:35.09ms step:215/1575 train_time:7540ms step_avg:35.07ms step:216/1575 train_time:7579ms step_avg:35.09ms step:217/1575 train_time:7610ms step_avg:35.07ms step:218/1575 train_time:7649ms step_avg:35.09ms step:219/1575 train_time:7680ms step_avg:35.07ms step:220/1575 train_time:7718ms step_avg:35.08ms step:221/1575 train_time:7749ms step_avg:35.06ms step:222/1575 train_time:7788ms step_avg:35.08ms step:223/1575 train_time:7819ms step_avg:35.06ms step:224/1575 train_time:7857ms step_avg:35.08ms step:225/1575 train_time:7888ms step_avg:35.06ms step:226/1575 train_time:7926ms step_avg:35.07ms step:227/1575 train_time:7957ms step_avg:35.05ms step:228/1575 train_time:7996ms step_avg:35.07ms step:229/1575 train_time:8026ms step_avg:35.05ms step:230/1575 train_time:8065ms step_avg:35.07ms step:231/1575 train_time:8095ms step_avg:35.05ms step:232/1575 train_time:8134ms step_avg:35.06ms step:233/1575 train_time:8165ms step_avg:35.04ms step:234/1575 train_time:8204ms step_avg:35.06ms step:235/1575 train_time:8235ms step_avg:35.04ms step:236/1575 train_time:8274ms step_avg:35.06ms step:237/1575 train_time:8305ms step_avg:35.04ms step:238/1575 train_time:8344ms step_avg:35.06ms step:239/1575 train_time:8375ms step_avg:35.04ms step:240/1575 train_time:8413ms step_avg:35.05ms step:241/1575 train_time:8444ms step_avg:35.04ms step:242/1575 train_time:8483ms step_avg:35.05ms step:243/1575 train_time:8514ms step_avg:35.04ms step:244/1575 train_time:8553ms step_avg:35.05ms step:245/1575 train_time:8584ms step_avg:35.04ms step:246/1575 train_time:8622ms step_avg:35.05ms step:247/1575 train_time:8653ms step_avg:35.03ms step:248/1575 train_time:8691ms step_avg:35.05ms step:249/1575 train_time:8722ms step_avg:35.03ms step:250/1575 train_time:8761ms step_avg:35.05ms step:250/1575 val_loss:4.5770 train_time:8810ms step_avg:35.24ms step:251/1575 train_time:8830ms step_avg:35.18ms step:252/1575 train_time:8850ms step_avg:35.12ms step:253/1575 train_time:8867ms step_avg:35.05ms step:254/1575 train_time:8904ms step_avg:35.05ms step:255/1575 train_time:8936ms step_avg:35.04ms step:256/1575 train_time:8976ms step_avg:35.06ms step:257/1575 train_time:9008ms step_avg:35.05ms step:258/1575 train_time:9047ms step_avg:35.06ms step:259/1575 train_time:9078ms step_avg:35.05ms step:260/1575 train_time:9116ms step_avg:35.06ms step:261/1575 train_time:9147ms step_avg:35.05ms step:262/1575 train_time:9186ms step_avg:35.06ms step:263/1575 train_time:9216ms step_avg:35.04ms step:264/1575 train_time:9255ms step_avg:35.06ms step:265/1575 train_time:9286ms step_avg:35.04ms step:266/1575 train_time:9325ms step_avg:35.05ms step:267/1575 train_time:9356ms step_avg:35.04ms step:268/1575 train_time:9395ms step_avg:35.05ms step:269/1575 train_time:9426ms step_avg:35.04ms step:270/1575 train_time:9464ms step_avg:35.05ms step:271/1575 train_time:9495ms step_avg:35.04ms step:272/1575 train_time:9533ms step_avg:35.05ms step:273/1575 train_time:9564ms step_avg:35.03ms step:274/1575 train_time:9603ms step_avg:35.05ms step:275/1575 train_time:9633ms step_avg:35.03ms step:276/1575 train_time:9672ms step_avg:35.04ms step:277/1575 train_time:9703ms step_avg:35.03ms step:278/1575 train_time:9741ms step_avg:35.04ms step:279/1575 train_time:9772ms step_avg:35.03ms step:280/1575 train_time:9811ms step_avg:35.04ms step:281/1575 train_time:9841ms step_avg:35.02ms step:282/1575 train_time:9880ms step_avg:35.04ms step:283/1575 train_time:9911ms step_avg:35.02ms step:284/1575 train_time:9950ms step_avg:35.03ms step:285/1575 train_time:9980ms step_avg:35.02ms step:286/1575 train_time:10019ms step_avg:35.03ms step:287/1575 train_time:10050ms step_avg:35.02ms step:288/1575 train_time:10088ms step_avg:35.03ms step:289/1575 train_time:10119ms step_avg:35.01ms step:290/1575 train_time:10157ms step_avg:35.03ms step:291/1575 train_time:10188ms step_avg:35.01ms step:292/1575 train_time:10227ms step_avg:35.03ms step:293/1575 train_time:10258ms step_avg:35.01ms step:294/1575 train_time:10297ms step_avg:35.02ms step:295/1575 train_time:10328ms step_avg:35.01ms step:296/1575 train_time:10366ms step_avg:35.02ms step:297/1575 train_time:10397ms step_avg:35.01ms step:298/1575 train_time:10436ms step_avg:35.02ms step:299/1575 train_time:10467ms step_avg:35.01ms step:300/1575 train_time:10506ms step_avg:35.02ms step:301/1575 train_time:10537ms step_avg:35.01ms step:302/1575 train_time:10575ms step_avg:35.02ms step:303/1575 train_time:10606ms step_avg:35.00ms step:304/1575 train_time:10645ms step_avg:35.02ms step:305/1575 train_time:10676ms step_avg:35.00ms step:306/1575 train_time:10715ms step_avg:35.02ms step:307/1575 train_time:10745ms step_avg:35.00ms step:308/1575 train_time:10784ms step_avg:35.01ms step:309/1575 train_time:10814ms step_avg:35.00ms step:310/1575 train_time:10853ms step_avg:35.01ms step:311/1575 train_time:10884ms step_avg:35.00ms step:312/1575 train_time:10924ms step_avg:35.01ms step:313/1575 train_time:10954ms step_avg:35.00ms step:314/1575 train_time:10994ms step_avg:35.01ms step:315/1575 train_time:11024ms step_avg:35.00ms step:316/1575 train_time:11062ms step_avg:35.01ms step:317/1575 train_time:11093ms step_avg:34.99ms step:318/1575 train_time:11132ms step_avg:35.01ms step:319/1575 train_time:11163ms step_avg:34.99ms step:320/1575 train_time:11201ms step_avg:35.00ms step:321/1575 train_time:11232ms step_avg:34.99ms step:322/1575 train_time:11270ms step_avg:35.00ms step:323/1575 train_time:11301ms step_avg:34.99ms step:324/1575 train_time:11339ms step_avg:35.00ms step:325/1575 train_time:11370ms step_avg:34.98ms step:326/1575 train_time:11408ms step_avg:35.00ms step:327/1575 train_time:11439ms step_avg:34.98ms step:328/1575 train_time:11478ms step_avg:34.99ms step:329/1575 train_time:11509ms step_avg:34.98ms step:330/1575 train_time:11547ms step_avg:34.99ms step:331/1575 train_time:11577ms step_avg:34.98ms step:332/1575 train_time:11616ms step_avg:34.99ms step:333/1575 train_time:11647ms step_avg:34.98ms step:334/1575 train_time:11686ms step_avg:34.99ms step:335/1575 train_time:11717ms step_avg:34.97ms step:336/1575 train_time:11755ms step_avg:34.99ms step:337/1575 train_time:11786ms step_avg:34.97ms step:338/1575 train_time:11825ms step_avg:34.98ms step:339/1575 train_time:11855ms step_avg:34.97ms step:340/1575 train_time:11894ms step_avg:34.98ms step:341/1575 train_time:11925ms step_avg:34.97ms step:342/1575 train_time:11963ms step_avg:34.98ms step:343/1575 train_time:11994ms step_avg:34.97ms step:344/1575 train_time:12033ms step_avg:34.98ms step:345/1575 train_time:12064ms step_avg:34.97ms step:346/1575 train_time:12102ms step_avg:34.98ms step:347/1575 train_time:12133ms step_avg:34.96ms step:348/1575 train_time:12172ms step_avg:34.98ms step:349/1575 train_time:12202ms step_avg:34.96ms step:350/1575 train_time:12241ms step_avg:34.97ms step:351/1575 train_time:12272ms step_avg:34.96ms step:352/1575 train_time:12310ms step_avg:34.97ms step:353/1575 train_time:12341ms step_avg:34.96ms step:354/1575 train_time:12379ms step_avg:34.97ms step:355/1575 train_time:12410ms step_avg:34.96ms step:356/1575 train_time:12449ms step_avg:34.97ms step:357/1575 train_time:12480ms step_avg:34.96ms step:358/1575 train_time:12518ms step_avg:34.97ms step:359/1575 train_time:12549ms step_avg:34.95ms step:360/1575 train_time:12588ms step_avg:34.97ms step:361/1575 train_time:12618ms step_avg:34.95ms step:362/1575 train_time:12657ms step_avg:34.96ms step:363/1575 train_time:12687ms step_avg:34.95ms step:364/1575 train_time:12726ms step_avg:34.96ms step:365/1575 train_time:12756ms step_avg:34.95ms step:366/1575 train_time:12795ms step_avg:34.96ms step:367/1575 train_time:12826ms step_avg:34.95ms step:368/1575 train_time:12864ms step_avg:34.96ms step:369/1575 train_time:12895ms step_avg:34.95ms step:370/1575 train_time:12934ms step_avg:34.96ms step:371/1575 train_time:12965ms step_avg:34.95ms step:372/1575 train_time:13004ms step_avg:34.96ms step:373/1575 train_time:13035ms step_avg:34.95ms step:374/1575 train_time:13073ms step_avg:34.96ms step:375/1575 train_time:13104ms step_avg:34.94ms step:376/1575 train_time:13143ms step_avg:34.95ms step:377/1575 train_time:13173ms step_avg:34.94ms step:378/1575 train_time:13212ms step_avg:34.95ms step:379/1575 train_time:13242ms step_avg:34.94ms step:380/1575 train_time:13281ms step_avg:34.95ms step:381/1575 train_time:13312ms step_avg:34.94ms step:382/1575 train_time:13350ms step_avg:34.95ms step:383/1575 train_time:13381ms step_avg:34.94ms step:384/1575 train_time:13420ms step_avg:34.95ms step:385/1575 train_time:13451ms step_avg:34.94ms step:386/1575 train_time:13489ms step_avg:34.95ms step:387/1575 train_time:13520ms step_avg:34.94ms step:388/1575 train_time:13558ms step_avg:34.94ms step:389/1575 train_time:13589ms step_avg:34.93ms step:390/1575 train_time:13628ms step_avg:34.94ms step:391/1575 train_time:13659ms step_avg:34.93ms step:392/1575 train_time:13697ms step_avg:34.94ms step:393/1575 train_time:13728ms step_avg:34.93ms step:394/1575 train_time:13766ms step_avg:34.94ms step:395/1575 train_time:13797ms step_avg:34.93ms step:396/1575 train_time:13836ms step_avg:34.94ms step:397/1575 train_time:13867ms step_avg:34.93ms step:398/1575 train_time:13906ms step_avg:34.94ms step:399/1575 train_time:13936ms step_avg:34.93ms step:400/1575 train_time:13975ms step_avg:34.94ms step:401/1575 train_time:14006ms step_avg:34.93ms step:402/1575 train_time:14044ms step_avg:34.94ms step:403/1575 train_time:14075ms step_avg:34.93ms step:404/1575 train_time:14114ms step_avg:34.94ms step:405/1575 train_time:14145ms step_avg:34.93ms step:406/1575 train_time:14184ms step_avg:34.93ms step:407/1575 train_time:14215ms step_avg:34.93ms step:408/1575 train_time:14253ms step_avg:34.93ms step:409/1575 train_time:14284ms step_avg:34.92ms step:410/1575 train_time:14322ms step_avg:34.93ms step:411/1575 train_time:14353ms step_avg:34.92ms step:412/1575 train_time:14391ms step_avg:34.93ms step:413/1575 train_time:14422ms step_avg:34.92ms step:414/1575 train_time:14460ms step_avg:34.93ms step:415/1575 train_time:14491ms step_avg:34.92ms step:416/1575 train_time:14530ms step_avg:34.93ms step:417/1575 train_time:14561ms step_avg:34.92ms step:418/1575 train_time:14599ms step_avg:34.93ms step:419/1575 train_time:14630ms step_avg:34.92ms step:420/1575 train_time:14668ms step_avg:34.92ms step:421/1575 train_time:14699ms step_avg:34.91ms step:422/1575 train_time:14738ms step_avg:34.92ms step:423/1575 train_time:14768ms step_avg:34.91ms step:424/1575 train_time:14808ms step_avg:34.92ms step:425/1575 train_time:14838ms step_avg:34.91ms step:426/1575 train_time:14877ms step_avg:34.92ms step:427/1575 train_time:14908ms step_avg:34.91ms step:428/1575 train_time:14946ms step_avg:34.92ms step:429/1575 train_time:14977ms step_avg:34.91ms step:430/1575 train_time:15016ms step_avg:34.92ms step:431/1575 train_time:15047ms step_avg:34.91ms step:432/1575 train_time:15085ms step_avg:34.92ms step:433/1575 train_time:15116ms step_avg:34.91ms step:434/1575 train_time:15154ms step_avg:34.92ms step:435/1575 train_time:15185ms step_avg:34.91ms step:436/1575 train_time:15224ms step_avg:34.92ms step:437/1575 train_time:15255ms step_avg:34.91ms step:438/1575 train_time:15293ms step_avg:34.92ms step:439/1575 train_time:15324ms step_avg:34.91ms step:440/1575 train_time:15363ms step_avg:34.92ms step:441/1575 train_time:15394ms step_avg:34.91ms step:442/1575 train_time:15432ms step_avg:34.91ms step:443/1575 train_time:15463ms step_avg:34.91ms step:444/1575 train_time:15502ms step_avg:34.91ms step:445/1575 train_time:15532ms step_avg:34.90ms step:446/1575 train_time:15571ms step_avg:34.91ms step:447/1575 train_time:15601ms step_avg:34.90ms step:448/1575 train_time:15640ms step_avg:34.91ms step:449/1575 train_time:15671ms step_avg:34.90ms step:450/1575 train_time:15710ms step_avg:34.91ms step:451/1575 train_time:15740ms step_avg:34.90ms step:452/1575 train_time:15779ms step_avg:34.91ms step:453/1575 train_time:15810ms step_avg:34.90ms step:454/1575 train_time:15848ms step_avg:34.91ms step:455/1575 train_time:15879ms step_avg:34.90ms step:456/1575 train_time:15917ms step_avg:34.91ms step:457/1575 train_time:15948ms step_avg:34.90ms step:458/1575 train_time:15987ms step_avg:34.91ms step:459/1575 train_time:16017ms step_avg:34.90ms step:460/1575 train_time:16056ms step_avg:34.91ms step:461/1575 train_time:16087ms step_avg:34.89ms step:462/1575 train_time:16126ms step_avg:34.90ms step:463/1575 train_time:16156ms step_avg:34.89ms step:464/1575 train_time:16194ms step_avg:34.90ms step:465/1575 train_time:16225ms step_avg:34.89ms step:466/1575 train_time:16264ms step_avg:34.90ms step:467/1575 train_time:16295ms step_avg:34.89ms step:468/1575 train_time:16334ms step_avg:34.90ms step:469/1575 train_time:16364ms step_avg:34.89ms step:470/1575 train_time:16403ms step_avg:34.90ms step:471/1575 train_time:16434ms step_avg:34.89ms step:472/1575 train_time:16473ms step_avg:34.90ms step:473/1575 train_time:16503ms step_avg:34.89ms step:474/1575 train_time:16542ms step_avg:34.90ms step:475/1575 train_time:16573ms step_avg:34.89ms step:476/1575 train_time:16611ms step_avg:34.90ms step:477/1575 train_time:16642ms step_avg:34.89ms step:478/1575 train_time:16680ms step_avg:34.90ms step:479/1575 train_time:16711ms step_avg:34.89ms step:480/1575 train_time:16749ms step_avg:34.89ms step:481/1575 train_time:16780ms step_avg:34.89ms step:482/1575 train_time:16819ms step_avg:34.89ms step:483/1575 train_time:16850ms step_avg:34.89ms step:484/1575 train_time:16888ms step_avg:34.89ms step:485/1575 train_time:16919ms step_avg:34.88ms step:486/1575 train_time:16957ms step_avg:34.89ms step:487/1575 train_time:16989ms step_avg:34.88ms step:488/1575 train_time:17028ms step_avg:34.89ms step:489/1575 train_time:17058ms step_avg:34.88ms step:490/1575 train_time:17097ms step_avg:34.89ms step:491/1575 train_time:17128ms step_avg:34.88ms step:492/1575 train_time:17167ms step_avg:34.89ms step:493/1575 train_time:17197ms step_avg:34.88ms step:494/1575 train_time:17236ms step_avg:34.89ms step:495/1575 train_time:17266ms step_avg:34.88ms step:496/1575 train_time:17305ms step_avg:34.89ms step:497/1575 train_time:17335ms step_avg:34.88ms step:498/1575 train_time:17374ms step_avg:34.89ms step:499/1575 train_time:17405ms step_avg:34.88ms step:500/1575 train_time:17444ms step_avg:34.89ms step:500/1575 val_loss:4.2260 train_time:17492ms step_avg:34.98ms step:501/1575 train_time:17512ms step_avg:34.95ms step:502/1575 train_time:17532ms step_avg:34.92ms step:503/1575 train_time:17550ms step_avg:34.89ms step:504/1575 train_time:17587ms step_avg:34.89ms step:505/1575 train_time:17619ms step_avg:34.89ms step:506/1575 train_time:17659ms step_avg:34.90ms step:507/1575 train_time:17690ms step_avg:34.89ms step:508/1575 train_time:17729ms step_avg:34.90ms step:509/1575 train_time:17760ms step_avg:34.89ms step:510/1575 train_time:17798ms step_avg:34.90ms step:511/1575 train_time:17829ms step_avg:34.89ms step:512/1575 train_time:17867ms step_avg:34.90ms step:513/1575 train_time:17940ms step_avg:34.97ms step:514/1575 train_time:17996ms step_avg:35.01ms step:515/1575 train_time:18058ms step_avg:35.06ms step:516/1575 train_time:18117ms step_avg:35.11ms step:517/1575 train_time:18180ms step_avg:35.16ms step:518/1575 train_time:18238ms step_avg:35.21ms step:519/1575 train_time:18301ms step_avg:35.26ms step:520/1575 train_time:18360ms step_avg:35.31ms step:521/1575 train_time:18422ms step_avg:35.36ms step:522/1575 train_time:18483ms step_avg:35.41ms step:523/1575 train_time:18548ms step_avg:35.46ms step:524/1575 train_time:18609ms step_avg:35.51ms step:525/1575 train_time:18673ms step_avg:35.57ms step:526/1575 train_time:18733ms step_avg:35.61ms step:527/1575 train_time:18799ms step_avg:35.67ms step:528/1575 train_time:18859ms step_avg:35.72ms step:529/1575 train_time:18922ms step_avg:35.77ms step:530/1575 train_time:18981ms step_avg:35.81ms step:531/1575 train_time:19044ms step_avg:35.86ms step:532/1575 train_time:19103ms step_avg:35.91ms step:533/1575 train_time:19166ms step_avg:35.96ms step:534/1575 train_time:19225ms step_avg:36.00ms step:535/1575 train_time:19289ms step_avg:36.05ms step:536/1575 train_time:19349ms step_avg:36.10ms step:537/1575 train_time:19411ms step_avg:36.15ms step:538/1575 train_time:19471ms step_avg:36.19ms step:539/1575 train_time:19535ms step_avg:36.24ms step:540/1575 train_time:19594ms step_avg:36.29ms step:541/1575 train_time:19659ms step_avg:36.34ms step:542/1575 train_time:19719ms step_avg:36.38ms step:543/1575 train_time:19783ms step_avg:36.43ms step:544/1575 train_time:19843ms step_avg:36.48ms step:545/1575 train_time:19906ms step_avg:36.52ms step:546/1575 train_time:19966ms step_avg:36.57ms step:547/1575 train_time:20029ms step_avg:36.62ms step:548/1575 train_time:20088ms step_avg:36.66ms step:549/1575 train_time:20152ms step_avg:36.71ms step:550/1575 train_time:20211ms step_avg:36.75ms step:551/1575 train_time:20274ms step_avg:36.80ms step:552/1575 train_time:20335ms step_avg:36.84ms step:553/1575 train_time:20399ms step_avg:36.89ms step:554/1575 train_time:20456ms step_avg:36.92ms step:555/1575 train_time:20520ms step_avg:36.97ms step:556/1575 train_time:20579ms step_avg:37.01ms step:557/1575 train_time:20643ms step_avg:37.06ms step:558/1575 train_time:20702ms step_avg:37.10ms step:559/1575 train_time:20766ms step_avg:37.15ms step:560/1575 train_time:20826ms step_avg:37.19ms step:561/1575 train_time:20889ms step_avg:37.24ms step:562/1575 train_time:20949ms step_avg:37.28ms step:563/1575 train_time:21013ms step_avg:37.32ms step:564/1575 train_time:21072ms step_avg:37.36ms step:565/1575 train_time:21136ms step_avg:37.41ms step:566/1575 train_time:21196ms step_avg:37.45ms step:567/1575 train_time:21259ms step_avg:37.49ms step:568/1575 train_time:21318ms step_avg:37.53ms step:569/1575 train_time:21385ms step_avg:37.58ms step:570/1575 train_time:21442ms step_avg:37.62ms step:571/1575 train_time:21506ms step_avg:37.66ms step:572/1575 train_time:21565ms step_avg:37.70ms step:573/1575 train_time:21628ms step_avg:37.74ms step:574/1575 train_time:21686ms step_avg:37.78ms step:575/1575 train_time:21750ms step_avg:37.83ms step:576/1575 train_time:21810ms step_avg:37.86ms step:577/1575 train_time:21874ms step_avg:37.91ms step:578/1575 train_time:21932ms step_avg:37.95ms step:579/1575 train_time:21996ms step_avg:37.99ms step:580/1575 train_time:22055ms step_avg:38.03ms step:581/1575 train_time:22118ms step_avg:38.07ms step:582/1575 train_time:22177ms step_avg:38.11ms step:583/1575 train_time:22241ms step_avg:38.15ms step:584/1575 train_time:22300ms step_avg:38.19ms step:585/1575 train_time:22364ms step_avg:38.23ms step:586/1575 train_time:22423ms step_avg:38.26ms step:587/1575 train_time:22487ms step_avg:38.31ms step:588/1575 train_time:22546ms step_avg:38.34ms step:589/1575 train_time:22609ms step_avg:38.39ms step:590/1575 train_time:22668ms step_avg:38.42ms step:591/1575 train_time:22731ms step_avg:38.46ms step:592/1575 train_time:22791ms step_avg:38.50ms step:593/1575 train_time:22856ms step_avg:38.54ms step:594/1575 train_time:22914ms step_avg:38.58ms step:595/1575 train_time:22977ms step_avg:38.62ms step:596/1575 train_time:23036ms step_avg:38.65ms step:597/1575 train_time:23099ms step_avg:38.69ms step:598/1575 train_time:23158ms step_avg:38.73ms step:599/1575 train_time:23222ms step_avg:38.77ms step:600/1575 train_time:23282ms step_avg:38.80ms step:601/1575 train_time:23345ms step_avg:38.84ms step:602/1575 train_time:23404ms step_avg:38.88ms step:603/1575 train_time:23467ms step_avg:38.92ms step:604/1575 train_time:23527ms step_avg:38.95ms step:605/1575 train_time:23590ms step_avg:38.99ms step:606/1575 train_time:23649ms step_avg:39.02ms step:607/1575 train_time:23713ms step_avg:39.07ms step:608/1575 train_time:23772ms step_avg:39.10ms step:609/1575 train_time:23835ms step_avg:39.14ms step:610/1575 train_time:23895ms step_avg:39.17ms step:611/1575 train_time:23959ms step_avg:39.21ms step:612/1575 train_time:24017ms step_avg:39.24ms step:613/1575 train_time:24081ms step_avg:39.28ms step:614/1575 train_time:24140ms step_avg:39.32ms step:615/1575 train_time:24204ms step_avg:39.36ms step:616/1575 train_time:24263ms step_avg:39.39ms step:617/1575 train_time:24326ms step_avg:39.43ms step:618/1575 train_time:24385ms step_avg:39.46ms step:619/1575 train_time:24448ms step_avg:39.50ms step:620/1575 train_time:24507ms step_avg:39.53ms step:621/1575 train_time:24571ms step_avg:39.57ms step:622/1575 train_time:24630ms step_avg:39.60ms step:623/1575 train_time:24693ms step_avg:39.64ms step:624/1575 train_time:24752ms step_avg:39.67ms step:625/1575 train_time:24816ms step_avg:39.71ms step:626/1575 train_time:24876ms step_avg:39.74ms step:627/1575 train_time:24939ms step_avg:39.78ms step:628/1575 train_time:24998ms step_avg:39.81ms step:629/1575 train_time:25062ms step_avg:39.84ms step:630/1575 train_time:25121ms step_avg:39.87ms step:631/1575 train_time:25184ms step_avg:39.91ms step:632/1575 train_time:25243ms step_avg:39.94ms step:633/1575 train_time:25306ms step_avg:39.98ms step:634/1575 train_time:25367ms step_avg:40.01ms step:635/1575 train_time:25431ms step_avg:40.05ms step:636/1575 train_time:25490ms step_avg:40.08ms step:637/1575 train_time:25552ms step_avg:40.11ms step:638/1575 train_time:25611ms step_avg:40.14ms step:639/1575 train_time:25674ms step_avg:40.18ms step:640/1575 train_time:25734ms step_avg:40.21ms step:641/1575 train_time:25798ms step_avg:40.25ms step:642/1575 train_time:25858ms step_avg:40.28ms step:643/1575 train_time:25921ms step_avg:40.31ms step:644/1575 train_time:25980ms step_avg:40.34ms step:645/1575 train_time:26043ms step_avg:40.38ms step:646/1575 train_time:26102ms step_avg:40.41ms step:647/1575 train_time:26165ms step_avg:40.44ms step:648/1575 train_time:26225ms step_avg:40.47ms step:649/1575 train_time:26288ms step_avg:40.51ms step:650/1575 train_time:26347ms step_avg:40.53ms step:651/1575 train_time:26410ms step_avg:40.57ms step:652/1575 train_time:26470ms step_avg:40.60ms step:653/1575 train_time:26533ms step_avg:40.63ms step:654/1575 train_time:26592ms step_avg:40.66ms step:655/1575 train_time:26657ms step_avg:40.70ms step:656/1575 train_time:26716ms step_avg:40.73ms step:657/1575 train_time:26779ms step_avg:40.76ms step:658/1575 train_time:26838ms step_avg:40.79ms step:659/1575 train_time:26902ms step_avg:40.82ms step:660/1575 train_time:26961ms step_avg:40.85ms step:661/1575 train_time:27024ms step_avg:40.88ms step:662/1575 train_time:27084ms step_avg:40.91ms step:663/1575 train_time:27147ms step_avg:40.95ms step:664/1575 train_time:27207ms step_avg:40.97ms step:665/1575 train_time:27270ms step_avg:41.01ms step:666/1575 train_time:27330ms step_avg:41.04ms step:667/1575 train_time:27394ms step_avg:41.07ms step:668/1575 train_time:27453ms step_avg:41.10ms step:669/1575 train_time:27517ms step_avg:41.13ms step:670/1575 train_time:27576ms step_avg:41.16ms step:671/1575 train_time:27639ms step_avg:41.19ms step:672/1575 train_time:27698ms step_avg:41.22ms step:673/1575 train_time:27761ms step_avg:41.25ms step:674/1575 train_time:27821ms step_avg:41.28ms step:675/1575 train_time:27885ms step_avg:41.31ms step:676/1575 train_time:27944ms step_avg:41.34ms step:677/1575 train_time:28007ms step_avg:41.37ms step:678/1575 train_time:28067ms step_avg:41.40ms step:679/1575 train_time:28130ms step_avg:41.43ms step:680/1575 train_time:28189ms step_avg:41.45ms step:681/1575 train_time:28253ms step_avg:41.49ms step:682/1575 train_time:28312ms step_avg:41.51ms step:683/1575 train_time:28376ms step_avg:41.55ms step:684/1575 train_time:28435ms step_avg:41.57ms step:685/1575 train_time:28499ms step_avg:41.60ms step:686/1575 train_time:28558ms step_avg:41.63ms step:687/1575 train_time:28621ms step_avg:41.66ms step:688/1575 train_time:28680ms step_avg:41.69ms step:689/1575 train_time:28743ms step_avg:41.72ms step:690/1575 train_time:28803ms step_avg:41.74ms step:691/1575 train_time:28867ms step_avg:41.78ms step:692/1575 train_time:28926ms step_avg:41.80ms step:693/1575 train_time:28990ms step_avg:41.83ms step:694/1575 train_time:29048ms step_avg:41.86ms step:695/1575 train_time:29111ms step_avg:41.89ms step:696/1575 train_time:29170ms step_avg:41.91ms step:697/1575 train_time:29233ms step_avg:41.94ms step:698/1575 train_time:29293ms step_avg:41.97ms step:699/1575 train_time:29357ms step_avg:42.00ms step:700/1575 train_time:29416ms step_avg:42.02ms step:701/1575 train_time:29479ms step_avg:42.05ms step:702/1575 train_time:29538ms step_avg:42.08ms step:703/1575 train_time:29602ms step_avg:42.11ms step:704/1575 train_time:29661ms step_avg:42.13ms step:705/1575 train_time:29725ms step_avg:42.16ms step:706/1575 train_time:29784ms step_avg:42.19ms step:707/1575 train_time:29848ms step_avg:42.22ms step:708/1575 train_time:29907ms step_avg:42.24ms step:709/1575 train_time:29970ms step_avg:42.27ms step:710/1575 train_time:30029ms step_avg:42.29ms step:711/1575 train_time:30094ms step_avg:42.33ms step:712/1575 train_time:30153ms step_avg:42.35ms step:713/1575 train_time:30217ms step_avg:42.38ms step:714/1575 train_time:30276ms step_avg:42.40ms step:715/1575 train_time:30339ms step_avg:42.43ms step:716/1575 train_time:30399ms step_avg:42.46ms step:717/1575 train_time:30462ms step_avg:42.49ms step:718/1575 train_time:30522ms step_avg:42.51ms step:719/1575 train_time:30588ms step_avg:42.54ms step:720/1575 train_time:30646ms step_avg:42.56ms step:721/1575 train_time:30708ms step_avg:42.59ms step:722/1575 train_time:30768ms step_avg:42.61ms step:723/1575 train_time:30832ms step_avg:42.64ms step:724/1575 train_time:30891ms step_avg:42.67ms step:725/1575 train_time:30954ms step_avg:42.70ms step:726/1575 train_time:31014ms step_avg:42.72ms step:727/1575 train_time:31078ms step_avg:42.75ms step:728/1575 train_time:31137ms step_avg:42.77ms step:729/1575 train_time:31201ms step_avg:42.80ms step:730/1575 train_time:31260ms step_avg:42.82ms step:731/1575 train_time:31322ms step_avg:42.85ms step:732/1575 train_time:31381ms step_avg:42.87ms step:733/1575 train_time:31445ms step_avg:42.90ms step:734/1575 train_time:31505ms step_avg:42.92ms step:735/1575 train_time:31567ms step_avg:42.95ms step:736/1575 train_time:31627ms step_avg:42.97ms step:737/1575 train_time:31690ms step_avg:43.00ms step:738/1575 train_time:31749ms step_avg:43.02ms step:739/1575 train_time:31812ms step_avg:43.05ms step:740/1575 train_time:31872ms step_avg:43.07ms step:741/1575 train_time:31934ms step_avg:43.10ms step:742/1575 train_time:31994ms step_avg:43.12ms step:743/1575 train_time:32057ms step_avg:43.15ms step:744/1575 train_time:32117ms step_avg:43.17ms step:745/1575 train_time:32180ms step_avg:43.19ms step:746/1575 train_time:32241ms step_avg:43.22ms step:747/1575 train_time:32303ms step_avg:43.24ms step:748/1575 train_time:32364ms step_avg:43.27ms step:749/1575 train_time:32427ms step_avg:43.29ms step:750/1575 train_time:32485ms step_avg:43.31ms step:750/1575 val_loss:3.8823 train_time:32531ms step_avg:43.37ms step:751/1575 train_time:32552ms step_avg:43.34ms step:752/1575 train_time:32610ms step_avg:43.37ms step:753/1575 train_time:32678ms step_avg:43.40ms step:754/1575 train_time:32741ms step_avg:43.42ms step:755/1575 train_time:32804ms step_avg:43.45ms step:756/1575 train_time:32864ms step_avg:43.47ms step:757/1575 train_time:32928ms step_avg:43.50ms step:758/1575 train_time:32987ms step_avg:43.52ms step:759/1575 train_time:33049ms step_avg:43.54ms step:760/1575 train_time:33108ms step_avg:43.56ms step:761/1575 train_time:33179ms step_avg:43.60ms step:762/1575 train_time:33240ms step_avg:43.62ms step:763/1575 train_time:33303ms step_avg:43.65ms step:764/1575 train_time:33360ms step_avg:43.66ms step:765/1575 train_time:33423ms step_avg:43.69ms step:766/1575 train_time:33481ms step_avg:43.71ms step:767/1575 train_time:33543ms step_avg:43.73ms step:768/1575 train_time:33604ms step_avg:43.76ms step:769/1575 train_time:33670ms step_avg:43.78ms step:770/1575 train_time:33730ms step_avg:43.81ms step:771/1575 train_time:33796ms step_avg:43.83ms step:772/1575 train_time:33847ms step_avg:43.84ms step:773/1575 train_time:33910ms step_avg:43.87ms step:774/1575 train_time:33969ms step_avg:43.89ms step:775/1575 train_time:34032ms step_avg:43.91ms step:776/1575 train_time:34091ms step_avg:43.93ms step:777/1575 train_time:34155ms step_avg:43.96ms step:778/1575 train_time:34214ms step_avg:43.98ms step:779/1575 train_time:34276ms step_avg:44.00ms step:780/1575 train_time:34335ms step_avg:44.02ms step:781/1575 train_time:34398ms step_avg:44.04ms step:782/1575 train_time:34457ms step_avg:44.06ms step:783/1575 train_time:34520ms step_avg:44.09ms step:784/1575 train_time:34580ms step_avg:44.11ms step:785/1575 train_time:34644ms step_avg:44.13ms step:786/1575 train_time:34704ms step_avg:44.15ms step:787/1575 train_time:34767ms step_avg:44.18ms step:788/1575 train_time:34828ms step_avg:44.20ms step:789/1575 train_time:34892ms step_avg:44.22ms step:790/1575 train_time:34950ms step_avg:44.24ms step:791/1575 train_time:35013ms step_avg:44.26ms step:792/1575 train_time:35072ms step_avg:44.28ms step:793/1575 train_time:35136ms step_avg:44.31ms step:794/1575 train_time:35195ms step_avg:44.33ms step:795/1575 train_time:35257ms step_avg:44.35ms step:796/1575 train_time:35317ms step_avg:44.37ms step:797/1575 train_time:35380ms step_avg:44.39ms step:798/1575 train_time:35440ms step_avg:44.41ms step:799/1575 train_time:35503ms step_avg:44.43ms step:800/1575 train_time:35561ms step_avg:44.45ms step:801/1575 train_time:35625ms step_avg:44.48ms step:802/1575 train_time:35685ms step_avg:44.50ms step:803/1575 train_time:35749ms step_avg:44.52ms step:804/1575 train_time:35808ms step_avg:44.54ms step:805/1575 train_time:35872ms step_avg:44.56ms step:806/1575 train_time:35934ms step_avg:44.58ms step:807/1575 train_time:35995ms step_avg:44.60ms step:808/1575 train_time:36055ms step_avg:44.62ms step:809/1575 train_time:36117ms step_avg:44.64ms step:810/1575 train_time:36176ms step_avg:44.66ms step:811/1575 train_time:36239ms step_avg:44.68ms step:812/1575 train_time:36298ms step_avg:44.70ms step:813/1575 train_time:36361ms step_avg:44.72ms step:814/1575 train_time:36420ms step_avg:44.74ms step:815/1575 train_time:36484ms step_avg:44.77ms step:816/1575 train_time:36544ms step_avg:44.78ms step:817/1575 train_time:36606ms step_avg:44.81ms step:818/1575 train_time:36666ms step_avg:44.82ms step:819/1575 train_time:36729ms step_avg:44.85ms step:820/1575 train_time:36789ms step_avg:44.86ms step:821/1575 train_time:36853ms step_avg:44.89ms step:822/1575 train_time:36912ms step_avg:44.91ms step:823/1575 train_time:36975ms step_avg:44.93ms step:824/1575 train_time:37035ms step_avg:44.95ms step:825/1575 train_time:37098ms step_avg:44.97ms step:826/1575 train_time:37158ms step_avg:44.99ms step:827/1575 train_time:37221ms step_avg:45.01ms step:828/1575 train_time:37281ms step_avg:45.03ms step:829/1575 train_time:37343ms step_avg:45.05ms step:830/1575 train_time:37403ms step_avg:45.06ms step:831/1575 train_time:37466ms step_avg:45.09ms step:832/1575 train_time:37525ms step_avg:45.10ms step:833/1575 train_time:37589ms step_avg:45.12ms step:834/1575 train_time:37648ms step_avg:45.14ms step:835/1575 train_time:37712ms step_avg:45.16ms step:836/1575 train_time:37771ms step_avg:45.18ms step:837/1575 train_time:37835ms step_avg:45.20ms step:838/1575 train_time:37895ms step_avg:45.22ms step:839/1575 train_time:37958ms step_avg:45.24ms step:840/1575 train_time:38017ms step_avg:45.26ms step:841/1575 train_time:38080ms step_avg:45.28ms step:842/1575 train_time:38140ms step_avg:45.30ms step:843/1575 train_time:38203ms step_avg:45.32ms step:844/1575 train_time:38262ms step_avg:45.33ms step:845/1575 train_time:38325ms step_avg:45.36ms step:846/1575 train_time:38384ms step_avg:45.37ms step:847/1575 train_time:38448ms step_avg:45.39ms step:848/1575 train_time:38508ms step_avg:45.41ms step:849/1575 train_time:38571ms step_avg:45.43ms step:850/1575 train_time:38630ms step_avg:45.45ms step:851/1575 train_time:38693ms step_avg:45.47ms step:852/1575 train_time:38753ms step_avg:45.48ms step:853/1575 train_time:38816ms step_avg:45.50ms step:854/1575 train_time:38875ms step_avg:45.52ms step:855/1575 train_time:38938ms step_avg:45.54ms step:856/1575 train_time:38997ms step_avg:45.56ms step:857/1575 train_time:39061ms step_avg:45.58ms step:858/1575 train_time:39121ms step_avg:45.60ms step:859/1575 train_time:39184ms step_avg:45.62ms step:860/1575 train_time:39243ms step_avg:45.63ms step:861/1575 train_time:39307ms step_avg:45.65ms step:862/1575 train_time:39367ms step_avg:45.67ms step:863/1575 train_time:39430ms step_avg:45.69ms step:864/1575 train_time:39489ms step_avg:45.70ms step:865/1575 train_time:39552ms step_avg:45.73ms step:866/1575 train_time:39611ms step_avg:45.74ms step:867/1575 train_time:39675ms step_avg:45.76ms step:868/1575 train_time:39734ms step_avg:45.78ms step:869/1575 train_time:39797ms step_avg:45.80ms step:870/1575 train_time:39857ms step_avg:45.81ms step:871/1575 train_time:39921ms step_avg:45.83ms step:872/1575 train_time:39980ms step_avg:45.85ms step:873/1575 train_time:40043ms step_avg:45.87ms step:874/1575 train_time:40101ms step_avg:45.88ms step:875/1575 train_time:40165ms step_avg:45.90ms step:876/1575 train_time:40224ms step_avg:45.92ms step:877/1575 train_time:40288ms step_avg:45.94ms step:878/1575 train_time:40347ms step_avg:45.95ms step:879/1575 train_time:40410ms step_avg:45.97ms step:880/1575 train_time:40469ms step_avg:45.99ms step:881/1575 train_time:40533ms step_avg:46.01ms step:882/1575 train_time:40592ms step_avg:46.02ms step:883/1575 train_time:40655ms step_avg:46.04ms step:884/1575 train_time:40715ms step_avg:46.06ms step:885/1575 train_time:40780ms step_avg:46.08ms step:886/1575 train_time:40845ms step_avg:46.10ms step:887/1575 train_time:40904ms step_avg:46.11ms step:888/1575 train_time:40964ms step_avg:46.13ms step:889/1575 train_time:41026ms step_avg:46.15ms step:890/1575 train_time:41085ms step_avg:46.16ms step:891/1575 train_time:41147ms step_avg:46.18ms step:892/1575 train_time:41206ms step_avg:46.20ms step:893/1575 train_time:41271ms step_avg:46.22ms step:894/1575 train_time:41330ms step_avg:46.23ms step:895/1575 train_time:41393ms step_avg:46.25ms step:896/1575 train_time:41452ms step_avg:46.26ms step:897/1575 train_time:41516ms step_avg:46.28ms step:898/1575 train_time:41575ms step_avg:46.30ms step:899/1575 train_time:41638ms step_avg:46.32ms step:900/1575 train_time:41697ms step_avg:46.33ms step:901/1575 train_time:41760ms step_avg:46.35ms step:902/1575 train_time:41819ms step_avg:46.36ms step:903/1575 train_time:41883ms step_avg:46.38ms step:904/1575 train_time:41942ms step_avg:46.40ms step:905/1575 train_time:42007ms step_avg:46.42ms step:906/1575 train_time:42065ms step_avg:46.43ms step:907/1575 train_time:42129ms step_avg:46.45ms step:908/1575 train_time:42189ms step_avg:46.46ms step:909/1575 train_time:42252ms step_avg:46.48ms step:910/1575 train_time:42311ms step_avg:46.50ms step:911/1575 train_time:42375ms step_avg:46.51ms step:912/1575 train_time:42435ms step_avg:46.53ms step:913/1575 train_time:42498ms step_avg:46.55ms step:914/1575 train_time:42558ms step_avg:46.56ms step:915/1575 train_time:42621ms step_avg:46.58ms step:916/1575 train_time:42681ms step_avg:46.59ms step:917/1575 train_time:42744ms step_avg:46.61ms step:918/1575 train_time:42803ms step_avg:46.63ms step:919/1575 train_time:42868ms step_avg:46.65ms step:920/1575 train_time:42927ms step_avg:46.66ms step:921/1575 train_time:42990ms step_avg:46.68ms step:922/1575 train_time:43050ms step_avg:46.69ms step:923/1575 train_time:43113ms step_avg:46.71ms step:924/1575 train_time:43173ms step_avg:46.72ms step:925/1575 train_time:43236ms step_avg:46.74ms step:926/1575 train_time:43296ms step_avg:46.76ms step:927/1575 train_time:43359ms step_avg:46.77ms step:928/1575 train_time:43418ms step_avg:46.79ms step:929/1575 train_time:43482ms step_avg:46.81ms step:930/1575 train_time:43542ms step_avg:46.82ms step:931/1575 train_time:43606ms step_avg:46.84ms step:932/1575 train_time:43665ms step_avg:46.85ms step:933/1575 train_time:43730ms step_avg:46.87ms step:934/1575 train_time:43789ms step_avg:46.88ms step:935/1575 train_time:43853ms step_avg:46.90ms step:936/1575 train_time:43912ms step_avg:46.91ms step:937/1575 train_time:43976ms step_avg:46.93ms step:938/1575 train_time:44035ms step_avg:46.95ms step:939/1575 train_time:44099ms step_avg:46.96ms step:940/1575 train_time:44158ms step_avg:46.98ms step:941/1575 train_time:44222ms step_avg:46.99ms step:942/1575 train_time:44281ms step_avg:47.01ms step:943/1575 train_time:44345ms step_avg:47.03ms step:944/1575 train_time:44404ms step_avg:47.04ms step:945/1575 train_time:44468ms step_avg:47.06ms step:946/1575 train_time:44528ms step_avg:47.07ms step:947/1575 train_time:44591ms step_avg:47.09ms step:948/1575 train_time:44650ms step_avg:47.10ms step:949/1575 train_time:44714ms step_avg:47.12ms step:950/1575 train_time:44774ms step_avg:47.13ms step:951/1575 train_time:44837ms step_avg:47.15ms step:952/1575 train_time:44897ms step_avg:47.16ms step:953/1575 train_time:44961ms step_avg:47.18ms step:954/1575 train_time:45021ms step_avg:47.19ms step:955/1575 train_time:45085ms step_avg:47.21ms step:956/1575 train_time:45143ms step_avg:47.22ms step:957/1575 train_time:45207ms step_avg:47.24ms step:958/1575 train_time:45266ms step_avg:47.25ms step:959/1575 train_time:45330ms step_avg:47.27ms step:960/1575 train_time:45389ms step_avg:47.28ms step:961/1575 train_time:45452ms step_avg:47.30ms step:962/1575 train_time:45512ms step_avg:47.31ms step:963/1575 train_time:45575ms step_avg:47.33ms step:964/1575 train_time:45635ms step_avg:47.34ms step:965/1575 train_time:45698ms step_avg:47.36ms step:966/1575 train_time:45757ms step_avg:47.37ms step:967/1575 train_time:45819ms step_avg:47.38ms step:968/1575 train_time:45879ms step_avg:47.40ms step:969/1575 train_time:45942ms step_avg:47.41ms step:970/1575 train_time:46001ms step_avg:47.42ms step:971/1575 train_time:46065ms step_avg:47.44ms step:972/1575 train_time:46124ms step_avg:47.45ms step:973/1575 train_time:46188ms step_avg:47.47ms step:974/1575 train_time:46247ms step_avg:47.48ms step:975/1575 train_time:46311ms step_avg:47.50ms step:976/1575 train_time:46370ms step_avg:47.51ms step:977/1575 train_time:46434ms step_avg:47.53ms step:978/1575 train_time:46493ms step_avg:47.54ms step:979/1575 train_time:46556ms step_avg:47.55ms step:980/1575 train_time:46615ms step_avg:47.57ms step:981/1575 train_time:46679ms step_avg:47.58ms step:982/1575 train_time:46739ms step_avg:47.60ms step:983/1575 train_time:46802ms step_avg:47.61ms step:984/1575 train_time:46861ms step_avg:47.62ms step:985/1575 train_time:46924ms step_avg:47.64ms step:986/1575 train_time:46983ms step_avg:47.65ms step:987/1575 train_time:47047ms step_avg:47.67ms step:988/1575 train_time:47106ms step_avg:47.68ms step:989/1575 train_time:47171ms step_avg:47.70ms step:990/1575 train_time:47230ms step_avg:47.71ms step:991/1575 train_time:47293ms step_avg:47.72ms step:992/1575 train_time:47353ms step_avg:47.73ms step:993/1575 train_time:47417ms step_avg:47.75ms step:994/1575 train_time:47476ms step_avg:47.76ms step:995/1575 train_time:47539ms step_avg:47.78ms step:996/1575 train_time:47598ms step_avg:47.79ms step:997/1575 train_time:47661ms step_avg:47.80ms step:998/1575 train_time:47720ms step_avg:47.82ms step:999/1575 train_time:47784ms step_avg:47.83ms step:1000/1575 train_time:47844ms step_avg:47.84ms step:1000/1575 val_loss:3.5812 train_time:47889ms step_avg:47.89ms step:1001/1575 train_time:47910ms step_avg:47.86ms step:1002/1575 train_time:47969ms step_avg:47.87ms step:1003/1575 train_time:48035ms step_avg:47.89ms step:1004/1575 train_time:48097ms step_avg:47.91ms step:1005/1575 train_time:48159ms step_avg:47.92ms step:1006/1575 train_time:48218ms step_avg:47.93ms step:1007/1575 train_time:48281ms step_avg:47.95ms step:1008/1575 train_time:48341ms step_avg:47.96ms step:1009/1575 train_time:48404ms step_avg:47.97ms step:1010/1575 train_time:48462ms step_avg:47.98ms step:1011/1575 train_time:48525ms step_avg:48.00ms step:1012/1575 train_time:48584ms step_avg:48.01ms step:1013/1575 train_time:48648ms step_avg:48.02ms step:1014/1575 train_time:48705ms step_avg:48.03ms step:1015/1575 train_time:48768ms step_avg:48.05ms step:1016/1575 train_time:48827ms step_avg:48.06ms step:1017/1575 train_time:48891ms step_avg:48.07ms step:1018/1575 train_time:48952ms step_avg:48.09ms step:1019/1575 train_time:49017ms step_avg:48.10ms step:1020/1575 train_time:49077ms step_avg:48.11ms step:1021/1575 train_time:49142ms step_avg:48.13ms step:1022/1575 train_time:49200ms step_avg:48.14ms step:1023/1575 train_time:49265ms step_avg:48.16ms step:1024/1575 train_time:49324ms step_avg:48.17ms step:1025/1575 train_time:49396ms step_avg:48.19ms step:1026/1575 train_time:49477ms step_avg:48.22ms step:1027/1575 train_time:49567ms step_avg:48.26ms step:1028/1575 train_time:49653ms step_avg:48.30ms step:1029/1575 train_time:49741ms step_avg:48.34ms step:1030/1575 train_time:49826ms step_avg:48.38ms step:1031/1575 train_time:49917ms step_avg:48.42ms step:1032/1575 train_time:50003ms step_avg:48.45ms step:1033/1575 train_time:50096ms step_avg:48.50ms step:1034/1575 train_time:50182ms step_avg:48.53ms step:1035/1575 train_time:50271ms step_avg:48.57ms step:1036/1575 train_time:50357ms step_avg:48.61ms step:1037/1575 train_time:50447ms step_avg:48.65ms step:1038/1575 train_time:50531ms step_avg:48.68ms step:1039/1575 train_time:50620ms step_avg:48.72ms step:1040/1575 train_time:50705ms step_avg:48.75ms step:1041/1575 train_time:50794ms step_avg:48.79ms step:1042/1575 train_time:50880ms step_avg:48.83ms step:1043/1575 train_time:50971ms step_avg:48.87ms step:1044/1575 train_time:51057ms step_avg:48.90ms step:1045/1575 train_time:51148ms step_avg:48.95ms step:1046/1575 train_time:51233ms step_avg:48.98ms step:1047/1575 train_time:51322ms step_avg:49.02ms step:1048/1575 train_time:51408ms step_avg:49.05ms step:1049/1575 train_time:51498ms step_avg:49.09ms step:1050/1575 train_time:51584ms step_avg:49.13ms step:1051/1575 train_time:51672ms step_avg:49.16ms step:1052/1575 train_time:51757ms step_avg:49.20ms step:1053/1575 train_time:51846ms step_avg:49.24ms step:1054/1575 train_time:51933ms step_avg:49.27ms step:1055/1575 train_time:52023ms step_avg:49.31ms step:1056/1575 train_time:52111ms step_avg:49.35ms step:1057/1575 train_time:52200ms step_avg:49.39ms step:1058/1575 train_time:52286ms step_avg:49.42ms step:1059/1575 train_time:52376ms step_avg:49.46ms step:1060/1575 train_time:52463ms step_avg:49.49ms step:1061/1575 train_time:52552ms step_avg:49.53ms step:1062/1575 train_time:52637ms step_avg:49.56ms step:1063/1575 train_time:52725ms step_avg:49.60ms step:1064/1575 train_time:52811ms step_avg:49.63ms step:1065/1575 train_time:52900ms step_avg:49.67ms step:1066/1575 train_time:52987ms step_avg:49.71ms step:1067/1575 train_time:53076ms step_avg:49.74ms step:1068/1575 train_time:53163ms step_avg:49.78ms step:1069/1575 train_time:53252ms step_avg:49.82ms step:1070/1575 train_time:53339ms step_avg:49.85ms step:1071/1575 train_time:53429ms step_avg:49.89ms step:1072/1575 train_time:53514ms step_avg:49.92ms step:1073/1575 train_time:53603ms step_avg:49.96ms step:1074/1575 train_time:53689ms step_avg:49.99ms step:1075/1575 train_time:53778ms step_avg:50.03ms step:1076/1575 train_time:53864ms step_avg:50.06ms step:1077/1575 train_time:53954ms step_avg:50.10ms step:1078/1575 train_time:54040ms step_avg:50.13ms step:1079/1575 train_time:54131ms step_avg:50.17ms step:1080/1575 train_time:54216ms step_avg:50.20ms step:1081/1575 train_time:54306ms step_avg:50.24ms step:1082/1575 train_time:54393ms step_avg:50.27ms step:1083/1575 train_time:54482ms step_avg:50.31ms step:1084/1575 train_time:54568ms step_avg:50.34ms step:1085/1575 train_time:54657ms step_avg:50.38ms step:1086/1575 train_time:54743ms step_avg:50.41ms step:1087/1575 train_time:54832ms step_avg:50.44ms step:1088/1575 train_time:54920ms step_avg:50.48ms step:1089/1575 train_time:55008ms step_avg:50.51ms step:1090/1575 train_time:55094ms step_avg:50.55ms step:1091/1575 train_time:55185ms step_avg:50.58ms step:1092/1575 train_time:55270ms step_avg:50.61ms step:1093/1575 train_time:55359ms step_avg:50.65ms step:1094/1575 train_time:55445ms step_avg:50.68ms step:1095/1575 train_time:55536ms step_avg:50.72ms step:1096/1575 train_time:55621ms step_avg:50.75ms step:1097/1575 train_time:55710ms step_avg:50.78ms step:1098/1575 train_time:55795ms step_avg:50.82ms step:1099/1575 train_time:55885ms step_avg:50.85ms step:1100/1575 train_time:55971ms step_avg:50.88ms step:1101/1575 train_time:56061ms step_avg:50.92ms step:1102/1575 train_time:56147ms step_avg:50.95ms step:1103/1575 train_time:56236ms step_avg:50.98ms step:1104/1575 train_time:56322ms step_avg:51.02ms step:1105/1575 train_time:56412ms step_avg:51.05ms step:1106/1575 train_time:56497ms step_avg:51.08ms step:1107/1575 train_time:56588ms step_avg:51.12ms step:1108/1575 train_time:56673ms step_avg:51.15ms step:1109/1575 train_time:56762ms step_avg:51.18ms step:1110/1575 train_time:56849ms step_avg:51.21ms step:1111/1575 train_time:56938ms step_avg:51.25ms step:1112/1575 train_time:57029ms step_avg:51.28ms step:1113/1575 train_time:57119ms step_avg:51.32ms step:1114/1575 train_time:57201ms step_avg:51.35ms step:1115/1575 train_time:57290ms step_avg:51.38ms step:1116/1575 train_time:57375ms step_avg:51.41ms step:1117/1575 train_time:57465ms step_avg:51.45ms step:1118/1575 train_time:57550ms step_avg:51.48ms step:1119/1575 train_time:57639ms step_avg:51.51ms step:1120/1575 train_time:57726ms step_avg:51.54ms step:1121/1575 train_time:57814ms step_avg:51.57ms step:1122/1575 train_time:57900ms step_avg:51.60ms step:1123/1575 train_time:57990ms step_avg:51.64ms step:1124/1575 train_time:58075ms step_avg:51.67ms step:1125/1575 train_time:58165ms step_avg:51.70ms step:1126/1575 train_time:58251ms step_avg:51.73ms step:1127/1575 train_time:58340ms step_avg:51.77ms step:1128/1575 train_time:58426ms step_avg:51.80ms step:1129/1575 train_time:58515ms step_avg:51.83ms step:1130/1575 train_time:58601ms step_avg:51.86ms step:1131/1575 train_time:58691ms step_avg:51.89ms step:1132/1575 train_time:58776ms step_avg:51.92ms step:1133/1575 train_time:58865ms step_avg:51.96ms step:1134/1575 train_time:58950ms step_avg:51.98ms step:1135/1575 train_time:59040ms step_avg:52.02ms step:1136/1575 train_time:59127ms step_avg:52.05ms step:1137/1575 train_time:59217ms step_avg:52.08ms step:1138/1575 train_time:59302ms step_avg:52.11ms step:1139/1575 train_time:59392ms step_avg:52.14ms step:1140/1575 train_time:59477ms step_avg:52.17ms step:1141/1575 train_time:59568ms step_avg:52.21ms step:1142/1575 train_time:59653ms step_avg:52.24ms step:1143/1575 train_time:59743ms step_avg:52.27ms step:1144/1575 train_time:59828ms step_avg:52.30ms step:1145/1575 train_time:59918ms step_avg:52.33ms step:1146/1575 train_time:60008ms step_avg:52.36ms step:1147/1575 train_time:60096ms step_avg:52.39ms step:1148/1575 train_time:60183ms step_avg:52.42ms step:1149/1575 train_time:60272ms step_avg:52.46ms step:1150/1575 train_time:60358ms step_avg:52.48ms step:1151/1575 train_time:60447ms step_avg:52.52ms step:1152/1575 train_time:60533ms step_avg:52.55ms step:1153/1575 train_time:60622ms step_avg:52.58ms step:1154/1575 train_time:60708ms step_avg:52.61ms step:1155/1575 train_time:60797ms step_avg:52.64ms step:1156/1575 train_time:60884ms step_avg:52.67ms step:1157/1575 train_time:60973ms step_avg:52.70ms step:1158/1575 train_time:61059ms step_avg:52.73ms step:1159/1575 train_time:61150ms step_avg:52.76ms step:1160/1575 train_time:61236ms step_avg:52.79ms step:1161/1575 train_time:61326ms step_avg:52.82ms step:1162/1575 train_time:61412ms step_avg:52.85ms step:1163/1575 train_time:61501ms step_avg:52.88ms step:1164/1575 train_time:61587ms step_avg:52.91ms step:1165/1575 train_time:61676ms step_avg:52.94ms step:1166/1575 train_time:61763ms step_avg:52.97ms step:1167/1575 train_time:61851ms step_avg:53.00ms step:1168/1575 train_time:61937ms step_avg:53.03ms step:1169/1575 train_time:62027ms step_avg:53.06ms step:1170/1575 train_time:62113ms step_avg:53.09ms step:1171/1575 train_time:62203ms step_avg:53.12ms step:1172/1575 train_time:62288ms step_avg:53.15ms step:1173/1575 train_time:62378ms step_avg:53.18ms step:1174/1575 train_time:62464ms step_avg:53.21ms step:1175/1575 train_time:62553ms step_avg:53.24ms step:1176/1575 train_time:62639ms step_avg:53.26ms step:1177/1575 train_time:62729ms step_avg:53.30ms step:1178/1575 train_time:62814ms step_avg:53.32ms step:1179/1575 train_time:62904ms step_avg:53.35ms step:1180/1575 train_time:62990ms step_avg:53.38ms step:1181/1575 train_time:63079ms step_avg:53.41ms step:1182/1575 train_time:63166ms step_avg:53.44ms step:1183/1575 train_time:63253ms step_avg:53.47ms step:1184/1575 train_time:63340ms step_avg:53.50ms step:1185/1575 train_time:63430ms step_avg:53.53ms step:1186/1575 train_time:63515ms step_avg:53.55ms step:1187/1575 train_time:63604ms step_avg:53.58ms step:1188/1575 train_time:63691ms step_avg:53.61ms step:1189/1575 train_time:63780ms step_avg:53.64ms step:1190/1575 train_time:63865ms step_avg:53.67ms step:1191/1575 train_time:63955ms step_avg:53.70ms step:1192/1575 train_time:64041ms step_avg:53.73ms step:1193/1575 train_time:64131ms step_avg:53.76ms step:1194/1575 train_time:64216ms step_avg:53.78ms step:1195/1575 train_time:64307ms step_avg:53.81ms step:1196/1575 train_time:64393ms step_avg:53.84ms step:1197/1575 train_time:64483ms step_avg:53.87ms step:1198/1575 train_time:64568ms step_avg:53.90ms step:1199/1575 train_time:64660ms step_avg:53.93ms step:1200/1575 train_time:64745ms step_avg:53.95ms step:1201/1575 train_time:64833ms step_avg:53.98ms step:1202/1575 train_time:64919ms step_avg:54.01ms step:1203/1575 train_time:65009ms step_avg:54.04ms step:1204/1575 train_time:65096ms step_avg:54.07ms step:1205/1575 train_time:65185ms step_avg:54.10ms step:1206/1575 train_time:65270ms step_avg:54.12ms step:1207/1575 train_time:65360ms step_avg:54.15ms step:1208/1575 train_time:65446ms step_avg:54.18ms step:1209/1575 train_time:65535ms step_avg:54.21ms step:1210/1575 train_time:65622ms step_avg:54.23ms step:1211/1575 train_time:65712ms step_avg:54.26ms step:1212/1575 train_time:65797ms step_avg:54.29ms step:1213/1575 train_time:65887ms step_avg:54.32ms step:1214/1575 train_time:65972ms step_avg:54.34ms step:1215/1575 train_time:66062ms step_avg:54.37ms step:1216/1575 train_time:66148ms step_avg:54.40ms step:1217/1575 train_time:66238ms step_avg:54.43ms step:1218/1575 train_time:66325ms step_avg:54.45ms step:1219/1575 train_time:66413ms step_avg:54.48ms step:1220/1575 train_time:66499ms step_avg:54.51ms step:1221/1575 train_time:66588ms step_avg:54.54ms step:1222/1575 train_time:66675ms step_avg:54.56ms step:1223/1575 train_time:66764ms step_avg:54.59ms step:1224/1575 train_time:66850ms step_avg:54.62ms step:1225/1575 train_time:66939ms step_avg:54.64ms step:1226/1575 train_time:67024ms step_avg:54.67ms step:1227/1575 train_time:67114ms step_avg:54.70ms step:1228/1575 train_time:67200ms step_avg:54.72ms step:1229/1575 train_time:67291ms step_avg:54.75ms step:1230/1575 train_time:67376ms step_avg:54.78ms step:1231/1575 train_time:67465ms step_avg:54.81ms step:1232/1575 train_time:67551ms step_avg:54.83ms step:1233/1575 train_time:67641ms step_avg:54.86ms step:1234/1575 train_time:67727ms step_avg:54.88ms step:1235/1575 train_time:67816ms step_avg:54.91ms step:1236/1575 train_time:67902ms step_avg:54.94ms step:1237/1575 train_time:67991ms step_avg:54.96ms step:1238/1575 train_time:68077ms step_avg:54.99ms step:1239/1575 train_time:68168ms step_avg:55.02ms step:1240/1575 train_time:68252ms step_avg:55.04ms step:1241/1575 train_time:68342ms step_avg:55.07ms step:1242/1575 train_time:68427ms step_avg:55.09ms step:1243/1575 train_time:68516ms step_avg:55.12ms step:1244/1575 train_time:68603ms step_avg:55.15ms step:1245/1575 train_time:68693ms step_avg:55.17ms step:1246/1575 train_time:68779ms step_avg:55.20ms step:1247/1575 train_time:68868ms step_avg:55.23ms step:1248/1575 train_time:68954ms step_avg:55.25ms step:1249/1575 train_time:69043ms step_avg:55.28ms step:1250/1575 train_time:69129ms step_avg:55.30ms step:1250/1575 val_loss:3.4043 train_time:69202ms step_avg:55.36ms step:1251/1575 train_time:69223ms step_avg:55.33ms step:1252/1575 train_time:69310ms step_avg:55.36ms step:1253/1575 train_time:69402ms step_avg:55.39ms step:1254/1575 train_time:69489ms step_avg:55.41ms step:1255/1575 train_time:69578ms step_avg:55.44ms step:1256/1575 train_time:69662ms step_avg:55.46ms step:1257/1575 train_time:69749ms step_avg:55.49ms step:1258/1575 train_time:69835ms step_avg:55.51ms step:1259/1575 train_time:69924ms step_avg:55.54ms step:1260/1575 train_time:70009ms step_avg:55.56ms step:1261/1575 train_time:70098ms step_avg:55.59ms step:1262/1575 train_time:70184ms step_avg:55.61ms step:1263/1575 train_time:70275ms step_avg:55.64ms step:1264/1575 train_time:70363ms step_avg:55.67ms step:1265/1575 train_time:70453ms step_avg:55.69ms step:1266/1575 train_time:70540ms step_avg:55.72ms step:1267/1575 train_time:70628ms step_avg:55.74ms step:1268/1575 train_time:70713ms step_avg:55.77ms step:1269/1575 train_time:70803ms step_avg:55.79ms step:1270/1575 train_time:70888ms step_avg:55.82ms step:1271/1575 train_time:70976ms step_avg:55.84ms step:1272/1575 train_time:71062ms step_avg:55.87ms step:1273/1575 train_time:71150ms step_avg:55.89ms step:1274/1575 train_time:71238ms step_avg:55.92ms step:1275/1575 train_time:71328ms step_avg:55.94ms step:1276/1575 train_time:71415ms step_avg:55.97ms step:1277/1575 train_time:71506ms step_avg:56.00ms step:1278/1575 train_time:71592ms step_avg:56.02ms step:1279/1575 train_time:71681ms step_avg:56.04ms step:1280/1575 train_time:71766ms step_avg:56.07ms step:1281/1575 train_time:71855ms step_avg:56.09ms step:1282/1575 train_time:71940ms step_avg:56.12ms step:1283/1575 train_time:72029ms step_avg:56.14ms step:1284/1575 train_time:72114ms step_avg:56.16ms step:1285/1575 train_time:72204ms step_avg:56.19ms step:1286/1575 train_time:72292ms step_avg:56.21ms step:1287/1575 train_time:72381ms step_avg:56.24ms step:1288/1575 train_time:72467ms step_avg:56.26ms step:1289/1575 train_time:72557ms step_avg:56.29ms step:1290/1575 train_time:72643ms step_avg:56.31ms step:1291/1575 train_time:72732ms step_avg:56.34ms step:1292/1575 train_time:72818ms step_avg:56.36ms step:1293/1575 train_time:72907ms step_avg:56.39ms step:1294/1575 train_time:72992ms step_avg:56.41ms step:1295/1575 train_time:73083ms step_avg:56.43ms step:1296/1575 train_time:73168ms step_avg:56.46ms step:1297/1575 train_time:73257ms step_avg:56.48ms step:1298/1575 train_time:73344ms step_avg:56.51ms step:1299/1575 train_time:73434ms step_avg:56.53ms step:1300/1575 train_time:73520ms step_avg:56.55ms step:1301/1575 train_time:73609ms step_avg:56.58ms step:1302/1575 train_time:73696ms step_avg:56.60ms step:1303/1575 train_time:73787ms step_avg:56.63ms step:1304/1575 train_time:73871ms step_avg:56.65ms step:1305/1575 train_time:73961ms step_avg:56.68ms step:1306/1575 train_time:74048ms step_avg:56.70ms step:1307/1575 train_time:74136ms step_avg:56.72ms step:1308/1575 train_time:74222ms step_avg:56.74ms step:1309/1575 train_time:74312ms step_avg:56.77ms step:1310/1575 train_time:74399ms step_avg:56.79ms step:1311/1575 train_time:74488ms step_avg:56.82ms step:1312/1575 train_time:74574ms step_avg:56.84ms step:1313/1575 train_time:74664ms step_avg:56.86ms step:1314/1575 train_time:74750ms step_avg:56.89ms step:1315/1575 train_time:74840ms step_avg:56.91ms step:1316/1575 train_time:74925ms step_avg:56.93ms step:1317/1575 train_time:75014ms step_avg:56.96ms step:1318/1575 train_time:75100ms step_avg:56.98ms step:1319/1575 train_time:75190ms step_avg:57.01ms step:1320/1575 train_time:75276ms step_avg:57.03ms step:1321/1575 train_time:75367ms step_avg:57.05ms step:1322/1575 train_time:75453ms step_avg:57.07ms step:1323/1575 train_time:75543ms step_avg:57.10ms step:1324/1575 train_time:75629ms step_avg:57.12ms step:1325/1575 train_time:75719ms step_avg:57.15ms step:1326/1575 train_time:75804ms step_avg:57.17ms step:1327/1575 train_time:75894ms step_avg:57.19ms step:1328/1575 train_time:75980ms step_avg:57.21ms step:1329/1575 train_time:76069ms step_avg:57.24ms step:1330/1575 train_time:76155ms step_avg:57.26ms step:1331/1575 train_time:76244ms step_avg:57.28ms step:1332/1575 train_time:76330ms step_avg:57.30ms step:1333/1575 train_time:76420ms step_avg:57.33ms step:1334/1575 train_time:76505ms step_avg:57.35ms step:1335/1575 train_time:76595ms step_avg:57.37ms step:1336/1575 train_time:76681ms step_avg:57.40ms step:1337/1575 train_time:76770ms step_avg:57.42ms step:1338/1575 train_time:76855ms step_avg:57.44ms step:1339/1575 train_time:76945ms step_avg:57.46ms step:1340/1575 train_time:77031ms step_avg:57.49ms step:1341/1575 train_time:77122ms step_avg:57.51ms step:1342/1575 train_time:77207ms step_avg:57.53ms step:1343/1575 train_time:77296ms step_avg:57.56ms step:1344/1575 train_time:77381ms step_avg:57.58ms step:1345/1575 train_time:77471ms step_avg:57.60ms step:1346/1575 train_time:77559ms step_avg:57.62ms step:1347/1575 train_time:77648ms step_avg:57.64ms step:1348/1575 train_time:77734ms step_avg:57.67ms step:1349/1575 train_time:77823ms step_avg:57.69ms step:1350/1575 train_time:77909ms step_avg:57.71ms step:1351/1575 train_time:77999ms step_avg:57.73ms step:1352/1575 train_time:78084ms step_avg:57.75ms step:1353/1575 train_time:78174ms step_avg:57.78ms step:1354/1575 train_time:78260ms step_avg:57.80ms step:1355/1575 train_time:78349ms step_avg:57.82ms step:1356/1575 train_time:78435ms step_avg:57.84ms step:1357/1575 train_time:78525ms step_avg:57.87ms step:1358/1575 train_time:78615ms step_avg:57.89ms step:1359/1575 train_time:78703ms step_avg:57.91ms step:1360/1575 train_time:78789ms step_avg:57.93ms step:1361/1575 train_time:78877ms step_avg:57.96ms step:1362/1575 train_time:78963ms step_avg:57.98ms step:1363/1575 train_time:79052ms step_avg:58.00ms step:1364/1575 train_time:79141ms step_avg:58.02ms step:1365/1575 train_time:79228ms step_avg:58.04ms step:1366/1575 train_time:79313ms step_avg:58.06ms step:1367/1575 train_time:79403ms step_avg:58.09ms step:1368/1575 train_time:79489ms step_avg:58.11ms step:1369/1575 train_time:79579ms step_avg:58.13ms step:1370/1575 train_time:79664ms step_avg:58.15ms step:1371/1575 train_time:79754ms step_avg:58.17ms step:1372/1575 train_time:79840ms step_avg:58.19ms step:1373/1575 train_time:79929ms step_avg:58.21ms step:1374/1575 train_time:80014ms step_avg:58.23ms step:1375/1575 train_time:80103ms step_avg:58.26ms step:1376/1575 train_time:80189ms step_avg:58.28ms step:1377/1575 train_time:80279ms step_avg:58.30ms step:1378/1575 train_time:80364ms step_avg:58.32ms step:1379/1575 train_time:80453ms step_avg:58.34ms step:1380/1575 train_time:80540ms step_avg:58.36ms step:1381/1575 train_time:80629ms step_avg:58.38ms step:1382/1575 train_time:80715ms step_avg:58.40ms step:1383/1575 train_time:80805ms step_avg:58.43ms step:1384/1575 train_time:80891ms step_avg:58.45ms step:1385/1575 train_time:80980ms step_avg:58.47ms step:1386/1575 train_time:81065ms step_avg:58.49ms step:1387/1575 train_time:81156ms step_avg:58.51ms step:1388/1575 train_time:81242ms step_avg:58.53ms step:1389/1575 train_time:81331ms step_avg:58.55ms step:1390/1575 train_time:81417ms step_avg:58.57ms step:1391/1575 train_time:81507ms step_avg:58.60ms step:1392/1575 train_time:81595ms step_avg:58.62ms step:1393/1575 train_time:81684ms step_avg:58.64ms step:1394/1575 train_time:81770ms step_avg:58.66ms step:1395/1575 train_time:81860ms step_avg:58.68ms step:1396/1575 train_time:81945ms step_avg:58.70ms step:1397/1575 train_time:82034ms step_avg:58.72ms step:1398/1575 train_time:82120ms step_avg:58.74ms step:1399/1575 train_time:82209ms step_avg:58.76ms step:1400/1575 train_time:82296ms step_avg:58.78ms step:1401/1575 train_time:82385ms step_avg:58.80ms step:1402/1575 train_time:82471ms step_avg:58.82ms step:1403/1575 train_time:82561ms step_avg:58.85ms step:1404/1575 train_time:82647ms step_avg:58.87ms step:1405/1575 train_time:82737ms step_avg:58.89ms step:1406/1575 train_time:82822ms step_avg:58.91ms step:1407/1575 train_time:82912ms step_avg:58.93ms step:1408/1575 train_time:82998ms step_avg:58.95ms step:1409/1575 train_time:83088ms step_avg:58.97ms step:1410/1575 train_time:83175ms step_avg:58.99ms step:1411/1575 train_time:83263ms step_avg:59.01ms step:1412/1575 train_time:83349ms step_avg:59.03ms step:1413/1575 train_time:83438ms step_avg:59.05ms step:1414/1575 train_time:83524ms step_avg:59.07ms step:1415/1575 train_time:83614ms step_avg:59.09ms step:1416/1575 train_time:83700ms step_avg:59.11ms step:1417/1575 train_time:83789ms step_avg:59.13ms step:1418/1575 train_time:83875ms step_avg:59.15ms step:1419/1575 train_time:83965ms step_avg:59.17ms step:1420/1575 train_time:84053ms step_avg:59.19ms step:1421/1575 train_time:84140ms step_avg:59.21ms step:1422/1575 train_time:84228ms step_avg:59.23ms step:1423/1575 train_time:84316ms step_avg:59.25ms step:1424/1575 train_time:84401ms step_avg:59.27ms step:1425/1575 train_time:84491ms step_avg:59.29ms step:1426/1575 train_time:84577ms step_avg:59.31ms step:1427/1575 train_time:84667ms step_avg:59.33ms step:1428/1575 train_time:84753ms step_avg:59.35ms step:1429/1575 train_time:84843ms step_avg:59.37ms step:1430/1575 train_time:84929ms step_avg:59.39ms step:1431/1575 train_time:85019ms step_avg:59.41ms step:1432/1575 train_time:85104ms step_avg:59.43ms step:1433/1575 train_time:85193ms step_avg:59.45ms step:1434/1575 train_time:85279ms step_avg:59.47ms step:1435/1575 train_time:85368ms step_avg:59.49ms step:1436/1575 train_time:85455ms step_avg:59.51ms step:1437/1575 train_time:85546ms step_avg:59.53ms step:1438/1575 train_time:85632ms step_avg:59.55ms step:1439/1575 train_time:85723ms step_avg:59.57ms step:1440/1575 train_time:85809ms step_avg:59.59ms step:1441/1575 train_time:85898ms step_avg:59.61ms step:1442/1575 train_time:85983ms step_avg:59.63ms step:1443/1575 train_time:86073ms step_avg:59.65ms step:1444/1575 train_time:86159ms step_avg:59.67ms step:1445/1575 train_time:86249ms step_avg:59.69ms step:1446/1575 train_time:86334ms step_avg:59.71ms step:1447/1575 train_time:86424ms step_avg:59.73ms step:1448/1575 train_time:86510ms step_avg:59.74ms step:1449/1575 train_time:86600ms step_avg:59.77ms step:1450/1575 train_time:86686ms step_avg:59.78ms step:1451/1575 train_time:86775ms step_avg:59.80ms step:1452/1575 train_time:86860ms step_avg:59.82ms step:1453/1575 train_time:86949ms step_avg:59.84ms step:1454/1575 train_time:87036ms step_avg:59.86ms step:1455/1575 train_time:87126ms step_avg:59.88ms step:1456/1575 train_time:87212ms step_avg:59.90ms step:1457/1575 train_time:87303ms step_avg:59.92ms step:1458/1575 train_time:87387ms step_avg:59.94ms step:1459/1575 train_time:87477ms step_avg:59.96ms step:1460/1575 train_time:87563ms step_avg:59.97ms step:1461/1575 train_time:87652ms step_avg:59.99ms step:1462/1575 train_time:87738ms step_avg:60.01ms step:1463/1575 train_time:87826ms step_avg:60.03ms step:1464/1575 train_time:87913ms step_avg:60.05ms step:1465/1575 train_time:88003ms step_avg:60.07ms step:1466/1575 train_time:88088ms step_avg:60.09ms step:1467/1575 train_time:88178ms step_avg:60.11ms step:1468/1575 train_time:88263ms step_avg:60.12ms step:1469/1575 train_time:88352ms step_avg:60.14ms step:1470/1575 train_time:88438ms step_avg:60.16ms step:1471/1575 train_time:88528ms step_avg:60.18ms step:1472/1575 train_time:88614ms step_avg:60.20ms step:1473/1575 train_time:88705ms step_avg:60.22ms step:1474/1575 train_time:88790ms step_avg:60.24ms step:1475/1575 train_time:88880ms step_avg:60.26ms step:1476/1575 train_time:88966ms step_avg:60.27ms step:1477/1575 train_time:89055ms step_avg:60.29ms step:1478/1575 train_time:89140ms step_avg:60.31ms step:1479/1575 train_time:89229ms step_avg:60.33ms step:1480/1575 train_time:89315ms step_avg:60.35ms step:1481/1575 train_time:89407ms step_avg:60.37ms step:1482/1575 train_time:89491ms step_avg:60.39ms step:1483/1575 train_time:89581ms step_avg:60.40ms step:1484/1575 train_time:89666ms step_avg:60.42ms step:1485/1575 train_time:89755ms step_avg:60.44ms step:1486/1575 train_time:89841ms step_avg:60.46ms step:1487/1575 train_time:89931ms step_avg:60.48ms step:1488/1575 train_time:90017ms step_avg:60.50ms step:1489/1575 train_time:90107ms step_avg:60.51ms step:1490/1575 train_time:90194ms step_avg:60.53ms step:1491/1575 train_time:90283ms step_avg:60.55ms step:1492/1575 train_time:90369ms step_avg:60.57ms step:1493/1575 train_time:90458ms step_avg:60.59ms step:1494/1575 train_time:90543ms step_avg:60.60ms step:1495/1575 train_time:90632ms step_avg:60.62ms step:1496/1575 train_time:90719ms step_avg:60.64ms step:1497/1575 train_time:90810ms step_avg:60.66ms step:1498/1575 train_time:90895ms step_avg:60.68ms step:1499/1575 train_time:90986ms step_avg:60.70ms step:1500/1575 train_time:91071ms step_avg:60.71ms step:1500/1575 val_loss:3.2986 train_time:91144ms step_avg:60.76ms step:1501/1575 train_time:91165ms step_avg:60.74ms step:1502/1575 train_time:91253ms step_avg:60.75ms step:1503/1575 train_time:91345ms step_avg:60.78ms step:1504/1575 train_time:91435ms step_avg:60.79ms step:1505/1575 train_time:91524ms step_avg:60.81ms step:1506/1575 train_time:91608ms step_avg:60.83ms step:1507/1575 train_time:91696ms step_avg:60.85ms step:1508/1575 train_time:91781ms step_avg:60.86ms step:1509/1575 train_time:91870ms step_avg:60.88ms step:1510/1575 train_time:91956ms step_avg:60.90ms step:1511/1575 train_time:92045ms step_avg:60.92ms step:1512/1575 train_time:92131ms step_avg:60.93ms step:1513/1575 train_time:92223ms step_avg:60.95ms step:1514/1575 train_time:92310ms step_avg:60.97ms step:1515/1575 train_time:92402ms step_avg:60.99ms step:1516/1575 train_time:92487ms step_avg:61.01ms step:1517/1575 train_time:92577ms step_avg:61.03ms step:1518/1575 train_time:92662ms step_avg:61.04ms step:1519/1575 train_time:92750ms step_avg:61.06ms step:1520/1575 train_time:92835ms step_avg:61.08ms step:1521/1575 train_time:92924ms step_avg:61.09ms step:1522/1575 train_time:93009ms step_avg:61.11ms step:1523/1575 train_time:93098ms step_avg:61.13ms step:1524/1575 train_time:93187ms step_avg:61.15ms step:1525/1575 train_time:93277ms step_avg:61.17ms step:1526/1575 train_time:93366ms step_avg:61.18ms step:1527/1575 train_time:93456ms step_avg:61.20ms step:1528/1575 train_time:93541ms step_avg:61.22ms step:1529/1575 train_time:93632ms step_avg:61.24ms step:1530/1575 train_time:93717ms step_avg:61.25ms step:1531/1575 train_time:93805ms step_avg:61.27ms step:1532/1575 train_time:93890ms step_avg:61.29ms step:1533/1575 train_time:93979ms step_avg:61.30ms step:1534/1575 train_time:94065ms step_avg:61.32ms step:1535/1575 train_time:94155ms step_avg:61.34ms step:1536/1575 train_time:94249ms step_avg:61.36ms step:1537/1575 train_time:94338ms step_avg:61.38ms step:1538/1575 train_time:94425ms step_avg:61.39ms step:1539/1575 train_time:94515ms step_avg:61.41ms step:1540/1575 train_time:94601ms step_avg:61.43ms step:1541/1575 train_time:94690ms step_avg:61.45ms step:1542/1575 train_time:94776ms step_avg:61.46ms step:1543/1575 train_time:94864ms step_avg:61.48ms step:1544/1575 train_time:94950ms step_avg:61.50ms step:1545/1575 train_time:95040ms step_avg:61.51ms step:1546/1575 train_time:95126ms step_avg:61.53ms step:1547/1575 train_time:95217ms step_avg:61.55ms step:1548/1575 train_time:95303ms step_avg:61.57ms step:1549/1575 train_time:95395ms step_avg:61.58ms step:1550/1575 train_time:95481ms step_avg:61.60ms step:1551/1575 train_time:95572ms step_avg:61.62ms step:1552/1575 train_time:95659ms step_avg:61.64ms step:1553/1575 train_time:95747ms step_avg:61.65ms step:1554/1575 train_time:95833ms step_avg:61.67ms step:1555/1575 train_time:95924ms step_avg:61.69ms step:1556/1575 train_time:96010ms step_avg:61.70ms step:1557/1575 train_time:96100ms step_avg:61.72ms step:1558/1575 train_time:96187ms step_avg:61.74ms step:1559/1575 train_time:96278ms step_avg:61.76ms step:1560/1575 train_time:96364ms step_avg:61.77ms step:1561/1575 train_time:96453ms step_avg:61.79ms step:1562/1575 train_time:96540ms step_avg:61.81ms step:1563/1575 train_time:96630ms step_avg:61.82ms step:1564/1575 train_time:96716ms step_avg:61.84ms step:1565/1575 train_time:96806ms step_avg:61.86ms step:1566/1575 train_time:96892ms step_avg:61.87ms step:1567/1575 train_time:96983ms step_avg:61.89ms step:1568/1575 train_time:97068ms step_avg:61.91ms step:1569/1575 train_time:97162ms step_avg:61.93ms step:1570/1575 train_time:97247ms step_avg:61.94ms step:1571/1575 train_time:97337ms step_avg:61.96ms step:1572/1575 train_time:97422ms step_avg:61.97ms step:1573/1575 train_time:97512ms step_avg:61.99ms step:1574/1575 train_time:97598ms step_avg:62.01ms step:1575/1575 train_time:97688ms step_avg:62.02ms step:1575/1575 val_loss:3.2771 train_time:97755ms step_avg:62.07ms peak memory allocated: 31016 MiB reserved: 46918 MiB