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(3)]) 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 # 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 layers 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] # 012 ... 012 structure on token value embeddings by @YouJiacheng, improved on @leloykun's U-net structure # dropping first layer updates this to .12 ... 012 ve = [ve[1], ve[2]] + [None] * (self.num_layers - 5) + [ve[0], ve[1], ve[2]] 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}, "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", "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 = 1560 # 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 03:03:19 2026 +-----------------------------------------------------------------------------------------+ | NVIDIA-SMI 570.148.08 Driver Version: 570.148.08 CUDA Version: 12.8 | |-----------------------------------------+------------------------+----------------------+ | GPU Name Persistence-M | Bus-Id Disp.A | Volatile Uncorr. ECC | | Fan Temp Perf Pwr:Usage/Cap | Memory-Usage | GPU-Util Compute M. | | | | MIG M. | |=========================================+========================+======================| | 0 NVIDIA H100 80GB HBM3 On | 00000000:61:00.0 Off | 0 | | N/A 35C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:62:00.0 Off | 0 | | N/A 39C P0 122W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 41C P0 122W / 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 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 42C P0 130W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 40C P0 124W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 36C P0 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 281831 C /usr/bin/python3 1510MiB | | 1 N/A N/A 281832 C /usr/bin/python3 1510MiB | | 2 N/A N/A 281833 C /usr/bin/python3 1510MiB | | 3 N/A N/A 281834 C /usr/bin/python3 1510MiB | | 4 N/A N/A 281835 C /usr/bin/python3 1510MiB | | 5 N/A N/A 281836 C /usr/bin/python3 1510MiB | | 6 N/A N/A 281837 C /usr/bin/python3 1510MiB | | 7 N/A N/A 281838 C /usr/bin/python3 1510MiB | +-----------------------------------------------------------------------------------------+ ==================================================================================================== Compiling model and warming up kernels (~7 minutes on first execution) Sampling steps [0, 1, 2, 519, 520, 521, 1039, 1040, 1041, 1559, 1560, 1561] for warmup Resetting Model step:0/1600 val_loss:10.8275 train_time:0ms step_avg:0.03ms step:1/1600 train_time:73ms step_avg:72.94ms step:2/1600 train_time:96ms step_avg:47.78ms step:3/1600 train_time:115ms step_avg:38.33ms step:4/1600 train_time:137ms step_avg:34.30ms step:5/1600 train_time:167ms step_avg:33.43ms step:6/1600 train_time:274ms step_avg:45.60ms step:7/1600 train_time:292ms step_avg:41.70ms step:8/1600 train_time:369ms step_avg:46.16ms step:9/1600 train_time:400ms step_avg:44.39ms step:10/1600 train_time:437ms step_avg:43.66ms step:11/1600 train_time:467ms step_avg:42.46ms step:12/1600 train_time:505ms step_avg:42.05ms step:13/1600 train_time:535ms step_avg:41.16ms step:14/1600 train_time:572ms step_avg:40.89ms step:15/1600 train_time:603ms step_avg:40.21ms step:16/1600 train_time:641ms step_avg:40.06ms step:17/1600 train_time:672ms step_avg:39.51ms step:18/1600 train_time:709ms step_avg:39.39ms step:19/1600 train_time:740ms step_avg:38.92ms step:20/1600 train_time:777ms step_avg:38.84ms step:21/1600 train_time:807ms step_avg:38.43ms step:22/1600 train_time:845ms step_avg:38.39ms step:23/1600 train_time:875ms step_avg:38.04ms step:24/1600 train_time:912ms step_avg:38.01ms step:25/1600 train_time:943ms step_avg:37.71ms step:26/1600 train_time:981ms step_avg:37.72ms step:27/1600 train_time:1011ms step_avg:37.45ms step:28/1600 train_time:1048ms step_avg:37.44ms step:29/1600 train_time:1079ms step_avg:37.21ms step:30/1600 train_time:1116ms step_avg:37.20ms step:31/1600 train_time:1147ms step_avg:36.99ms step:32/1600 train_time:1185ms step_avg:37.03ms step:33/1600 train_time:1215ms step_avg:36.82ms step:34/1600 train_time:1252ms step_avg:36.83ms step:35/1600 train_time:1284ms step_avg:36.68ms step:36/1600 train_time:1321ms step_avg:36.70ms step:37/1600 train_time:1352ms step_avg:36.54ms step:38/1600 train_time:1390ms step_avg:36.57ms step:39/1600 train_time:1420ms step_avg:36.42ms step:40/1600 train_time:1458ms step_avg:36.45ms step:41/1600 train_time:1488ms step_avg:36.30ms step:42/1600 train_time:1526ms step_avg:36.33ms step:43/1600 train_time:1557ms step_avg:36.20ms step:44/1600 train_time:1594ms step_avg:36.22ms step:45/1600 train_time:1624ms step_avg:36.09ms step:46/1600 train_time:1662ms step_avg:36.13ms step:47/1600 train_time:1692ms step_avg:36.01ms step:48/1600 train_time:1730ms step_avg:36.04ms step:49/1600 train_time:1761ms step_avg:35.93ms step:50/1600 train_time:1798ms step_avg:35.97ms step:51/1600 train_time:1829ms step_avg:35.86ms step:52/1600 train_time:1866ms step_avg:35.88ms step:53/1600 train_time:1896ms step_avg:35.78ms step:54/1600 train_time:1933ms step_avg:35.80ms step:55/1600 train_time:1964ms step_avg:35.71ms step:56/1600 train_time:2001ms step_avg:35.74ms step:57/1600 train_time:2032ms step_avg:35.65ms step:58/1600 train_time:2070ms step_avg:35.68ms step:59/1600 train_time:2100ms step_avg:35.60ms step:60/1600 train_time:2137ms step_avg:35.62ms step:61/1600 train_time:2168ms step_avg:35.54ms step:62/1600 train_time:2206ms step_avg:35.57ms step:63/1600 train_time:2236ms step_avg:35.49ms step:64/1600 train_time:2273ms step_avg:35.52ms step:65/1600 train_time:2304ms step_avg:35.45ms step:66/1600 train_time:2342ms step_avg:35.49ms step:67/1600 train_time:2373ms step_avg:35.41ms step:68/1600 train_time:2410ms step_avg:35.44ms step:69/1600 train_time:2440ms step_avg:35.37ms step:70/1600 train_time:2477ms step_avg:35.39ms step:71/1600 train_time:2508ms step_avg:35.32ms step:72/1600 train_time:2545ms step_avg:35.35ms step:73/1600 train_time:2576ms step_avg:35.29ms step:74/1600 train_time:2613ms step_avg:35.31ms step:75/1600 train_time:2644ms step_avg:35.25ms step:76/1600 train_time:2681ms step_avg:35.27ms step:77/1600 train_time:2712ms step_avg:35.22ms step:78/1600 train_time:2749ms step_avg:35.25ms step:79/1600 train_time:2780ms step_avg:35.19ms step:80/1600 train_time:2818ms step_avg:35.22ms step:81/1600 train_time:2848ms step_avg:35.16ms step:82/1600 train_time:2886ms step_avg:35.19ms step:83/1600 train_time:2916ms step_avg:35.13ms step:84/1600 train_time:2953ms step_avg:35.16ms step:85/1600 train_time:2984ms step_avg:35.11ms step:86/1600 train_time:3022ms step_avg:35.14ms step:87/1600 train_time:3052ms step_avg:35.09ms step:88/1600 train_time:3090ms step_avg:35.11ms step:89/1600 train_time:3121ms step_avg:35.06ms step:90/1600 train_time:3158ms step_avg:35.09ms step:91/1600 train_time:3189ms step_avg:35.05ms step:92/1600 train_time:3226ms step_avg:35.07ms step:93/1600 train_time:3257ms step_avg:35.02ms step:94/1600 train_time:3294ms step_avg:35.04ms step:95/1600 train_time:3325ms step_avg:35.00ms step:96/1600 train_time:3363ms step_avg:35.03ms step:97/1600 train_time:3394ms step_avg:34.99ms step:98/1600 train_time:3431ms step_avg:35.01ms step:99/1600 train_time:3462ms step_avg:34.97ms step:100/1600 train_time:3499ms step_avg:34.99ms step:101/1600 train_time:3530ms step_avg:34.95ms step:102/1600 train_time:3567ms step_avg:34.97ms step:103/1600 train_time:3598ms step_avg:34.93ms step:104/1600 train_time:3635ms step_avg:34.95ms step:105/1600 train_time:3666ms step_avg:34.91ms step:106/1600 train_time:3704ms step_avg:34.94ms step:107/1600 train_time:3735ms step_avg:34.90ms step:108/1600 train_time:3772ms step_avg:34.92ms step:109/1600 train_time:3802ms step_avg:34.88ms step:110/1600 train_time:3840ms step_avg:34.90ms step:111/1600 train_time:3870ms step_avg:34.87ms step:112/1600 train_time:3907ms step_avg:34.89ms step:113/1600 train_time:3938ms step_avg:34.85ms step:114/1600 train_time:3975ms step_avg:34.87ms step:115/1600 train_time:4005ms step_avg:34.83ms step:116/1600 train_time:4043ms step_avg:34.85ms step:117/1600 train_time:4073ms step_avg:34.82ms step:118/1600 train_time:4111ms step_avg:34.84ms step:119/1600 train_time:4142ms step_avg:34.80ms step:120/1600 train_time:4179ms step_avg:34.83ms step:121/1600 train_time:4210ms step_avg:34.79ms step:122/1600 train_time:4247ms step_avg:34.81ms step:123/1600 train_time:4278ms step_avg:34.78ms step:124/1600 train_time:4315ms step_avg:34.80ms step:125/1600 train_time:4346ms step_avg:34.77ms step:126/1600 train_time:4383ms step_avg:34.79ms step:127/1600 train_time:4415ms step_avg:34.76ms step:128/1600 train_time:4452ms step_avg:34.78ms step:129/1600 train_time:4482ms step_avg:34.75ms step:130/1600 train_time:4520ms step_avg:34.77ms step:131/1600 train_time:4550ms step_avg:34.73ms step:132/1600 train_time:4588ms step_avg:34.75ms step:133/1600 train_time:4618ms step_avg:34.72ms step:134/1600 train_time:4656ms step_avg:34.74ms step:135/1600 train_time:4686ms step_avg:34.71ms step:136/1600 train_time:4723ms step_avg:34.73ms step:137/1600 train_time:4754ms step_avg:34.70ms step:138/1600 train_time:4791ms step_avg:34.72ms step:139/1600 train_time:4822ms step_avg:34.69ms step:140/1600 train_time:4859ms step_avg:34.71ms step:141/1600 train_time:4890ms step_avg:34.68ms step:142/1600 train_time:4927ms step_avg:34.70ms step:143/1600 train_time:4957ms step_avg:34.67ms step:144/1600 train_time:4994ms step_avg:34.68ms step:145/1600 train_time:5025ms step_avg:34.65ms step:146/1600 train_time:5062ms step_avg:34.67ms step:147/1600 train_time:5093ms step_avg:34.64ms step:148/1600 train_time:5130ms step_avg:34.67ms step:149/1600 train_time:5161ms step_avg:34.64ms step:150/1600 train_time:5198ms step_avg:34.65ms step:151/1600 train_time:5228ms step_avg:34.62ms step:152/1600 train_time:5266ms step_avg:34.64ms step:153/1600 train_time:5296ms step_avg:34.62ms step:154/1600 train_time:5334ms step_avg:34.63ms step:155/1600 train_time:5365ms step_avg:34.61ms step:156/1600 train_time:5402ms step_avg:34.63ms step:157/1600 train_time:5433ms step_avg:34.60ms step:158/1600 train_time:5470ms step_avg:34.62ms step:159/1600 train_time:5500ms step_avg:34.59ms step:160/1600 train_time:5538ms step_avg:34.61ms step:161/1600 train_time:5568ms step_avg:34.59ms step:162/1600 train_time:5606ms step_avg:34.60ms step:163/1600 train_time:5636ms step_avg:34.58ms step:164/1600 train_time:5673ms step_avg:34.59ms step:165/1600 train_time:5705ms step_avg:34.57ms step:166/1600 train_time:5742ms step_avg:34.59ms step:167/1600 train_time:5773ms step_avg:34.57ms step:168/1600 train_time:5811ms step_avg:34.59ms step:169/1600 train_time:5841ms step_avg:34.56ms step:170/1600 train_time:5879ms step_avg:34.58ms step:171/1600 train_time:5909ms step_avg:34.56ms step:172/1600 train_time:5946ms step_avg:34.57ms step:173/1600 train_time:5977ms step_avg:34.55ms step:174/1600 train_time:6014ms step_avg:34.56ms step:175/1600 train_time:6045ms step_avg:34.54ms step:176/1600 train_time:6083ms step_avg:34.56ms step:177/1600 train_time:6113ms step_avg:34.54ms step:178/1600 train_time:6151ms step_avg:34.55ms step:179/1600 train_time:6182ms step_avg:34.53ms step:180/1600 train_time:6219ms step_avg:34.55ms step:181/1600 train_time:6250ms step_avg:34.53ms step:182/1600 train_time:6287ms step_avg:34.54ms step:183/1600 train_time:6317ms step_avg:34.52ms step:184/1600 train_time:6354ms step_avg:34.53ms step:185/1600 train_time:6385ms step_avg:34.51ms step:186/1600 train_time:6422ms step_avg:34.53ms step:187/1600 train_time:6453ms step_avg:34.51ms step:188/1600 train_time:6490ms step_avg:34.52ms step:189/1600 train_time:6520ms step_avg:34.50ms step:190/1600 train_time:6558ms step_avg:34.51ms step:191/1600 train_time:6588ms step_avg:34.49ms step:192/1600 train_time:6626ms step_avg:34.51ms step:193/1600 train_time:6656ms step_avg:34.49ms step:194/1600 train_time:6693ms step_avg:34.50ms step:195/1600 train_time:6724ms step_avg:34.48ms step:196/1600 train_time:6762ms step_avg:34.50ms step:197/1600 train_time:6793ms step_avg:34.48ms step:198/1600 train_time:6830ms step_avg:34.50ms step:199/1600 train_time:6861ms step_avg:34.48ms step:200/1600 train_time:6898ms step_avg:34.49ms step:201/1600 train_time:6928ms step_avg:34.47ms step:202/1600 train_time:6966ms step_avg:34.48ms step:203/1600 train_time:6996ms step_avg:34.46ms step:204/1600 train_time:7033ms step_avg:34.48ms step:205/1600 train_time:7064ms step_avg:34.46ms step:206/1600 train_time:7102ms step_avg:34.48ms step:207/1600 train_time:7133ms step_avg:34.46ms step:208/1600 train_time:7171ms step_avg:34.47ms step:209/1600 train_time:7201ms step_avg:34.45ms step:210/1600 train_time:7239ms step_avg:34.47ms step:211/1600 train_time:7269ms step_avg:34.45ms step:212/1600 train_time:7307ms step_avg:34.46ms step:213/1600 train_time:7337ms step_avg:34.45ms step:214/1600 train_time:7374ms step_avg:34.46ms step:215/1600 train_time:7405ms step_avg:34.44ms step:216/1600 train_time:7443ms step_avg:34.46ms step:217/1600 train_time:7473ms step_avg:34.44ms step:218/1600 train_time:7510ms step_avg:34.45ms step:219/1600 train_time:7541ms step_avg:34.43ms step:220/1600 train_time:7579ms step_avg:34.45ms step:221/1600 train_time:7609ms step_avg:34.43ms step:222/1600 train_time:7646ms step_avg:34.44ms step:223/1600 train_time:7676ms step_avg:34.42ms step:224/1600 train_time:7713ms step_avg:34.43ms step:225/1600 train_time:7743ms step_avg:34.42ms step:226/1600 train_time:7781ms step_avg:34.43ms step:227/1600 train_time:7811ms step_avg:34.41ms step:228/1600 train_time:7848ms step_avg:34.42ms step:229/1600 train_time:7879ms step_avg:34.41ms step:230/1600 train_time:7916ms step_avg:34.42ms step:231/1600 train_time:7947ms step_avg:34.40ms step:232/1600 train_time:7984ms step_avg:34.42ms step:233/1600 train_time:8015ms step_avg:34.40ms step:234/1600 train_time:8052ms step_avg:34.41ms step:235/1600 train_time:8083ms step_avg:34.39ms step:236/1600 train_time:8120ms step_avg:34.41ms step:237/1600 train_time:8150ms step_avg:34.39ms step:238/1600 train_time:8188ms step_avg:34.40ms step:239/1600 train_time:8218ms step_avg:34.38ms step:240/1600 train_time:8255ms step_avg:34.40ms step:241/1600 train_time:8286ms step_avg:34.38ms step:242/1600 train_time:8324ms step_avg:34.39ms step:243/1600 train_time:8354ms step_avg:34.38ms step:244/1600 train_time:8392ms step_avg:34.39ms step:245/1600 train_time:8422ms step_avg:34.38ms step:246/1600 train_time:8459ms step_avg:34.39ms step:247/1600 train_time:8490ms step_avg:34.37ms step:248/1600 train_time:8527ms step_avg:34.39ms step:249/1600 train_time:8558ms step_avg:34.37ms step:250/1600 train_time:8595ms step_avg:34.38ms step:250/1600 val_loss:4.5815 train_time:8642ms step_avg:34.57ms step:251/1600 train_time:8663ms step_avg:34.51ms step:252/1600 train_time:8683ms step_avg:34.45ms step:253/1600 train_time:8700ms step_avg:34.39ms step:254/1600 train_time:8733ms step_avg:34.38ms step:255/1600 train_time:8765ms step_avg:34.37ms step:256/1600 train_time:8804ms step_avg:34.39ms step:257/1600 train_time:8835ms step_avg:34.38ms step:258/1600 train_time:8873ms step_avg:34.39ms step:259/1600 train_time:8904ms step_avg:34.38ms step:260/1600 train_time:8941ms step_avg:34.39ms step:261/1600 train_time:8972ms step_avg:34.37ms step:262/1600 train_time:9009ms step_avg:34.39ms step:263/1600 train_time:9040ms step_avg:34.37ms step:264/1600 train_time:9077ms step_avg:34.38ms step:265/1600 train_time:9108ms step_avg:34.37ms step:266/1600 train_time:9145ms step_avg:34.38ms step:267/1600 train_time:9176ms step_avg:34.37ms step:268/1600 train_time:9213ms step_avg:34.38ms step:269/1600 train_time:9243ms step_avg:34.36ms step:270/1600 train_time:9280ms step_avg:34.37ms step:271/1600 train_time:9311ms step_avg:34.36ms step:272/1600 train_time:9348ms step_avg:34.37ms step:273/1600 train_time:9378ms step_avg:34.35ms step:274/1600 train_time:9415ms step_avg:34.36ms step:275/1600 train_time:9446ms step_avg:34.35ms step:276/1600 train_time:9483ms step_avg:34.36ms step:277/1600 train_time:9513ms step_avg:34.34ms step:278/1600 train_time:9551ms step_avg:34.36ms step:279/1600 train_time:9581ms step_avg:34.34ms step:280/1600 train_time:9618ms step_avg:34.35ms step:281/1600 train_time:9649ms step_avg:34.34ms step:282/1600 train_time:9686ms step_avg:34.35ms step:283/1600 train_time:9716ms step_avg:34.33ms step:284/1600 train_time:9754ms step_avg:34.34ms step:285/1600 train_time:9784ms step_avg:34.33ms step:286/1600 train_time:9821ms step_avg:34.34ms step:287/1600 train_time:9853ms step_avg:34.33ms step:288/1600 train_time:9890ms step_avg:34.34ms step:289/1600 train_time:9921ms step_avg:34.33ms step:290/1600 train_time:9958ms step_avg:34.34ms step:291/1600 train_time:9989ms step_avg:34.33ms step:292/1600 train_time:10026ms step_avg:34.34ms step:293/1600 train_time:10057ms step_avg:34.32ms step:294/1600 train_time:10094ms step_avg:34.33ms step:295/1600 train_time:10125ms step_avg:34.32ms step:296/1600 train_time:10162ms step_avg:34.33ms step:297/1600 train_time:10192ms step_avg:34.32ms step:298/1600 train_time:10230ms step_avg:34.33ms step:299/1600 train_time:10260ms step_avg:34.31ms step:300/1600 train_time:10297ms step_avg:34.32ms step:301/1600 train_time:10328ms step_avg:34.31ms step:302/1600 train_time:10365ms step_avg:34.32ms step:303/1600 train_time:10396ms step_avg:34.31ms step:304/1600 train_time:10433ms step_avg:34.32ms step:305/1600 train_time:10463ms step_avg:34.31ms step:306/1600 train_time:10501ms step_avg:34.32ms step:307/1600 train_time:10531ms step_avg:34.30ms step:308/1600 train_time:10569ms step_avg:34.31ms step:309/1600 train_time:10599ms step_avg:34.30ms step:310/1600 train_time:10636ms step_avg:34.31ms step:311/1600 train_time:10667ms step_avg:34.30ms step:312/1600 train_time:10704ms step_avg:34.31ms step:313/1600 train_time:10734ms step_avg:34.30ms step:314/1600 train_time:10772ms step_avg:34.31ms step:315/1600 train_time:10803ms step_avg:34.29ms step:316/1600 train_time:10840ms step_avg:34.30ms step:317/1600 train_time:10871ms step_avg:34.29ms step:318/1600 train_time:10908ms step_avg:34.30ms step:319/1600 train_time:10939ms step_avg:34.29ms step:320/1600 train_time:10977ms step_avg:34.30ms step:321/1600 train_time:11007ms step_avg:34.29ms step:322/1600 train_time:11044ms step_avg:34.30ms step:323/1600 train_time:11074ms step_avg:34.29ms step:324/1600 train_time:11112ms step_avg:34.30ms step:325/1600 train_time:11142ms step_avg:34.28ms step:326/1600 train_time:11179ms step_avg:34.29ms step:327/1600 train_time:11210ms step_avg:34.28ms step:328/1600 train_time:11247ms step_avg:34.29ms step:329/1600 train_time:11277ms step_avg:34.28ms step:330/1600 train_time:11315ms step_avg:34.29ms step:331/1600 train_time:11345ms step_avg:34.28ms step:332/1600 train_time:11382ms step_avg:34.28ms step:333/1600 train_time:11413ms step_avg:34.27ms step:334/1600 train_time:11451ms step_avg:34.28ms step:335/1600 train_time:11481ms step_avg:34.27ms step:336/1600 train_time:11518ms step_avg:34.28ms step:337/1600 train_time:11549ms step_avg:34.27ms step:338/1600 train_time:11586ms step_avg:34.28ms step:339/1600 train_time:11616ms step_avg:34.27ms step:340/1600 train_time:11654ms step_avg:34.28ms step:341/1600 train_time:11684ms step_avg:34.26ms step:342/1600 train_time:11721ms step_avg:34.27ms step:343/1600 train_time:11752ms step_avg:34.26ms step:344/1600 train_time:11790ms step_avg:34.27ms step:345/1600 train_time:11820ms step_avg:34.26ms step:346/1600 train_time:11858ms step_avg:34.27ms step:347/1600 train_time:11889ms step_avg:34.26ms step:348/1600 train_time:11926ms step_avg:34.27ms step:349/1600 train_time:11957ms step_avg:34.26ms step:350/1600 train_time:11994ms step_avg:34.27ms step:351/1600 train_time:12024ms step_avg:34.26ms step:352/1600 train_time:12061ms step_avg:34.27ms step:353/1600 train_time:12092ms step_avg:34.25ms step:354/1600 train_time:12129ms step_avg:34.26ms step:355/1600 train_time:12160ms step_avg:34.25ms step:356/1600 train_time:12197ms step_avg:34.26ms step:357/1600 train_time:12228ms step_avg:34.25ms step:358/1600 train_time:12265ms step_avg:34.26ms step:359/1600 train_time:12295ms step_avg:34.25ms step:360/1600 train_time:12333ms step_avg:34.26ms step:361/1600 train_time:12364ms step_avg:34.25ms step:362/1600 train_time:12401ms step_avg:34.26ms step:363/1600 train_time:12431ms step_avg:34.25ms step:364/1600 train_time:12469ms step_avg:34.25ms step:365/1600 train_time:12499ms step_avg:34.24ms step:366/1600 train_time:12536ms step_avg:34.25ms step:367/1600 train_time:12567ms step_avg:34.24ms step:368/1600 train_time:12604ms step_avg:34.25ms step:369/1600 train_time:12635ms step_avg:34.24ms step:370/1600 train_time:12672ms step_avg:34.25ms step:371/1600 train_time:12702ms step_avg:34.24ms step:372/1600 train_time:12739ms step_avg:34.24ms step:373/1600 train_time:12769ms step_avg:34.23ms step:374/1600 train_time:12807ms step_avg:34.24ms step:375/1600 train_time:12837ms step_avg:34.23ms step:376/1600 train_time:12874ms step_avg:34.24ms step:377/1600 train_time:12904ms step_avg:34.23ms step:378/1600 train_time:12942ms step_avg:34.24ms step:379/1600 train_time:12973ms step_avg:34.23ms step:380/1600 train_time:13011ms step_avg:34.24ms step:381/1600 train_time:13041ms step_avg:34.23ms step:382/1600 train_time:13078ms step_avg:34.24ms step:383/1600 train_time:13109ms step_avg:34.23ms step:384/1600 train_time:13146ms step_avg:34.23ms step:385/1600 train_time:13177ms step_avg:34.23ms step:386/1600 train_time:13214ms step_avg:34.23ms step:387/1600 train_time:13244ms step_avg:34.22ms step:388/1600 train_time:13281ms step_avg:34.23ms step:389/1600 train_time:13311ms step_avg:34.22ms step:390/1600 train_time:13349ms step_avg:34.23ms step:391/1600 train_time:13380ms step_avg:34.22ms step:392/1600 train_time:13417ms step_avg:34.23ms step:393/1600 train_time:13447ms step_avg:34.22ms step:394/1600 train_time:13485ms step_avg:34.23ms step:395/1600 train_time:13515ms step_avg:34.22ms step:396/1600 train_time:13552ms step_avg:34.22ms step:397/1600 train_time:13583ms step_avg:34.21ms step:398/1600 train_time:13620ms step_avg:34.22ms step:399/1600 train_time:13651ms step_avg:34.21ms step:400/1600 train_time:13688ms step_avg:34.22ms step:401/1600 train_time:13719ms step_avg:34.21ms step:402/1600 train_time:13756ms step_avg:34.22ms step:403/1600 train_time:13786ms step_avg:34.21ms step:404/1600 train_time:13823ms step_avg:34.22ms step:405/1600 train_time:13854ms step_avg:34.21ms step:406/1600 train_time:13891ms step_avg:34.22ms step:407/1600 train_time:13922ms step_avg:34.21ms step:408/1600 train_time:13959ms step_avg:34.21ms step:409/1600 train_time:13990ms step_avg:34.21ms step:410/1600 train_time:14027ms step_avg:34.21ms step:411/1600 train_time:14058ms step_avg:34.20ms step:412/1600 train_time:14095ms step_avg:34.21ms step:413/1600 train_time:14126ms step_avg:34.20ms step:414/1600 train_time:14163ms step_avg:34.21ms step:415/1600 train_time:14193ms step_avg:34.20ms step:416/1600 train_time:14231ms step_avg:34.21ms step:417/1600 train_time:14261ms step_avg:34.20ms step:418/1600 train_time:14299ms step_avg:34.21ms step:419/1600 train_time:14329ms step_avg:34.20ms step:420/1600 train_time:14366ms step_avg:34.21ms step:421/1600 train_time:14397ms step_avg:34.20ms step:422/1600 train_time:14434ms step_avg:34.20ms step:423/1600 train_time:14465ms step_avg:34.20ms step:424/1600 train_time:14502ms step_avg:34.20ms step:425/1600 train_time:14532ms step_avg:34.19ms step:426/1600 train_time:14570ms step_avg:34.20ms step:427/1600 train_time:14601ms step_avg:34.19ms step:428/1600 train_time:14638ms step_avg:34.20ms step:429/1600 train_time:14669ms step_avg:34.19ms step:430/1600 train_time:14706ms step_avg:34.20ms step:431/1600 train_time:14737ms step_avg:34.19ms step:432/1600 train_time:14774ms step_avg:34.20ms step:433/1600 train_time:14805ms step_avg:34.19ms step:434/1600 train_time:14842ms step_avg:34.20ms step:435/1600 train_time:14872ms step_avg:34.19ms step:436/1600 train_time:14910ms step_avg:34.20ms step:437/1600 train_time:14940ms step_avg:34.19ms step:438/1600 train_time:14978ms step_avg:34.20ms step:439/1600 train_time:15008ms step_avg:34.19ms step:440/1600 train_time:15045ms step_avg:34.19ms step:441/1600 train_time:15075ms step_avg:34.18ms step:442/1600 train_time:15113ms step_avg:34.19ms step:443/1600 train_time:15144ms step_avg:34.18ms step:444/1600 train_time:15181ms step_avg:34.19ms step:445/1600 train_time:15211ms step_avg:34.18ms step:446/1600 train_time:15249ms step_avg:34.19ms step:447/1600 train_time:15279ms step_avg:34.18ms step:448/1600 train_time:15316ms step_avg:34.19ms step:449/1600 train_time:15347ms step_avg:34.18ms step:450/1600 train_time:15384ms step_avg:34.19ms step:451/1600 train_time:15415ms step_avg:34.18ms step:452/1600 train_time:15452ms step_avg:34.19ms step:453/1600 train_time:15482ms step_avg:34.18ms step:454/1600 train_time:15520ms step_avg:34.18ms step:455/1600 train_time:15551ms step_avg:34.18ms step:456/1600 train_time:15588ms step_avg:34.18ms step:457/1600 train_time:15618ms step_avg:34.18ms step:458/1600 train_time:15656ms step_avg:34.18ms step:459/1600 train_time:15686ms step_avg:34.17ms step:460/1600 train_time:15723ms step_avg:34.18ms step:461/1600 train_time:15754ms step_avg:34.17ms step:462/1600 train_time:15791ms step_avg:34.18ms step:463/1600 train_time:15822ms step_avg:34.17ms step:464/1600 train_time:15859ms step_avg:34.18ms step:465/1600 train_time:15889ms step_avg:34.17ms step:466/1600 train_time:15926ms step_avg:34.18ms step:467/1600 train_time:15957ms step_avg:34.17ms step:468/1600 train_time:15994ms step_avg:34.18ms step:469/1600 train_time:16025ms step_avg:34.17ms step:470/1600 train_time:16061ms step_avg:34.17ms step:471/1600 train_time:16092ms step_avg:34.17ms step:472/1600 train_time:16130ms step_avg:34.17ms step:473/1600 train_time:16161ms step_avg:34.17ms step:474/1600 train_time:16198ms step_avg:34.17ms step:475/1600 train_time:16229ms step_avg:34.17ms step:476/1600 train_time:16267ms step_avg:34.17ms step:477/1600 train_time:16297ms step_avg:34.16ms step:478/1600 train_time:16334ms step_avg:34.17ms step:479/1600 train_time:16364ms step_avg:34.16ms step:480/1600 train_time:16401ms step_avg:34.17ms step:481/1600 train_time:16432ms step_avg:34.16ms step:482/1600 train_time:16469ms step_avg:34.17ms step:483/1600 train_time:16499ms step_avg:34.16ms step:484/1600 train_time:16536ms step_avg:34.17ms step:485/1600 train_time:16567ms step_avg:34.16ms step:486/1600 train_time:16605ms step_avg:34.17ms step:487/1600 train_time:16636ms step_avg:34.16ms step:488/1600 train_time:16673ms step_avg:34.17ms step:489/1600 train_time:16703ms step_avg:34.16ms step:490/1600 train_time:16740ms step_avg:34.16ms step:491/1600 train_time:16771ms step_avg:34.16ms step:492/1600 train_time:16808ms step_avg:34.16ms step:493/1600 train_time:16839ms step_avg:34.16ms step:494/1600 train_time:16877ms step_avg:34.16ms step:495/1600 train_time:16907ms step_avg:34.16ms step:496/1600 train_time:16944ms step_avg:34.16ms step:497/1600 train_time:16974ms step_avg:34.15ms step:498/1600 train_time:17012ms step_avg:34.16ms step:499/1600 train_time:17043ms step_avg:34.15ms step:500/1600 train_time:17080ms step_avg:34.16ms step:500/1600 val_loss:4.2462 train_time:17128ms step_avg:34.26ms step:501/1600 train_time:17148ms step_avg:34.23ms step:502/1600 train_time:17168ms step_avg:34.20ms step:503/1600 train_time:17186ms step_avg:34.17ms step:504/1600 train_time:17222ms step_avg:34.17ms step:505/1600 train_time:17254ms step_avg:34.17ms step:506/1600 train_time:17292ms step_avg:34.17ms step:507/1600 train_time:17324ms step_avg:34.17ms step:508/1600 train_time:17361ms step_avg:34.18ms step:509/1600 train_time:17392ms step_avg:34.17ms step:510/1600 train_time:17429ms step_avg:34.17ms step:511/1600 train_time:17459ms step_avg:34.17ms step:512/1600 train_time:17497ms step_avg:34.17ms step:513/1600 train_time:17527ms step_avg:34.17ms step:514/1600 train_time:17564ms step_avg:34.17ms step:515/1600 train_time:17595ms step_avg:34.16ms step:516/1600 train_time:17632ms step_avg:34.17ms step:517/1600 train_time:17662ms step_avg:34.16ms step:518/1600 train_time:17699ms step_avg:34.17ms step:519/1600 train_time:17729ms step_avg:34.16ms step:520/1600 train_time:17766ms step_avg:34.17ms step:521/1600 train_time:17838ms step_avg:34.24ms step:522/1600 train_time:17895ms step_avg:34.28ms step:523/1600 train_time:17956ms step_avg:34.33ms step:524/1600 train_time:18014ms step_avg:34.38ms step:525/1600 train_time:18075ms step_avg:34.43ms step:526/1600 train_time:18134ms step_avg:34.48ms step:527/1600 train_time:18198ms step_avg:34.53ms step:528/1600 train_time:18259ms step_avg:34.58ms step:529/1600 train_time:18323ms step_avg:34.64ms step:530/1600 train_time:18381ms step_avg:34.68ms step:531/1600 train_time:18446ms step_avg:34.74ms step:532/1600 train_time:18504ms step_avg:34.78ms step:533/1600 train_time:18567ms step_avg:34.83ms step:534/1600 train_time:18627ms step_avg:34.88ms step:535/1600 train_time:18689ms step_avg:34.93ms step:536/1600 train_time:18748ms step_avg:34.98ms step:537/1600 train_time:18810ms step_avg:35.03ms step:538/1600 train_time:18869ms step_avg:35.07ms step:539/1600 train_time:18931ms step_avg:35.12ms step:540/1600 train_time:18991ms step_avg:35.17ms step:541/1600 train_time:19052ms step_avg:35.22ms step:542/1600 train_time:19110ms step_avg:35.26ms step:543/1600 train_time:19172ms step_avg:35.31ms step:544/1600 train_time:19231ms step_avg:35.35ms step:545/1600 train_time:19293ms step_avg:35.40ms step:546/1600 train_time:19352ms step_avg:35.44ms step:547/1600 train_time:19414ms step_avg:35.49ms step:548/1600 train_time:19473ms step_avg:35.53ms step:549/1600 train_time:19536ms step_avg:35.59ms step:550/1600 train_time:19597ms step_avg:35.63ms step:551/1600 train_time:19660ms step_avg:35.68ms step:552/1600 train_time:19719ms step_avg:35.72ms step:553/1600 train_time:19782ms step_avg:35.77ms step:554/1600 train_time:19841ms step_avg:35.81ms step:555/1600 train_time:19903ms step_avg:35.86ms step:556/1600 train_time:19961ms step_avg:35.90ms step:557/1600 train_time:20024ms step_avg:35.95ms step:558/1600 train_time:20084ms step_avg:35.99ms step:559/1600 train_time:20147ms step_avg:36.04ms step:560/1600 train_time:20207ms step_avg:36.08ms step:561/1600 train_time:20269ms step_avg:36.13ms step:562/1600 train_time:20328ms step_avg:36.17ms step:563/1600 train_time:20391ms step_avg:36.22ms step:564/1600 train_time:20449ms step_avg:36.26ms step:565/1600 train_time:20512ms step_avg:36.30ms step:566/1600 train_time:20571ms step_avg:36.35ms step:567/1600 train_time:20633ms step_avg:36.39ms step:568/1600 train_time:20692ms step_avg:36.43ms step:569/1600 train_time:20754ms step_avg:36.48ms step:570/1600 train_time:20813ms step_avg:36.51ms step:571/1600 train_time:20875ms step_avg:36.56ms step:572/1600 train_time:20936ms step_avg:36.60ms step:573/1600 train_time:21002ms step_avg:36.65ms step:574/1600 train_time:21058ms step_avg:36.69ms step:575/1600 train_time:21121ms step_avg:36.73ms step:576/1600 train_time:21179ms step_avg:36.77ms step:577/1600 train_time:21240ms step_avg:36.81ms step:578/1600 train_time:21300ms step_avg:36.85ms step:579/1600 train_time:21362ms step_avg:36.90ms step:580/1600 train_time:21421ms step_avg:36.93ms step:581/1600 train_time:21483ms step_avg:36.98ms step:582/1600 train_time:21543ms step_avg:37.02ms step:583/1600 train_time:21604ms step_avg:37.06ms step:584/1600 train_time:21665ms step_avg:37.10ms step:585/1600 train_time:21727ms step_avg:37.14ms step:586/1600 train_time:21786ms step_avg:37.18ms step:587/1600 train_time:21848ms step_avg:37.22ms step:588/1600 train_time:21908ms step_avg:37.26ms step:589/1600 train_time:21971ms step_avg:37.30ms step:590/1600 train_time:22029ms step_avg:37.34ms step:591/1600 train_time:22091ms step_avg:37.38ms step:592/1600 train_time:22150ms step_avg:37.42ms step:593/1600 train_time:22212ms step_avg:37.46ms step:594/1600 train_time:22271ms step_avg:37.49ms step:595/1600 train_time:22333ms step_avg:37.53ms step:596/1600 train_time:22393ms step_avg:37.57ms step:597/1600 train_time:22455ms step_avg:37.61ms step:598/1600 train_time:22513ms step_avg:37.65ms step:599/1600 train_time:22576ms step_avg:37.69ms step:600/1600 train_time:22635ms step_avg:37.73ms step:601/1600 train_time:22698ms step_avg:37.77ms step:602/1600 train_time:22758ms step_avg:37.80ms step:603/1600 train_time:22820ms step_avg:37.84ms step:604/1600 train_time:22880ms step_avg:37.88ms step:605/1600 train_time:22942ms step_avg:37.92ms step:606/1600 train_time:23002ms step_avg:37.96ms step:607/1600 train_time:23063ms step_avg:38.00ms step:608/1600 train_time:23122ms step_avg:38.03ms step:609/1600 train_time:23185ms step_avg:38.07ms step:610/1600 train_time:23245ms step_avg:38.11ms step:611/1600 train_time:23307ms step_avg:38.15ms step:612/1600 train_time:23366ms step_avg:38.18ms step:613/1600 train_time:23429ms step_avg:38.22ms step:614/1600 train_time:23488ms step_avg:38.25ms step:615/1600 train_time:23551ms step_avg:38.29ms step:616/1600 train_time:23610ms step_avg:38.33ms step:617/1600 train_time:23672ms step_avg:38.37ms step:618/1600 train_time:23731ms step_avg:38.40ms step:619/1600 train_time:23793ms step_avg:38.44ms step:620/1600 train_time:23851ms step_avg:38.47ms step:621/1600 train_time:23913ms step_avg:38.51ms step:622/1600 train_time:23971ms step_avg:38.54ms step:623/1600 train_time:24033ms step_avg:38.58ms step:624/1600 train_time:24092ms step_avg:38.61ms step:625/1600 train_time:24153ms step_avg:38.65ms step:626/1600 train_time:24212ms step_avg:38.68ms step:627/1600 train_time:24275ms step_avg:38.72ms step:628/1600 train_time:24334ms step_avg:38.75ms step:629/1600 train_time:24396ms step_avg:38.79ms step:630/1600 train_time:24455ms step_avg:38.82ms step:631/1600 train_time:24518ms step_avg:38.86ms step:632/1600 train_time:24578ms step_avg:38.89ms step:633/1600 train_time:24641ms step_avg:38.93ms step:634/1600 train_time:24700ms step_avg:38.96ms step:635/1600 train_time:24763ms step_avg:39.00ms step:636/1600 train_time:24822ms step_avg:39.03ms step:637/1600 train_time:24884ms step_avg:39.07ms step:638/1600 train_time:24944ms step_avg:39.10ms step:639/1600 train_time:25007ms step_avg:39.14ms step:640/1600 train_time:25066ms step_avg:39.17ms step:641/1600 train_time:25129ms step_avg:39.20ms step:642/1600 train_time:25187ms step_avg:39.23ms step:643/1600 train_time:25250ms step_avg:39.27ms step:644/1600 train_time:25309ms step_avg:39.30ms step:645/1600 train_time:25371ms step_avg:39.34ms step:646/1600 train_time:25430ms step_avg:39.37ms step:647/1600 train_time:25493ms step_avg:39.40ms step:648/1600 train_time:25551ms step_avg:39.43ms step:649/1600 train_time:25612ms step_avg:39.46ms step:650/1600 train_time:25671ms step_avg:39.49ms step:651/1600 train_time:25733ms step_avg:39.53ms step:652/1600 train_time:25792ms step_avg:39.56ms step:653/1600 train_time:25855ms step_avg:39.59ms step:654/1600 train_time:25914ms step_avg:39.62ms step:655/1600 train_time:25976ms step_avg:39.66ms step:656/1600 train_time:26036ms step_avg:39.69ms step:657/1600 train_time:26098ms step_avg:39.72ms step:658/1600 train_time:26157ms step_avg:39.75ms step:659/1600 train_time:26222ms step_avg:39.79ms step:660/1600 train_time:26279ms step_avg:39.82ms step:661/1600 train_time:26342ms step_avg:39.85ms step:662/1600 train_time:26401ms step_avg:39.88ms step:663/1600 train_time:26464ms step_avg:39.92ms step:664/1600 train_time:26523ms step_avg:39.94ms step:665/1600 train_time:26586ms step_avg:39.98ms step:666/1600 train_time:26645ms step_avg:40.01ms step:667/1600 train_time:26707ms step_avg:40.04ms step:668/1600 train_time:26766ms step_avg:40.07ms step:669/1600 train_time:26829ms step_avg:40.10ms step:670/1600 train_time:26888ms step_avg:40.13ms step:671/1600 train_time:26950ms step_avg:40.16ms step:672/1600 train_time:27010ms step_avg:40.19ms step:673/1600 train_time:27071ms step_avg:40.23ms step:674/1600 train_time:27130ms step_avg:40.25ms step:675/1600 train_time:27191ms step_avg:40.28ms step:676/1600 train_time:27250ms step_avg:40.31ms step:677/1600 train_time:27313ms step_avg:40.34ms step:678/1600 train_time:27371ms step_avg:40.37ms step:679/1600 train_time:27432ms step_avg:40.40ms step:680/1600 train_time:27491ms step_avg:40.43ms step:681/1600 train_time:27553ms step_avg:40.46ms step:682/1600 train_time:27612ms step_avg:40.49ms step:683/1600 train_time:27674ms step_avg:40.52ms step:684/1600 train_time:27732ms step_avg:40.54ms step:685/1600 train_time:27795ms step_avg:40.58ms step:686/1600 train_time:27854ms step_avg:40.60ms step:687/1600 train_time:27917ms step_avg:40.64ms step:688/1600 train_time:27976ms step_avg:40.66ms step:689/1600 train_time:28038ms step_avg:40.69ms step:690/1600 train_time:28097ms step_avg:40.72ms step:691/1600 train_time:28159ms step_avg:40.75ms step:692/1600 train_time:28219ms step_avg:40.78ms step:693/1600 train_time:28281ms step_avg:40.81ms step:694/1600 train_time:28341ms step_avg:40.84ms step:695/1600 train_time:28404ms step_avg:40.87ms step:696/1600 train_time:28463ms step_avg:40.90ms step:697/1600 train_time:28525ms step_avg:40.93ms step:698/1600 train_time:28584ms step_avg:40.95ms step:699/1600 train_time:28647ms step_avg:40.98ms step:700/1600 train_time:28706ms step_avg:41.01ms step:701/1600 train_time:28769ms step_avg:41.04ms step:702/1600 train_time:28828ms step_avg:41.07ms step:703/1600 train_time:28891ms step_avg:41.10ms step:704/1600 train_time:28950ms step_avg:41.12ms step:705/1600 train_time:29012ms step_avg:41.15ms step:706/1600 train_time:29071ms step_avg:41.18ms step:707/1600 train_time:29133ms step_avg:41.21ms step:708/1600 train_time:29192ms step_avg:41.23ms step:709/1600 train_time:29254ms step_avg:41.26ms step:710/1600 train_time:29312ms step_avg:41.28ms step:711/1600 train_time:29375ms step_avg:41.31ms step:712/1600 train_time:29433ms step_avg:41.34ms step:713/1600 train_time:29496ms step_avg:41.37ms step:714/1600 train_time:29553ms step_avg:41.39ms step:715/1600 train_time:29616ms step_avg:41.42ms step:716/1600 train_time:29675ms step_avg:41.45ms step:717/1600 train_time:29738ms step_avg:41.48ms step:718/1600 train_time:29798ms step_avg:41.50ms step:719/1600 train_time:29861ms step_avg:41.53ms step:720/1600 train_time:29920ms step_avg:41.56ms step:721/1600 train_time:29982ms step_avg:41.58ms step:722/1600 train_time:30041ms step_avg:41.61ms step:723/1600 train_time:30104ms step_avg:41.64ms step:724/1600 train_time:30163ms step_avg:41.66ms step:725/1600 train_time:30225ms step_avg:41.69ms step:726/1600 train_time:30284ms step_avg:41.71ms step:727/1600 train_time:30347ms step_avg:41.74ms step:728/1600 train_time:30406ms step_avg:41.77ms step:729/1600 train_time:30469ms step_avg:41.80ms step:730/1600 train_time:30527ms step_avg:41.82ms step:731/1600 train_time:30590ms step_avg:41.85ms step:732/1600 train_time:30648ms step_avg:41.87ms step:733/1600 train_time:30711ms step_avg:41.90ms step:734/1600 train_time:30770ms step_avg:41.92ms step:735/1600 train_time:30832ms step_avg:41.95ms step:736/1600 train_time:30891ms step_avg:41.97ms step:737/1600 train_time:30952ms step_avg:42.00ms step:738/1600 train_time:31013ms step_avg:42.02ms step:739/1600 train_time:31075ms step_avg:42.05ms step:740/1600 train_time:31133ms step_avg:42.07ms step:741/1600 train_time:31193ms step_avg:42.10ms step:742/1600 train_time:31251ms step_avg:42.12ms step:743/1600 train_time:31315ms step_avg:42.15ms step:744/1600 train_time:31374ms step_avg:42.17ms step:745/1600 train_time:31435ms step_avg:42.19ms step:746/1600 train_time:31494ms step_avg:42.22ms step:747/1600 train_time:31556ms step_avg:42.24ms step:748/1600 train_time:31615ms step_avg:42.27ms step:749/1600 train_time:31677ms step_avg:42.29ms step:750/1600 train_time:31737ms step_avg:42.32ms step:750/1600 val_loss:3.8946 train_time:31783ms step_avg:42.38ms step:751/1600 train_time:31804ms step_avg:42.35ms step:752/1600 train_time:31865ms step_avg:42.37ms step:753/1600 train_time:31929ms step_avg:42.40ms step:754/1600 train_time:31990ms step_avg:42.43ms step:755/1600 train_time:32052ms step_avg:42.45ms step:756/1600 train_time:32111ms step_avg:42.48ms step:757/1600 train_time:32173ms step_avg:42.50ms step:758/1600 train_time:32232ms step_avg:42.52ms step:759/1600 train_time:32294ms step_avg:42.55ms step:760/1600 train_time:32353ms step_avg:42.57ms step:761/1600 train_time:32415ms step_avg:42.59ms step:762/1600 train_time:32474ms step_avg:42.62ms step:763/1600 train_time:32535ms step_avg:42.64ms step:764/1600 train_time:32594ms step_avg:42.66ms step:765/1600 train_time:32654ms step_avg:42.69ms step:766/1600 train_time:32714ms step_avg:42.71ms step:767/1600 train_time:32778ms step_avg:42.74ms step:768/1600 train_time:32839ms step_avg:42.76ms step:769/1600 train_time:32902ms step_avg:42.79ms step:770/1600 train_time:32961ms step_avg:42.81ms step:771/1600 train_time:33023ms step_avg:42.83ms step:772/1600 train_time:33081ms step_avg:42.85ms step:773/1600 train_time:33144ms step_avg:42.88ms step:774/1600 train_time:33202ms step_avg:42.90ms step:775/1600 train_time:33265ms step_avg:42.92ms step:776/1600 train_time:33324ms step_avg:42.94ms step:777/1600 train_time:33386ms step_avg:42.97ms step:778/1600 train_time:33445ms step_avg:42.99ms step:779/1600 train_time:33507ms step_avg:43.01ms step:780/1600 train_time:33567ms step_avg:43.03ms step:781/1600 train_time:33629ms step_avg:43.06ms step:782/1600 train_time:33688ms step_avg:43.08ms step:783/1600 train_time:33751ms step_avg:43.10ms step:784/1600 train_time:33810ms step_avg:43.13ms step:785/1600 train_time:33875ms step_avg:43.15ms step:786/1600 train_time:33935ms step_avg:43.17ms step:787/1600 train_time:33998ms step_avg:43.20ms step:788/1600 train_time:34057ms step_avg:43.22ms step:789/1600 train_time:34119ms step_avg:43.24ms step:790/1600 train_time:34178ms step_avg:43.26ms step:791/1600 train_time:34240ms step_avg:43.29ms step:792/1600 train_time:34298ms step_avg:43.31ms step:793/1600 train_time:34361ms step_avg:43.33ms step:794/1600 train_time:34419ms step_avg:43.35ms step:795/1600 train_time:34482ms step_avg:43.37ms step:796/1600 train_time:34540ms step_avg:43.39ms step:797/1600 train_time:34602ms step_avg:43.42ms step:798/1600 train_time:34660ms step_avg:43.43ms step:799/1600 train_time:34722ms step_avg:43.46ms step:800/1600 train_time:34782ms step_avg:43.48ms step:801/1600 train_time:34845ms step_avg:43.50ms step:802/1600 train_time:34905ms step_avg:43.52ms step:803/1600 train_time:34968ms step_avg:43.55ms step:804/1600 train_time:35029ms step_avg:43.57ms step:805/1600 train_time:35091ms step_avg:43.59ms step:806/1600 train_time:35150ms step_avg:43.61ms step:807/1600 train_time:35211ms step_avg:43.63ms step:808/1600 train_time:35271ms step_avg:43.65ms step:809/1600 train_time:35335ms step_avg:43.68ms step:810/1600 train_time:35394ms step_avg:43.70ms step:811/1600 train_time:35456ms step_avg:43.72ms step:812/1600 train_time:35514ms step_avg:43.74ms step:813/1600 train_time:35578ms step_avg:43.76ms step:814/1600 train_time:35638ms step_avg:43.78ms step:815/1600 train_time:35699ms step_avg:43.80ms step:816/1600 train_time:35757ms step_avg:43.82ms step:817/1600 train_time:35820ms step_avg:43.84ms step:818/1600 train_time:35878ms step_avg:43.86ms step:819/1600 train_time:35940ms step_avg:43.88ms step:820/1600 train_time:36000ms step_avg:43.90ms step:821/1600 train_time:36061ms step_avg:43.92ms step:822/1600 train_time:36120ms step_avg:43.94ms step:823/1600 train_time:36182ms step_avg:43.96ms step:824/1600 train_time:36240ms step_avg:43.98ms step:825/1600 train_time:36303ms step_avg:44.00ms step:826/1600 train_time:36362ms step_avg:44.02ms step:827/1600 train_time:36425ms step_avg:44.04ms step:828/1600 train_time:36484ms step_avg:44.06ms step:829/1600 train_time:36547ms step_avg:44.09ms step:830/1600 train_time:36606ms step_avg:44.10ms step:831/1600 train_time:36669ms step_avg:44.13ms step:832/1600 train_time:36728ms step_avg:44.14ms step:833/1600 train_time:36791ms step_avg:44.17ms step:834/1600 train_time:36849ms step_avg:44.18ms step:835/1600 train_time:36913ms step_avg:44.21ms step:836/1600 train_time:36972ms step_avg:44.22ms step:837/1600 train_time:37035ms step_avg:44.25ms step:838/1600 train_time:37094ms step_avg:44.27ms step:839/1600 train_time:37156ms step_avg:44.29ms step:840/1600 train_time:37216ms step_avg:44.30ms step:841/1600 train_time:37279ms step_avg:44.33ms step:842/1600 train_time:37338ms step_avg:44.34ms step:843/1600 train_time:37399ms step_avg:44.36ms step:844/1600 train_time:37458ms step_avg:44.38ms step:845/1600 train_time:37520ms step_avg:44.40ms step:846/1600 train_time:37578ms step_avg:44.42ms step:847/1600 train_time:37643ms step_avg:44.44ms step:848/1600 train_time:37700ms step_avg:44.46ms step:849/1600 train_time:37761ms step_avg:44.48ms step:850/1600 train_time:37819ms step_avg:44.49ms step:851/1600 train_time:37881ms step_avg:44.51ms step:852/1600 train_time:37941ms step_avg:44.53ms step:853/1600 train_time:38005ms step_avg:44.55ms step:854/1600 train_time:38063ms step_avg:44.57ms step:855/1600 train_time:38126ms step_avg:44.59ms step:856/1600 train_time:38186ms step_avg:44.61ms step:857/1600 train_time:38248ms step_avg:44.63ms step:858/1600 train_time:38307ms step_avg:44.65ms step:859/1600 train_time:38370ms step_avg:44.67ms step:860/1600 train_time:38430ms step_avg:44.69ms step:861/1600 train_time:38491ms step_avg:44.71ms step:862/1600 train_time:38550ms step_avg:44.72ms step:863/1600 train_time:38612ms step_avg:44.74ms step:864/1600 train_time:38672ms step_avg:44.76ms step:865/1600 train_time:38734ms step_avg:44.78ms step:866/1600 train_time:38793ms step_avg:44.80ms step:867/1600 train_time:38856ms step_avg:44.82ms step:868/1600 train_time:38915ms step_avg:44.83ms step:869/1600 train_time:38978ms step_avg:44.85ms step:870/1600 train_time:39038ms step_avg:44.87ms step:871/1600 train_time:39100ms step_avg:44.89ms step:872/1600 train_time:39159ms step_avg:44.91ms step:873/1600 train_time:39222ms step_avg:44.93ms step:874/1600 train_time:39280ms step_avg:44.94ms step:875/1600 train_time:39342ms step_avg:44.96ms step:876/1600 train_time:39401ms step_avg:44.98ms step:877/1600 train_time:39462ms step_avg:45.00ms step:878/1600 train_time:39521ms step_avg:45.01ms step:879/1600 train_time:39584ms step_avg:45.03ms step:880/1600 train_time:39644ms step_avg:45.05ms step:881/1600 train_time:39707ms step_avg:45.07ms step:882/1600 train_time:39766ms step_avg:45.09ms step:883/1600 train_time:39828ms step_avg:45.11ms step:884/1600 train_time:39888ms step_avg:45.12ms step:885/1600 train_time:39951ms step_avg:45.14ms step:886/1600 train_time:40010ms step_avg:45.16ms step:887/1600 train_time:40072ms step_avg:45.18ms step:888/1600 train_time:40132ms step_avg:45.19ms step:889/1600 train_time:40195ms step_avg:45.21ms step:890/1600 train_time:40259ms step_avg:45.23ms step:891/1600 train_time:40318ms step_avg:45.25ms step:892/1600 train_time:40377ms step_avg:45.27ms step:893/1600 train_time:40440ms step_avg:45.29ms step:894/1600 train_time:40498ms step_avg:45.30ms step:895/1600 train_time:40560ms step_avg:45.32ms step:896/1600 train_time:40618ms step_avg:45.33ms step:897/1600 train_time:40682ms step_avg:45.35ms step:898/1600 train_time:40741ms step_avg:45.37ms step:899/1600 train_time:40803ms step_avg:45.39ms step:900/1600 train_time:40861ms step_avg:45.40ms step:901/1600 train_time:40923ms step_avg:45.42ms step:902/1600 train_time:40981ms step_avg:45.43ms step:903/1600 train_time:41044ms step_avg:45.45ms step:904/1600 train_time:41103ms step_avg:45.47ms step:905/1600 train_time:41165ms step_avg:45.49ms step:906/1600 train_time:41225ms step_avg:45.50ms step:907/1600 train_time:41288ms step_avg:45.52ms step:908/1600 train_time:41348ms step_avg:45.54ms step:909/1600 train_time:41410ms step_avg:45.56ms step:910/1600 train_time:41469ms step_avg:45.57ms step:911/1600 train_time:41531ms step_avg:45.59ms step:912/1600 train_time:41590ms step_avg:45.60ms step:913/1600 train_time:41654ms step_avg:45.62ms step:914/1600 train_time:41713ms step_avg:45.64ms step:915/1600 train_time:41776ms step_avg:45.66ms step:916/1600 train_time:41835ms step_avg:45.67ms step:917/1600 train_time:41898ms step_avg:45.69ms step:918/1600 train_time:41957ms step_avg:45.71ms step:919/1600 train_time:42020ms step_avg:45.72ms step:920/1600 train_time:42080ms step_avg:45.74ms step:921/1600 train_time:42141ms step_avg:45.76ms step:922/1600 train_time:42199ms step_avg:45.77ms step:923/1600 train_time:42260ms step_avg:45.79ms step:924/1600 train_time:42320ms step_avg:45.80ms step:925/1600 train_time:42382ms step_avg:45.82ms step:926/1600 train_time:42440ms step_avg:45.83ms step:927/1600 train_time:42503ms step_avg:45.85ms step:928/1600 train_time:42562ms step_avg:45.86ms step:929/1600 train_time:42624ms step_avg:45.88ms step:930/1600 train_time:42682ms step_avg:45.90ms step:931/1600 train_time:42746ms step_avg:45.91ms step:932/1600 train_time:42805ms step_avg:45.93ms step:933/1600 train_time:42868ms step_avg:45.95ms step:934/1600 train_time:42927ms step_avg:45.96ms step:935/1600 train_time:42990ms step_avg:45.98ms step:936/1600 train_time:43048ms step_avg:45.99ms step:937/1600 train_time:43111ms step_avg:46.01ms step:938/1600 train_time:43170ms step_avg:46.02ms step:939/1600 train_time:43233ms step_avg:46.04ms step:940/1600 train_time:43292ms step_avg:46.06ms step:941/1600 train_time:43354ms step_avg:46.07ms step:942/1600 train_time:43414ms step_avg:46.09ms step:943/1600 train_time:43477ms step_avg:46.10ms step:944/1600 train_time:43535ms step_avg:46.12ms step:945/1600 train_time:43597ms step_avg:46.13ms step:946/1600 train_time:43658ms step_avg:46.15ms step:947/1600 train_time:43720ms step_avg:46.17ms step:948/1600 train_time:43778ms step_avg:46.18ms step:949/1600 train_time:43841ms step_avg:46.20ms step:950/1600 train_time:43899ms step_avg:46.21ms step:951/1600 train_time:43961ms step_avg:46.23ms step:952/1600 train_time:44019ms step_avg:46.24ms step:953/1600 train_time:44081ms step_avg:46.25ms step:954/1600 train_time:44140ms step_avg:46.27ms step:955/1600 train_time:44201ms step_avg:46.28ms step:956/1600 train_time:44260ms step_avg:46.30ms step:957/1600 train_time:44322ms step_avg:46.31ms step:958/1600 train_time:44381ms step_avg:46.33ms step:959/1600 train_time:44444ms step_avg:46.34ms step:960/1600 train_time:44503ms step_avg:46.36ms step:961/1600 train_time:44566ms step_avg:46.38ms step:962/1600 train_time:44627ms step_avg:46.39ms step:963/1600 train_time:44690ms step_avg:46.41ms step:964/1600 train_time:44748ms step_avg:46.42ms step:965/1600 train_time:44811ms step_avg:46.44ms step:966/1600 train_time:44870ms step_avg:46.45ms step:967/1600 train_time:44934ms step_avg:46.47ms step:968/1600 train_time:44992ms step_avg:46.48ms step:969/1600 train_time:45055ms step_avg:46.50ms step:970/1600 train_time:45114ms step_avg:46.51ms step:971/1600 train_time:45177ms step_avg:46.53ms step:972/1600 train_time:45238ms step_avg:46.54ms step:973/1600 train_time:45299ms step_avg:46.56ms step:974/1600 train_time:45357ms step_avg:46.57ms step:975/1600 train_time:45419ms step_avg:46.58ms step:976/1600 train_time:45478ms step_avg:46.60ms step:977/1600 train_time:45540ms step_avg:46.61ms step:978/1600 train_time:45599ms step_avg:46.62ms step:979/1600 train_time:45662ms step_avg:46.64ms step:980/1600 train_time:45719ms step_avg:46.65ms step:981/1600 train_time:45782ms step_avg:46.67ms step:982/1600 train_time:45842ms step_avg:46.68ms step:983/1600 train_time:45904ms step_avg:46.70ms step:984/1600 train_time:45961ms step_avg:46.71ms step:985/1600 train_time:46023ms step_avg:46.72ms step:986/1600 train_time:46081ms step_avg:46.74ms step:987/1600 train_time:46143ms step_avg:46.75ms step:988/1600 train_time:46203ms step_avg:46.76ms step:989/1600 train_time:46266ms step_avg:46.78ms step:990/1600 train_time:46325ms step_avg:46.79ms step:991/1600 train_time:46388ms step_avg:46.81ms step:992/1600 train_time:46447ms step_avg:46.82ms step:993/1600 train_time:46509ms step_avg:46.84ms step:994/1600 train_time:46568ms step_avg:46.85ms step:995/1600 train_time:46631ms step_avg:46.87ms step:996/1600 train_time:46690ms step_avg:46.88ms step:997/1600 train_time:46753ms step_avg:46.89ms step:998/1600 train_time:46812ms step_avg:46.91ms step:999/1600 train_time:46875ms step_avg:46.92ms step:1000/1600 train_time:46934ms step_avg:46.93ms step:1000/1600 val_loss:3.5959 train_time:46980ms step_avg:46.98ms step:1001/1600 train_time:47001ms step_avg:46.95ms step:1002/1600 train_time:47059ms step_avg:46.97ms step:1003/1600 train_time:47124ms step_avg:46.98ms step:1004/1600 train_time:47185ms step_avg:47.00ms step:1005/1600 train_time:47246ms step_avg:47.01ms step:1006/1600 train_time:47306ms step_avg:47.02ms step:1007/1600 train_time:47367ms step_avg:47.04ms step:1008/1600 train_time:47425ms step_avg:47.05ms step:1009/1600 train_time:47486ms step_avg:47.06ms step:1010/1600 train_time:47544ms step_avg:47.07ms step:1011/1600 train_time:47606ms step_avg:47.09ms step:1012/1600 train_time:47664ms step_avg:47.10ms step:1013/1600 train_time:47725ms step_avg:47.11ms step:1014/1600 train_time:47785ms step_avg:47.13ms step:1015/1600 train_time:47845ms step_avg:47.14ms step:1016/1600 train_time:47903ms step_avg:47.15ms step:1017/1600 train_time:47966ms step_avg:47.16ms step:1018/1600 train_time:48026ms step_avg:47.18ms step:1019/1600 train_time:48088ms step_avg:47.19ms step:1020/1600 train_time:48148ms step_avg:47.20ms step:1021/1600 train_time:48211ms step_avg:47.22ms step:1022/1600 train_time:48270ms step_avg:47.23ms step:1023/1600 train_time:48333ms step_avg:47.25ms step:1024/1600 train_time:48392ms step_avg:47.26ms step:1025/1600 train_time:48455ms step_avg:47.27ms step:1026/1600 train_time:48514ms step_avg:47.28ms step:1027/1600 train_time:48578ms step_avg:47.30ms step:1028/1600 train_time:48636ms step_avg:47.31ms step:1029/1600 train_time:48698ms step_avg:47.33ms step:1030/1600 train_time:48757ms step_avg:47.34ms step:1031/1600 train_time:48818ms step_avg:47.35ms step:1032/1600 train_time:48877ms step_avg:47.36ms step:1033/1600 train_time:48940ms step_avg:47.38ms step:1034/1600 train_time:49000ms step_avg:47.39ms step:1035/1600 train_time:49064ms step_avg:47.40ms step:1036/1600 train_time:49123ms step_avg:47.42ms step:1037/1600 train_time:49188ms step_avg:47.43ms step:1038/1600 train_time:49246ms step_avg:47.44ms step:1039/1600 train_time:49308ms step_avg:47.46ms step:1040/1600 train_time:49367ms step_avg:47.47ms step:1041/1600 train_time:49436ms step_avg:47.49ms step:1042/1600 train_time:49519ms step_avg:47.52ms step:1043/1600 train_time:49607ms step_avg:47.56ms step:1044/1600 train_time:49691ms step_avg:47.60ms step:1045/1600 train_time:49779ms step_avg:47.64ms step:1046/1600 train_time:49864ms step_avg:47.67ms step:1047/1600 train_time:49953ms step_avg:47.71ms step:1048/1600 train_time:50040ms step_avg:47.75ms step:1049/1600 train_time:50131ms step_avg:47.79ms step:1050/1600 train_time:50215ms step_avg:47.82ms step:1051/1600 train_time:50304ms step_avg:47.86ms step:1052/1600 train_time:50391ms step_avg:47.90ms step:1053/1600 train_time:50479ms step_avg:47.94ms step:1054/1600 train_time:50564ms step_avg:47.97ms step:1055/1600 train_time:50652ms step_avg:48.01ms step:1056/1600 train_time:50736ms step_avg:48.05ms step:1057/1600 train_time:50825ms step_avg:48.08ms step:1058/1600 train_time:50909ms step_avg:48.12ms step:1059/1600 train_time:50997ms step_avg:48.16ms step:1060/1600 train_time:51083ms step_avg:48.19ms step:1061/1600 train_time:51173ms step_avg:48.23ms step:1062/1600 train_time:51259ms step_avg:48.27ms step:1063/1600 train_time:51347ms step_avg:48.30ms step:1064/1600 train_time:51433ms step_avg:48.34ms step:1065/1600 train_time:51520ms step_avg:48.38ms step:1066/1600 train_time:51605ms step_avg:48.41ms step:1067/1600 train_time:51693ms step_avg:48.45ms step:1068/1600 train_time:51778ms step_avg:48.48ms step:1069/1600 train_time:51866ms step_avg:48.52ms step:1070/1600 train_time:51951ms step_avg:48.55ms step:1071/1600 train_time:52040ms step_avg:48.59ms step:1072/1600 train_time:52126ms step_avg:48.63ms step:1073/1600 train_time:52215ms step_avg:48.66ms step:1074/1600 train_time:52301ms step_avg:48.70ms step:1075/1600 train_time:52389ms step_avg:48.73ms step:1076/1600 train_time:52473ms step_avg:48.77ms step:1077/1600 train_time:52561ms step_avg:48.80ms step:1078/1600 train_time:52646ms step_avg:48.84ms step:1079/1600 train_time:52734ms step_avg:48.87ms step:1080/1600 train_time:52819ms step_avg:48.91ms step:1081/1600 train_time:52908ms step_avg:48.94ms step:1082/1600 train_time:52993ms step_avg:48.98ms step:1083/1600 train_time:53082ms step_avg:49.01ms step:1084/1600 train_time:53167ms step_avg:49.05ms step:1085/1600 train_time:53256ms step_avg:49.08ms step:1086/1600 train_time:53343ms step_avg:49.12ms step:1087/1600 train_time:53432ms step_avg:49.16ms step:1088/1600 train_time:53516ms step_avg:49.19ms step:1089/1600 train_time:53603ms step_avg:49.22ms step:1090/1600 train_time:53689ms step_avg:49.26ms step:1091/1600 train_time:53776ms step_avg:49.29ms step:1092/1600 train_time:53862ms step_avg:49.32ms step:1093/1600 train_time:53951ms step_avg:49.36ms step:1094/1600 train_time:54035ms step_avg:49.39ms step:1095/1600 train_time:54124ms step_avg:49.43ms step:1096/1600 train_time:54209ms step_avg:49.46ms step:1097/1600 train_time:54300ms step_avg:49.50ms step:1098/1600 train_time:54384ms step_avg:49.53ms step:1099/1600 train_time:54472ms step_avg:49.56ms step:1100/1600 train_time:54556ms step_avg:49.60ms step:1101/1600 train_time:54644ms step_avg:49.63ms step:1102/1600 train_time:54729ms step_avg:49.66ms step:1103/1600 train_time:54817ms step_avg:49.70ms step:1104/1600 train_time:54902ms step_avg:49.73ms step:1105/1600 train_time:54991ms step_avg:49.77ms step:1106/1600 train_time:55076ms step_avg:49.80ms step:1107/1600 train_time:55165ms step_avg:49.83ms step:1108/1600 train_time:55252ms step_avg:49.87ms step:1109/1600 train_time:55338ms step_avg:49.90ms step:1110/1600 train_time:55423ms step_avg:49.93ms step:1111/1600 train_time:55512ms step_avg:49.97ms step:1112/1600 train_time:55596ms step_avg:50.00ms step:1113/1600 train_time:55684ms step_avg:50.03ms step:1114/1600 train_time:55768ms step_avg:50.06ms step:1115/1600 train_time:55858ms step_avg:50.10ms step:1116/1600 train_time:55944ms step_avg:50.13ms step:1117/1600 train_time:56032ms step_avg:50.16ms step:1118/1600 train_time:56118ms step_avg:50.20ms step:1119/1600 train_time:56207ms step_avg:50.23ms step:1120/1600 train_time:56292ms step_avg:50.26ms step:1121/1600 train_time:56379ms step_avg:50.29ms step:1122/1600 train_time:56465ms step_avg:50.33ms step:1123/1600 train_time:56554ms step_avg:50.36ms step:1124/1600 train_time:56638ms step_avg:50.39ms step:1125/1600 train_time:56728ms step_avg:50.42ms step:1126/1600 train_time:56812ms step_avg:50.45ms step:1127/1600 train_time:56900ms step_avg:50.49ms step:1128/1600 train_time:56986ms step_avg:50.52ms step:1129/1600 train_time:57074ms step_avg:50.55ms step:1130/1600 train_time:57159ms step_avg:50.58ms step:1131/1600 train_time:57247ms step_avg:50.62ms step:1132/1600 train_time:57333ms step_avg:50.65ms step:1133/1600 train_time:57420ms step_avg:50.68ms step:1134/1600 train_time:57506ms step_avg:50.71ms step:1135/1600 train_time:57595ms step_avg:50.74ms step:1136/1600 train_time:57684ms step_avg:50.78ms step:1137/1600 train_time:57770ms step_avg:50.81ms step:1138/1600 train_time:57856ms step_avg:50.84ms step:1139/1600 train_time:57943ms step_avg:50.87ms step:1140/1600 train_time:58027ms step_avg:50.90ms step:1141/1600 train_time:58115ms step_avg:50.93ms step:1142/1600 train_time:58200ms step_avg:50.96ms step:1143/1600 train_time:58288ms step_avg:51.00ms step:1144/1600 train_time:58373ms step_avg:51.03ms step:1145/1600 train_time:58461ms step_avg:51.06ms step:1146/1600 train_time:58546ms step_avg:51.09ms step:1147/1600 train_time:58635ms step_avg:51.12ms step:1148/1600 train_time:58721ms step_avg:51.15ms step:1149/1600 train_time:58810ms step_avg:51.18ms step:1150/1600 train_time:58894ms step_avg:51.21ms step:1151/1600 train_time:58983ms step_avg:51.25ms step:1152/1600 train_time:59068ms step_avg:51.27ms step:1153/1600 train_time:59157ms step_avg:51.31ms step:1154/1600 train_time:59247ms step_avg:51.34ms step:1155/1600 train_time:59334ms step_avg:51.37ms step:1156/1600 train_time:59417ms step_avg:51.40ms step:1157/1600 train_time:59505ms step_avg:51.43ms step:1158/1600 train_time:59590ms step_avg:51.46ms step:1159/1600 train_time:59677ms step_avg:51.49ms step:1160/1600 train_time:59762ms step_avg:51.52ms step:1161/1600 train_time:59851ms step_avg:51.55ms step:1162/1600 train_time:59935ms step_avg:51.58ms step:1163/1600 train_time:60024ms step_avg:51.61ms step:1164/1600 train_time:60109ms step_avg:51.64ms step:1165/1600 train_time:60198ms step_avg:51.67ms step:1166/1600 train_time:60284ms step_avg:51.70ms step:1167/1600 train_time:60373ms step_avg:51.73ms step:1168/1600 train_time:60457ms step_avg:51.76ms step:1169/1600 train_time:60545ms step_avg:51.79ms step:1170/1600 train_time:60631ms step_avg:51.82ms step:1171/1600 train_time:60718ms step_avg:51.85ms step:1172/1600 train_time:60805ms step_avg:51.88ms step:1173/1600 train_time:60893ms step_avg:51.91ms step:1174/1600 train_time:60979ms step_avg:51.94ms step:1175/1600 train_time:61067ms step_avg:51.97ms step:1176/1600 train_time:61152ms step_avg:52.00ms step:1177/1600 train_time:61239ms step_avg:52.03ms step:1178/1600 train_time:61324ms step_avg:52.06ms step:1179/1600 train_time:61413ms step_avg:52.09ms step:1180/1600 train_time:61498ms step_avg:52.12ms step:1181/1600 train_time:61587ms step_avg:52.15ms step:1182/1600 train_time:61671ms step_avg:52.18ms step:1183/1600 train_time:61760ms step_avg:52.21ms step:1184/1600 train_time:61845ms step_avg:52.23ms step:1185/1600 train_time:61934ms step_avg:52.26ms step:1186/1600 train_time:62019ms step_avg:52.29ms step:1187/1600 train_time:62108ms step_avg:52.32ms step:1188/1600 train_time:62193ms step_avg:52.35ms step:1189/1600 train_time:62281ms step_avg:52.38ms step:1190/1600 train_time:62366ms step_avg:52.41ms step:1191/1600 train_time:62455ms step_avg:52.44ms step:1192/1600 train_time:62540ms step_avg:52.47ms step:1193/1600 train_time:62628ms step_avg:52.50ms step:1194/1600 train_time:62713ms step_avg:52.52ms step:1195/1600 train_time:62801ms step_avg:52.55ms step:1196/1600 train_time:62888ms step_avg:52.58ms step:1197/1600 train_time:62976ms step_avg:52.61ms step:1198/1600 train_time:63061ms step_avg:52.64ms step:1199/1600 train_time:63151ms step_avg:52.67ms step:1200/1600 train_time:63236ms step_avg:52.70ms step:1201/1600 train_time:63324ms step_avg:52.73ms step:1202/1600 train_time:63409ms step_avg:52.75ms step:1203/1600 train_time:63498ms step_avg:52.78ms step:1204/1600 train_time:63583ms step_avg:52.81ms step:1205/1600 train_time:63671ms step_avg:52.84ms step:1206/1600 train_time:63756ms step_avg:52.87ms step:1207/1600 train_time:63844ms step_avg:52.89ms step:1208/1600 train_time:63930ms step_avg:52.92ms step:1209/1600 train_time:64018ms step_avg:52.95ms step:1210/1600 train_time:64103ms step_avg:52.98ms step:1211/1600 train_time:64194ms step_avg:53.01ms step:1212/1600 train_time:64278ms step_avg:53.03ms step:1213/1600 train_time:64366ms step_avg:53.06ms step:1214/1600 train_time:64451ms step_avg:53.09ms step:1215/1600 train_time:64539ms step_avg:53.12ms step:1216/1600 train_time:64624ms step_avg:53.14ms step:1217/1600 train_time:64712ms step_avg:53.17ms step:1218/1600 train_time:64797ms step_avg:53.20ms step:1219/1600 train_time:64885ms step_avg:53.23ms step:1220/1600 train_time:64970ms step_avg:53.25ms step:1221/1600 train_time:65057ms step_avg:53.28ms step:1222/1600 train_time:65143ms step_avg:53.31ms step:1223/1600 train_time:65231ms step_avg:53.34ms step:1224/1600 train_time:65317ms step_avg:53.36ms step:1225/1600 train_time:65406ms step_avg:53.39ms step:1226/1600 train_time:65492ms step_avg:53.42ms step:1227/1600 train_time:65579ms step_avg:53.45ms step:1228/1600 train_time:65664ms step_avg:53.47ms step:1229/1600 train_time:65752ms step_avg:53.50ms step:1230/1600 train_time:65837ms step_avg:53.53ms step:1231/1600 train_time:65925ms step_avg:53.55ms step:1232/1600 train_time:66010ms step_avg:53.58ms step:1233/1600 train_time:66099ms step_avg:53.61ms step:1234/1600 train_time:66184ms step_avg:53.63ms step:1235/1600 train_time:66273ms step_avg:53.66ms step:1236/1600 train_time:66358ms step_avg:53.69ms step:1237/1600 train_time:66446ms step_avg:53.72ms step:1238/1600 train_time:66531ms step_avg:53.74ms step:1239/1600 train_time:66619ms step_avg:53.77ms step:1240/1600 train_time:66704ms step_avg:53.79ms step:1241/1600 train_time:66792ms step_avg:53.82ms step:1242/1600 train_time:66877ms step_avg:53.85ms step:1243/1600 train_time:66965ms step_avg:53.87ms step:1244/1600 train_time:67050ms step_avg:53.90ms step:1245/1600 train_time:67139ms step_avg:53.93ms step:1246/1600 train_time:67225ms step_avg:53.95ms step:1247/1600 train_time:67313ms step_avg:53.98ms step:1248/1600 train_time:67397ms step_avg:54.00ms step:1249/1600 train_time:67486ms step_avg:54.03ms step:1250/1600 train_time:67571ms step_avg:54.06ms step:1250/1600 val_loss:3.4169 train_time:67642ms step_avg:54.11ms step:1251/1600 train_time:67663ms step_avg:54.09ms step:1252/1600 train_time:67751ms step_avg:54.11ms step:1253/1600 train_time:67845ms step_avg:54.15ms step:1254/1600 train_time:67931ms step_avg:54.17ms step:1255/1600 train_time:68017ms step_avg:54.20ms step:1256/1600 train_time:68102ms step_avg:54.22ms step:1257/1600 train_time:68189ms step_avg:54.25ms step:1258/1600 train_time:68273ms step_avg:54.27ms step:1259/1600 train_time:68360ms step_avg:54.30ms step:1260/1600 train_time:68446ms step_avg:54.32ms step:1261/1600 train_time:68532ms step_avg:54.35ms step:1262/1600 train_time:68617ms step_avg:54.37ms step:1263/1600 train_time:68708ms step_avg:54.40ms step:1264/1600 train_time:68795ms step_avg:54.43ms step:1265/1600 train_time:68885ms step_avg:54.45ms step:1266/1600 train_time:68972ms step_avg:54.48ms step:1267/1600 train_time:69059ms step_avg:54.51ms step:1268/1600 train_time:69145ms step_avg:54.53ms step:1269/1600 train_time:69234ms step_avg:54.56ms step:1270/1600 train_time:69318ms step_avg:54.58ms step:1271/1600 train_time:69405ms step_avg:54.61ms step:1272/1600 train_time:69491ms step_avg:54.63ms step:1273/1600 train_time:69578ms step_avg:54.66ms step:1274/1600 train_time:69667ms step_avg:54.68ms step:1275/1600 train_time:69757ms step_avg:54.71ms step:1276/1600 train_time:69842ms step_avg:54.73ms step:1277/1600 train_time:69931ms step_avg:54.76ms step:1278/1600 train_time:70016ms step_avg:54.79ms step:1279/1600 train_time:70104ms step_avg:54.81ms step:1280/1600 train_time:70188ms step_avg:54.83ms step:1281/1600 train_time:70276ms step_avg:54.86ms step:1282/1600 train_time:70360ms step_avg:54.88ms step:1283/1600 train_time:70449ms step_avg:54.91ms step:1284/1600 train_time:70533ms step_avg:54.93ms step:1285/1600 train_time:70621ms step_avg:54.96ms step:1286/1600 train_time:70708ms step_avg:54.98ms step:1287/1600 train_time:70797ms step_avg:55.01ms step:1288/1600 train_time:70883ms step_avg:55.03ms step:1289/1600 train_time:70972ms step_avg:55.06ms step:1290/1600 train_time:71058ms step_avg:55.08ms step:1291/1600 train_time:71145ms step_avg:55.11ms step:1292/1600 train_time:71230ms step_avg:55.13ms step:1293/1600 train_time:71317ms step_avg:55.16ms step:1294/1600 train_time:71402ms step_avg:55.18ms step:1295/1600 train_time:71490ms step_avg:55.20ms step:1296/1600 train_time:71574ms step_avg:55.23ms step:1297/1600 train_time:71664ms step_avg:55.25ms step:1298/1600 train_time:71750ms step_avg:55.28ms step:1299/1600 train_time:71839ms step_avg:55.30ms step:1300/1600 train_time:71924ms step_avg:55.33ms step:1301/1600 train_time:72013ms step_avg:55.35ms step:1302/1600 train_time:72098ms step_avg:55.38ms step:1303/1600 train_time:72186ms step_avg:55.40ms step:1304/1600 train_time:72272ms step_avg:55.42ms step:1305/1600 train_time:72359ms step_avg:55.45ms step:1306/1600 train_time:72446ms step_avg:55.47ms step:1307/1600 train_time:72532ms step_avg:55.50ms step:1308/1600 train_time:72619ms step_avg:55.52ms step:1309/1600 train_time:72707ms step_avg:55.54ms step:1310/1600 train_time:72792ms step_avg:55.57ms step:1311/1600 train_time:72881ms step_avg:55.59ms step:1312/1600 train_time:72966ms step_avg:55.61ms step:1313/1600 train_time:73055ms step_avg:55.64ms step:1314/1600 train_time:73140ms step_avg:55.66ms step:1315/1600 train_time:73229ms step_avg:55.69ms step:1316/1600 train_time:73313ms step_avg:55.71ms step:1317/1600 train_time:73402ms step_avg:55.73ms step:1318/1600 train_time:73486ms step_avg:55.76ms step:1319/1600 train_time:73574ms step_avg:55.78ms step:1320/1600 train_time:73659ms step_avg:55.80ms step:1321/1600 train_time:73748ms step_avg:55.83ms step:1322/1600 train_time:73834ms step_avg:55.85ms step:1323/1600 train_time:73922ms step_avg:55.87ms step:1324/1600 train_time:74009ms step_avg:55.90ms step:1325/1600 train_time:74095ms step_avg:55.92ms step:1326/1600 train_time:74180ms step_avg:55.94ms step:1327/1600 train_time:74269ms step_avg:55.97ms step:1328/1600 train_time:74353ms step_avg:55.99ms step:1329/1600 train_time:74442ms step_avg:56.01ms step:1330/1600 train_time:74526ms step_avg:56.03ms step:1331/1600 train_time:74615ms step_avg:56.06ms step:1332/1600 train_time:74699ms step_avg:56.08ms step:1333/1600 train_time:74790ms step_avg:56.11ms step:1334/1600 train_time:74874ms step_avg:56.13ms step:1335/1600 train_time:74962ms step_avg:56.15ms step:1336/1600 train_time:75047ms step_avg:56.17ms step:1337/1600 train_time:75135ms step_avg:56.20ms step:1338/1600 train_time:75220ms step_avg:56.22ms step:1339/1600 train_time:75308ms step_avg:56.24ms step:1340/1600 train_time:75392ms step_avg:56.26ms step:1341/1600 train_time:75481ms step_avg:56.29ms step:1342/1600 train_time:75566ms step_avg:56.31ms step:1343/1600 train_time:75655ms step_avg:56.33ms step:1344/1600 train_time:75740ms step_avg:56.35ms step:1345/1600 train_time:75830ms step_avg:56.38ms step:1346/1600 train_time:75914ms step_avg:56.40ms step:1347/1600 train_time:76002ms step_avg:56.42ms step:1348/1600 train_time:76087ms step_avg:56.44ms step:1349/1600 train_time:76174ms step_avg:56.47ms step:1350/1600 train_time:76259ms step_avg:56.49ms step:1351/1600 train_time:76347ms step_avg:56.51ms step:1352/1600 train_time:76432ms step_avg:56.53ms step:1353/1600 train_time:76520ms step_avg:56.56ms step:1354/1600 train_time:76606ms step_avg:56.58ms step:1355/1600 train_time:76695ms step_avg:56.60ms step:1356/1600 train_time:76780ms step_avg:56.62ms step:1357/1600 train_time:76870ms step_avg:56.65ms step:1358/1600 train_time:76955ms step_avg:56.67ms step:1359/1600 train_time:77043ms step_avg:56.69ms step:1360/1600 train_time:77128ms step_avg:56.71ms step:1361/1600 train_time:77217ms step_avg:56.74ms step:1362/1600 train_time:77302ms step_avg:56.76ms step:1363/1600 train_time:77390ms step_avg:56.78ms step:1364/1600 train_time:77475ms step_avg:56.80ms step:1365/1600 train_time:77562ms step_avg:56.82ms step:1366/1600 train_time:77652ms step_avg:56.85ms step:1367/1600 train_time:77739ms step_avg:56.87ms step:1368/1600 train_time:77824ms step_avg:56.89ms step:1369/1600 train_time:77912ms step_avg:56.91ms step:1370/1600 train_time:77997ms step_avg:56.93ms step:1371/1600 train_time:78085ms step_avg:56.95ms step:1372/1600 train_time:78170ms step_avg:56.98ms step:1373/1600 train_time:78261ms step_avg:57.00ms step:1374/1600 train_time:78345ms step_avg:57.02ms step:1375/1600 train_time:78433ms step_avg:57.04ms step:1376/1600 train_time:78517ms step_avg:57.06ms step:1377/1600 train_time:78606ms step_avg:57.08ms step:1378/1600 train_time:78691ms step_avg:57.11ms step:1379/1600 train_time:78778ms step_avg:57.13ms step:1380/1600 train_time:78865ms step_avg:57.15ms step:1381/1600 train_time:78954ms step_avg:57.17ms step:1382/1600 train_time:79038ms step_avg:57.19ms step:1383/1600 train_time:79127ms step_avg:57.21ms step:1384/1600 train_time:79212ms step_avg:57.23ms step:1385/1600 train_time:79300ms step_avg:57.26ms step:1386/1600 train_time:79385ms step_avg:57.28ms step:1387/1600 train_time:79475ms step_avg:57.30ms step:1388/1600 train_time:79561ms step_avg:57.32ms step:1389/1600 train_time:79649ms step_avg:57.34ms step:1390/1600 train_time:79734ms step_avg:57.36ms step:1391/1600 train_time:79904ms step_avg:57.44ms step:1392/1600 train_time:79937ms step_avg:57.43ms step:1393/1600 train_time:80011ms step_avg:57.44ms step:1394/1600 train_time:80094ms step_avg:57.46ms step:1395/1600 train_time:80180ms step_avg:57.48ms step:1396/1600 train_time:80266ms step_avg:57.50ms step:1397/1600 train_time:80354ms step_avg:57.52ms step:1398/1600 train_time:80438ms step_avg:57.54ms step:1399/1600 train_time:80528ms step_avg:57.56ms step:1400/1600 train_time:80613ms step_avg:57.58ms step:1401/1600 train_time:80701ms step_avg:57.60ms step:1402/1600 train_time:80787ms step_avg:57.62ms step:1403/1600 train_time:80877ms step_avg:57.65ms step:1404/1600 train_time:80963ms step_avg:57.67ms step:1405/1600 train_time:81052ms step_avg:57.69ms step:1406/1600 train_time:81136ms step_avg:57.71ms step:1407/1600 train_time:81224ms step_avg:57.73ms step:1408/1600 train_time:81309ms step_avg:57.75ms step:1409/1600 train_time:81396ms step_avg:57.77ms step:1410/1600 train_time:81482ms step_avg:57.79ms step:1411/1600 train_time:81570ms step_avg:57.81ms step:1412/1600 train_time:81655ms step_avg:57.83ms step:1413/1600 train_time:81744ms step_avg:57.85ms step:1414/1600 train_time:81829ms step_avg:57.87ms step:1415/1600 train_time:81918ms step_avg:57.89ms step:1416/1600 train_time:82004ms step_avg:57.91ms step:1417/1600 train_time:82092ms step_avg:57.93ms step:1418/1600 train_time:82177ms step_avg:57.95ms step:1419/1600 train_time:82265ms step_avg:57.97ms step:1420/1600 train_time:82350ms step_avg:57.99ms step:1421/1600 train_time:82439ms step_avg:58.01ms step:1422/1600 train_time:82525ms step_avg:58.03ms step:1423/1600 train_time:82613ms step_avg:58.06ms step:1424/1600 train_time:82698ms step_avg:58.07ms step:1425/1600 train_time:82787ms step_avg:58.10ms step:1426/1600 train_time:82872ms step_avg:58.12ms step:1427/1600 train_time:82961ms step_avg:58.14ms step:1428/1600 train_time:83046ms step_avg:58.16ms step:1429/1600 train_time:83135ms step_avg:58.18ms step:1430/1600 train_time:83219ms step_avg:58.19ms step:1431/1600 train_time:83307ms step_avg:58.22ms step:1432/1600 train_time:83392ms step_avg:58.23ms step:1433/1600 train_time:83479ms step_avg:58.25ms step:1434/1600 train_time:83564ms step_avg:58.27ms step:1435/1600 train_time:83653ms step_avg:58.29ms step:1436/1600 train_time:83738ms step_avg:58.31ms step:1437/1600 train_time:83827ms step_avg:58.33ms step:1438/1600 train_time:83912ms step_avg:58.35ms step:1439/1600 train_time:84000ms step_avg:58.37ms step:1440/1600 train_time:84085ms step_avg:58.39ms step:1441/1600 train_time:84173ms step_avg:58.41ms step:1442/1600 train_time:84259ms step_avg:58.43ms step:1443/1600 train_time:84348ms step_avg:58.45ms step:1444/1600 train_time:84433ms step_avg:58.47ms step:1445/1600 train_time:84520ms step_avg:58.49ms step:1446/1600 train_time:84606ms step_avg:58.51ms step:1447/1600 train_time:84695ms step_avg:58.53ms step:1448/1600 train_time:84780ms step_avg:58.55ms step:1449/1600 train_time:84868ms step_avg:58.57ms step:1450/1600 train_time:84954ms step_avg:58.59ms step:1451/1600 train_time:85043ms step_avg:58.61ms step:1452/1600 train_time:85126ms step_avg:58.63ms step:1453/1600 train_time:85215ms step_avg:58.65ms step:1454/1600 train_time:85301ms step_avg:58.67ms step:1455/1600 train_time:85389ms step_avg:58.69ms step:1456/1600 train_time:85474ms step_avg:58.70ms step:1457/1600 train_time:85562ms step_avg:58.73ms step:1458/1600 train_time:85647ms step_avg:58.74ms step:1459/1600 train_time:85736ms step_avg:58.76ms step:1460/1600 train_time:85821ms step_avg:58.78ms step:1461/1600 train_time:85910ms step_avg:58.80ms step:1462/1600 train_time:85995ms step_avg:58.82ms step:1463/1600 train_time:86083ms step_avg:58.84ms step:1464/1600 train_time:86168ms step_avg:58.86ms step:1465/1600 train_time:86257ms step_avg:58.88ms step:1466/1600 train_time:86343ms step_avg:58.90ms step:1467/1600 train_time:86431ms step_avg:58.92ms step:1468/1600 train_time:86516ms step_avg:58.93ms step:1469/1600 train_time:86604ms step_avg:58.95ms step:1470/1600 train_time:86689ms step_avg:58.97ms step:1471/1600 train_time:86777ms step_avg:58.99ms step:1472/1600 train_time:86862ms step_avg:59.01ms step:1473/1600 train_time:86951ms step_avg:59.03ms step:1474/1600 train_time:87036ms step_avg:59.05ms step:1475/1600 train_time:87124ms step_avg:59.07ms step:1476/1600 train_time:87211ms step_avg:59.09ms step:1477/1600 train_time:87298ms step_avg:59.11ms step:1478/1600 train_time:87384ms step_avg:59.12ms step:1479/1600 train_time:87472ms step_avg:59.14ms step:1480/1600 train_time:87557ms step_avg:59.16ms step:1481/1600 train_time:87645ms step_avg:59.18ms step:1482/1600 train_time:87731ms step_avg:59.20ms step:1483/1600 train_time:87819ms step_avg:59.22ms step:1484/1600 train_time:87904ms step_avg:59.23ms step:1485/1600 train_time:87993ms step_avg:59.25ms step:1486/1600 train_time:88079ms step_avg:59.27ms step:1487/1600 train_time:88169ms step_avg:59.29ms step:1488/1600 train_time:88253ms step_avg:59.31ms step:1489/1600 train_time:88341ms step_avg:59.33ms step:1490/1600 train_time:88427ms step_avg:59.35ms step:1491/1600 train_time:88516ms step_avg:59.37ms step:1492/1600 train_time:88599ms step_avg:59.38ms step:1493/1600 train_time:88688ms step_avg:59.40ms step:1494/1600 train_time:88774ms step_avg:59.42ms step:1495/1600 train_time:88862ms step_avg:59.44ms step:1496/1600 train_time:88947ms step_avg:59.46ms step:1497/1600 train_time:89036ms step_avg:59.48ms step:1498/1600 train_time:89122ms step_avg:59.49ms step:1499/1600 train_time:89211ms step_avg:59.51ms step:1500/1600 train_time:89296ms step_avg:59.53ms step:1500/1600 val_loss:3.3070 train_time:89367ms step_avg:59.58ms step:1501/1600 train_time:89389ms step_avg:59.55ms step:1502/1600 train_time:89474ms step_avg:59.57ms step:1503/1600 train_time:89566ms step_avg:59.59ms step:1504/1600 train_time:89652ms step_avg:59.61ms step:1505/1600 train_time:89739ms step_avg:59.63ms step:1506/1600 train_time:89824ms step_avg:59.64ms step:1507/1600 train_time:89912ms step_avg:59.66ms step:1508/1600 train_time:89996ms step_avg:59.68ms step:1509/1600 train_time:90083ms step_avg:59.70ms step:1510/1600 train_time:90166ms step_avg:59.71ms step:1511/1600 train_time:90253ms step_avg:59.73ms step:1512/1600 train_time:90341ms step_avg:59.75ms step:1513/1600 train_time:90432ms step_avg:59.77ms step:1514/1600 train_time:90520ms step_avg:59.79ms step:1515/1600 train_time:90610ms step_avg:59.81ms step:1516/1600 train_time:90696ms step_avg:59.83ms step:1517/1600 train_time:90783ms step_avg:59.84ms step:1518/1600 train_time:90868ms step_avg:59.86ms step:1519/1600 train_time:90956ms step_avg:59.88ms step:1520/1600 train_time:91039ms step_avg:59.89ms step:1521/1600 train_time:91128ms step_avg:59.91ms step:1522/1600 train_time:91212ms step_avg:59.93ms step:1523/1600 train_time:91300ms step_avg:59.95ms step:1524/1600 train_time:91387ms step_avg:59.97ms step:1525/1600 train_time:91477ms step_avg:59.99ms step:1526/1600 train_time:91563ms step_avg:60.00ms step:1527/1600 train_time:91652ms step_avg:60.02ms step:1528/1600 train_time:91738ms step_avg:60.04ms step:1529/1600 train_time:91827ms step_avg:60.06ms step:1530/1600 train_time:91911ms step_avg:60.07ms step:1531/1600 train_time:91999ms step_avg:60.09ms step:1532/1600 train_time:92086ms step_avg:60.11ms step:1533/1600 train_time:92172ms step_avg:60.13ms step:1534/1600 train_time:92256ms step_avg:60.14ms step:1535/1600 train_time:92345ms step_avg:60.16ms step:1536/1600 train_time:92432ms step_avg:60.18ms step:1537/1600 train_time:92520ms step_avg:60.20ms step:1538/1600 train_time:92608ms step_avg:60.21ms step:1539/1600 train_time:92696ms step_avg:60.23ms step:1540/1600 train_time:92782ms step_avg:60.25ms step:1541/1600 train_time:92870ms step_avg:60.27ms step:1542/1600 train_time:92955ms step_avg:60.28ms step:1543/1600 train_time:93043ms step_avg:60.30ms step:1544/1600 train_time:93128ms step_avg:60.32ms step:1545/1600 train_time:93215ms step_avg:60.33ms step:1546/1600 train_time:93300ms step_avg:60.35ms step:1547/1600 train_time:93389ms step_avg:60.37ms step:1548/1600 train_time:93475ms step_avg:60.38ms step:1549/1600 train_time:93564ms step_avg:60.40ms step:1550/1600 train_time:93650ms step_avg:60.42ms step:1551/1600 train_time:93738ms step_avg:60.44ms step:1552/1600 train_time:93823ms step_avg:60.45ms step:1553/1600 train_time:93912ms step_avg:60.47ms step:1554/1600 train_time:93998ms step_avg:60.49ms step:1555/1600 train_time:94086ms step_avg:60.51ms step:1556/1600 train_time:94170ms step_avg:60.52ms step:1557/1600 train_time:94258ms step_avg:60.54ms step:1558/1600 train_time:94344ms step_avg:60.55ms step:1559/1600 train_time:94434ms step_avg:60.57ms step:1560/1600 train_time:94518ms step_avg:60.59ms step:1561/1600 train_time:94613ms step_avg:60.61ms step:1562/1600 train_time:94697ms step_avg:60.63ms step:1563/1600 train_time:94785ms step_avg:60.64ms step:1564/1600 train_time:94871ms step_avg:60.66ms step:1565/1600 train_time:94960ms step_avg:60.68ms step:1566/1600 train_time:95045ms step_avg:60.69ms step:1567/1600 train_time:95133ms step_avg:60.71ms step:1568/1600 train_time:95218ms step_avg:60.73ms step:1569/1600 train_time:95307ms step_avg:60.74ms step:1570/1600 train_time:95393ms step_avg:60.76ms step:1571/1600 train_time:95482ms step_avg:60.78ms step:1572/1600 train_time:95569ms step_avg:60.79ms step:1573/1600 train_time:95658ms step_avg:60.81ms step:1574/1600 train_time:95744ms step_avg:60.83ms step:1575/1600 train_time:95833ms step_avg:60.85ms step:1576/1600 train_time:95918ms step_avg:60.86ms step:1577/1600 train_time:96011ms step_avg:60.88ms step:1578/1600 train_time:96096ms step_avg:60.90ms step:1579/1600 train_time:96183ms step_avg:60.91ms step:1580/1600 train_time:96268ms step_avg:60.93ms step:1581/1600 train_time:96356ms step_avg:60.95ms step:1582/1600 train_time:96442ms step_avg:60.96ms step:1583/1600 train_time:96531ms step_avg:60.98ms step:1584/1600 train_time:96617ms step_avg:61.00ms step:1585/1600 train_time:96705ms step_avg:61.01ms step:1586/1600 train_time:96791ms step_avg:61.03ms step:1587/1600 train_time:96879ms step_avg:61.05ms step:1588/1600 train_time:96965ms step_avg:61.06ms step:1589/1600 train_time:97053ms step_avg:61.08ms step:1590/1600 train_time:97138ms step_avg:61.09ms step:1591/1600 train_time:97227ms step_avg:61.11ms step:1592/1600 train_time:97313ms step_avg:61.13ms step:1593/1600 train_time:97401ms step_avg:61.14ms step:1594/1600 train_time:97487ms step_avg:61.16ms step:1595/1600 train_time:97576ms step_avg:61.18ms step:1596/1600 train_time:97662ms step_avg:61.19ms step:1597/1600 train_time:97751ms step_avg:61.21ms step:1598/1600 train_time:97836ms step_avg:61.22ms step:1599/1600 train_time:97924ms step_avg:61.24ms step:1600/1600 train_time:98010ms step_avg:61.26ms step:1600/1600 val_loss:3.2776 train_time:98081ms step_avg:61.30ms peak memory allocated: 30789 MiB reserved: 46238 MiB