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:19:58 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 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 41C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 35C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:6A:00.0 Off | 0 | | N/A 35C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 42C P0 131W / 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 37C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 299489 C /usr/bin/python3 1510MiB | | 1 N/A N/A 299490 C /usr/bin/python3 1510MiB | | 2 N/A N/A 299491 C /usr/bin/python3 1510MiB | | 3 N/A N/A 299492 C /usr/bin/python3 1510MiB | | 4 N/A N/A 299493 C /usr/bin/python3 1510MiB | | 5 N/A N/A 299494 C /usr/bin/python3 1510MiB | | 6 N/A N/A 299495 C /usr/bin/python3 1510MiB | | 7 N/A N/A 299496 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.8300 train_time:0ms step_avg:0.03ms step:1/1600 train_time:86ms step_avg:86.02ms step:2/1600 train_time:107ms step_avg:53.40ms step:3/1600 train_time:126ms step_avg:41.83ms step:4/1600 train_time:154ms step_avg:38.61ms step:5/1600 train_time:184ms step_avg:36.90ms step:6/1600 train_time:297ms step_avg:49.47ms step:7/1600 train_time:313ms step_avg:44.78ms step:8/1600 train_time:432ms step_avg:53.99ms step:9/1600 train_time:462ms step_avg:51.35ms step:10/1600 train_time:499ms step_avg:49.90ms step:11/1600 train_time:530ms step_avg:48.16ms step:12/1600 train_time:567ms step_avg:47.23ms step:13/1600 train_time:597ms step_avg:45.96ms step:14/1600 train_time:635ms step_avg:45.34ms step:15/1600 train_time:666ms step_avg:44.38ms step:16/1600 train_time:703ms step_avg:43.94ms step:17/1600 train_time:734ms step_avg:43.18ms step:18/1600 train_time:771ms step_avg:42.85ms step:19/1600 train_time:802ms step_avg:42.22ms step:20/1600 train_time:839ms step_avg:41.96ms step:21/1600 train_time:870ms step_avg:41.41ms step:22/1600 train_time:906ms step_avg:41.20ms step:23/1600 train_time:938ms step_avg:40.76ms step:24/1600 train_time:975ms step_avg:40.63ms step:25/1600 train_time:1006ms step_avg:40.24ms step:26/1600 train_time:1043ms step_avg:40.13ms step:27/1600 train_time:1074ms step_avg:39.77ms step:28/1600 train_time:1111ms step_avg:39.66ms step:29/1600 train_time:1142ms step_avg:39.37ms step:30/1600 train_time:1179ms step_avg:39.29ms step:31/1600 train_time:1209ms step_avg:39.01ms step:32/1600 train_time:1247ms step_avg:38.97ms step:33/1600 train_time:1278ms step_avg:38.72ms step:34/1600 train_time:1315ms step_avg:38.68ms step:35/1600 train_time:1346ms step_avg:38.46ms step:36/1600 train_time:1383ms step_avg:38.43ms step:37/1600 train_time:1414ms step_avg:38.22ms step:38/1600 train_time:1452ms step_avg:38.20ms step:39/1600 train_time:1482ms step_avg:38.01ms step:40/1600 train_time:1519ms step_avg:37.98ms step:41/1600 train_time:1550ms step_avg:37.80ms step:42/1600 train_time:1587ms step_avg:37.78ms step:43/1600 train_time:1618ms step_avg:37.62ms step:44/1600 train_time:1655ms step_avg:37.62ms step:45/1600 train_time:1685ms step_avg:37.45ms step:46/1600 train_time:1723ms step_avg:37.46ms step:47/1600 train_time:1753ms step_avg:37.31ms step:48/1600 train_time:1791ms step_avg:37.31ms step:49/1600 train_time:1822ms step_avg:37.18ms step:50/1600 train_time:1859ms step_avg:37.18ms step:51/1600 train_time:1889ms step_avg:37.05ms step:52/1600 train_time:1927ms step_avg:37.06ms step:53/1600 train_time:1958ms step_avg:36.94ms step:54/1600 train_time:1995ms step_avg:36.94ms step:55/1600 train_time:2026ms step_avg:36.84ms step:56/1600 train_time:2064ms step_avg:36.86ms step:57/1600 train_time:2095ms step_avg:36.75ms step:58/1600 train_time:2132ms step_avg:36.75ms step:59/1600 train_time:2163ms step_avg:36.65ms step:60/1600 train_time:2199ms step_avg:36.66ms step:61/1600 train_time:2230ms step_avg:36.56ms step:62/1600 train_time:2268ms step_avg:36.57ms step:63/1600 train_time:2299ms step_avg:36.49ms step:64/1600 train_time:2336ms step_avg:36.50ms step:65/1600 train_time:2367ms step_avg:36.41ms step:66/1600 train_time:2404ms step_avg:36.42ms step:67/1600 train_time:2435ms step_avg:36.34ms step:68/1600 train_time:2472ms step_avg:36.36ms step:69/1600 train_time:2503ms step_avg:36.28ms step:70/1600 train_time:2540ms step_avg:36.29ms step:71/1600 train_time:2571ms step_avg:36.21ms step:72/1600 train_time:2608ms step_avg:36.23ms step:73/1600 train_time:2640ms step_avg:36.17ms step:74/1600 train_time:2678ms step_avg:36.18ms step:75/1600 train_time:2708ms step_avg:36.11ms step:76/1600 train_time:2745ms step_avg:36.12ms step:77/1600 train_time:2776ms step_avg:36.05ms step:78/1600 train_time:2813ms step_avg:36.07ms step:79/1600 train_time:2844ms step_avg:36.00ms step:80/1600 train_time:2881ms step_avg:36.02ms step:81/1600 train_time:2913ms step_avg:35.96ms step:82/1600 train_time:2950ms step_avg:35.98ms step:83/1600 train_time:2981ms step_avg:35.92ms step:84/1600 train_time:3018ms step_avg:35.93ms step:85/1600 train_time:3049ms step_avg:35.87ms step:86/1600 train_time:3086ms step_avg:35.88ms step:87/1600 train_time:3117ms step_avg:35.82ms step:88/1600 train_time:3154ms step_avg:35.84ms step:89/1600 train_time:3184ms step_avg:35.78ms step:90/1600 train_time:3221ms step_avg:35.79ms step:91/1600 train_time:3252ms step_avg:35.74ms step:92/1600 train_time:3289ms step_avg:35.75ms step:93/1600 train_time:3321ms step_avg:35.71ms step:94/1600 train_time:3358ms step_avg:35.72ms step:95/1600 train_time:3389ms step_avg:35.67ms step:96/1600 train_time:3426ms step_avg:35.69ms step:97/1600 train_time:3457ms step_avg:35.64ms step:98/1600 train_time:3494ms step_avg:35.65ms step:99/1600 train_time:3525ms step_avg:35.60ms step:100/1600 train_time:3562ms step_avg:35.62ms step:101/1600 train_time:3593ms step_avg:35.57ms step:102/1600 train_time:3630ms step_avg:35.59ms step:103/1600 train_time:3661ms step_avg:35.54ms step:104/1600 train_time:3697ms step_avg:35.55ms step:105/1600 train_time:3728ms step_avg:35.51ms step:106/1600 train_time:3766ms step_avg:35.53ms step:107/1600 train_time:3797ms step_avg:35.49ms step:108/1600 train_time:3835ms step_avg:35.51ms step:109/1600 train_time:3866ms step_avg:35.47ms step:110/1600 train_time:3903ms step_avg:35.48ms step:111/1600 train_time:3934ms step_avg:35.44ms step:112/1600 train_time:3971ms step_avg:35.46ms step:113/1600 train_time:4002ms step_avg:35.42ms step:114/1600 train_time:4039ms step_avg:35.43ms step:115/1600 train_time:4070ms step_avg:35.39ms step:116/1600 train_time:4107ms step_avg:35.41ms step:117/1600 train_time:4138ms step_avg:35.37ms step:118/1600 train_time:4176ms step_avg:35.39ms step:119/1600 train_time:4207ms step_avg:35.35ms step:120/1600 train_time:4244ms step_avg:35.37ms step:121/1600 train_time:4275ms step_avg:35.33ms step:122/1600 train_time:4312ms step_avg:35.34ms step:123/1600 train_time:4343ms step_avg:35.31ms step:124/1600 train_time:4380ms step_avg:35.32ms step:125/1600 train_time:4411ms step_avg:35.29ms step:126/1600 train_time:4448ms step_avg:35.31ms step:127/1600 train_time:4479ms step_avg:35.27ms step:128/1600 train_time:4516ms step_avg:35.28ms step:129/1600 train_time:4548ms step_avg:35.25ms step:130/1600 train_time:4585ms step_avg:35.27ms step:131/1600 train_time:4615ms step_avg:35.23ms step:132/1600 train_time:4653ms step_avg:35.25ms step:133/1600 train_time:4683ms step_avg:35.21ms step:134/1600 train_time:4720ms step_avg:35.23ms step:135/1600 train_time:4751ms step_avg:35.19ms step:136/1600 train_time:4789ms step_avg:35.21ms step:137/1600 train_time:4820ms step_avg:35.18ms step:138/1600 train_time:4857ms step_avg:35.20ms step:139/1600 train_time:4887ms step_avg:35.16ms step:140/1600 train_time:4925ms step_avg:35.18ms step:141/1600 train_time:4956ms step_avg:35.15ms step:142/1600 train_time:4993ms step_avg:35.16ms step:143/1600 train_time:5024ms step_avg:35.13ms step:144/1600 train_time:5061ms step_avg:35.15ms step:145/1600 train_time:5092ms step_avg:35.12ms step:146/1600 train_time:5129ms step_avg:35.13ms step:147/1600 train_time:5160ms step_avg:35.10ms step:148/1600 train_time:5197ms step_avg:35.12ms step:149/1600 train_time:5228ms step_avg:35.09ms step:150/1600 train_time:5265ms step_avg:35.10ms step:151/1600 train_time:5296ms step_avg:35.07ms step:152/1600 train_time:5333ms step_avg:35.09ms step:153/1600 train_time:5364ms step_avg:35.06ms step:154/1600 train_time:5401ms step_avg:35.07ms step:155/1600 train_time:5433ms step_avg:35.05ms step:156/1600 train_time:5470ms step_avg:35.06ms step:157/1600 train_time:5500ms step_avg:35.03ms step:158/1600 train_time:5537ms step_avg:35.05ms step:159/1600 train_time:5568ms step_avg:35.02ms step:160/1600 train_time:5605ms step_avg:35.03ms step:161/1600 train_time:5636ms step_avg:35.01ms step:162/1600 train_time:5674ms step_avg:35.02ms step:163/1600 train_time:5704ms step_avg:34.99ms step:164/1600 train_time:5741ms step_avg:35.01ms step:165/1600 train_time:5772ms step_avg:34.98ms step:166/1600 train_time:5810ms step_avg:35.00ms step:167/1600 train_time:5841ms step_avg:34.97ms step:168/1600 train_time:5878ms step_avg:34.99ms step:169/1600 train_time:5909ms step_avg:34.96ms step:170/1600 train_time:5946ms step_avg:34.97ms step:171/1600 train_time:5977ms step_avg:34.95ms step:172/1600 train_time:6014ms step_avg:34.96ms step:173/1600 train_time:6045ms step_avg:34.94ms step:174/1600 train_time:6082ms step_avg:34.95ms step:175/1600 train_time:6113ms step_avg:34.93ms step:176/1600 train_time:6149ms step_avg:34.94ms step:177/1600 train_time:6180ms step_avg:34.92ms step:178/1600 train_time:6218ms step_avg:34.93ms step:179/1600 train_time:6249ms step_avg:34.91ms step:180/1600 train_time:6286ms step_avg:34.92ms step:181/1600 train_time:6317ms step_avg:34.90ms step:182/1600 train_time:6353ms step_avg:34.91ms step:183/1600 train_time:6384ms step_avg:34.89ms step:184/1600 train_time:6421ms step_avg:34.90ms step:185/1600 train_time:6452ms step_avg:34.88ms step:186/1600 train_time:6489ms step_avg:34.89ms step:187/1600 train_time:6520ms step_avg:34.87ms step:188/1600 train_time:6557ms step_avg:34.88ms step:189/1600 train_time:6587ms step_avg:34.85ms step:190/1600 train_time:6624ms step_avg:34.87ms step:191/1600 train_time:6655ms step_avg:34.84ms step:192/1600 train_time:6692ms step_avg:34.86ms step:193/1600 train_time:6723ms step_avg:34.83ms step:194/1600 train_time:6760ms step_avg:34.85ms step:195/1600 train_time:6790ms step_avg:34.82ms step:196/1600 train_time:6828ms step_avg:34.84ms step:197/1600 train_time:6858ms step_avg:34.81ms step:198/1600 train_time:6895ms step_avg:34.82ms step:199/1600 train_time:6926ms step_avg:34.80ms step:200/1600 train_time:6963ms step_avg:34.82ms step:201/1600 train_time:6994ms step_avg:34.80ms step:202/1600 train_time:7032ms step_avg:34.81ms step:203/1600 train_time:7062ms step_avg:34.79ms step:204/1600 train_time:7100ms step_avg:34.80ms step:205/1600 train_time:7130ms step_avg:34.78ms step:206/1600 train_time:7167ms step_avg:34.79ms step:207/1600 train_time:7198ms step_avg:34.77ms step:208/1600 train_time:7235ms step_avg:34.78ms step:209/1600 train_time:7266ms step_avg:34.76ms step:210/1600 train_time:7303ms step_avg:34.78ms step:211/1600 train_time:7334ms step_avg:34.76ms step:212/1600 train_time:7371ms step_avg:34.77ms step:213/1600 train_time:7402ms step_avg:34.75ms step:214/1600 train_time:7439ms step_avg:34.76ms step:215/1600 train_time:7470ms step_avg:34.75ms step:216/1600 train_time:7507ms step_avg:34.76ms step:217/1600 train_time:7538ms step_avg:34.74ms step:218/1600 train_time:7575ms step_avg:34.75ms step:219/1600 train_time:7606ms step_avg:34.73ms step:220/1600 train_time:7644ms step_avg:34.74ms step:221/1600 train_time:7674ms step_avg:34.73ms step:222/1600 train_time:7711ms step_avg:34.74ms step:223/1600 train_time:7742ms step_avg:34.72ms step:224/1600 train_time:7779ms step_avg:34.73ms step:225/1600 train_time:7811ms step_avg:34.71ms step:226/1600 train_time:7848ms step_avg:34.73ms step:227/1600 train_time:7879ms step_avg:34.71ms step:228/1600 train_time:7916ms step_avg:34.72ms step:229/1600 train_time:7947ms step_avg:34.70ms step:230/1600 train_time:7984ms step_avg:34.71ms step:231/1600 train_time:8015ms step_avg:34.70ms step:232/1600 train_time:8052ms step_avg:34.71ms step:233/1600 train_time:8083ms step_avg:34.69ms step:234/1600 train_time:8120ms step_avg:34.70ms step:235/1600 train_time:8151ms step_avg:34.69ms step:236/1600 train_time:8190ms step_avg:34.70ms step:237/1600 train_time:8220ms step_avg:34.68ms step:238/1600 train_time:8257ms step_avg:34.69ms step:239/1600 train_time:8288ms step_avg:34.68ms step:240/1600 train_time:8325ms step_avg:34.69ms step:241/1600 train_time:8355ms step_avg:34.67ms step:242/1600 train_time:8392ms step_avg:34.68ms step:243/1600 train_time:8423ms step_avg:34.66ms step:244/1600 train_time:8460ms step_avg:34.67ms step:245/1600 train_time:8491ms step_avg:34.66ms step:246/1600 train_time:8528ms step_avg:34.67ms step:247/1600 train_time:8558ms step_avg:34.65ms step:248/1600 train_time:8595ms step_avg:34.66ms step:249/1600 train_time:8626ms step_avg:34.64ms step:250/1600 train_time:8664ms step_avg:34.66ms step:250/1600 val_loss:4.5841 train_time:8711ms step_avg:34.84ms step:251/1600 train_time:8729ms step_avg:34.78ms step:252/1600 train_time:8748ms step_avg:34.71ms step:253/1600 train_time:8765ms step_avg:34.65ms step:254/1600 train_time:8803ms step_avg:34.66ms step:255/1600 train_time:8835ms step_avg:34.65ms step:256/1600 train_time:8873ms step_avg:34.66ms step:257/1600 train_time:8905ms step_avg:34.65ms step:258/1600 train_time:8942ms step_avg:34.66ms step:259/1600 train_time:8973ms step_avg:34.64ms step:260/1600 train_time:9011ms step_avg:34.66ms step:261/1600 train_time:9041ms step_avg:34.64ms step:262/1600 train_time:9078ms step_avg:34.65ms step:263/1600 train_time:9109ms step_avg:34.63ms step:264/1600 train_time:9145ms step_avg:34.64ms step:265/1600 train_time:9176ms step_avg:34.63ms step:266/1600 train_time:9213ms step_avg:34.64ms step:267/1600 train_time:9244ms step_avg:34.62ms step:268/1600 train_time:9281ms step_avg:34.63ms step:269/1600 train_time:9311ms step_avg:34.61ms step:270/1600 train_time:9348ms step_avg:34.62ms step:271/1600 train_time:9379ms step_avg:34.61ms step:272/1600 train_time:9416ms step_avg:34.62ms step:273/1600 train_time:9446ms step_avg:34.60ms step:274/1600 train_time:9484ms step_avg:34.61ms step:275/1600 train_time:9514ms step_avg:34.60ms step:276/1600 train_time:9551ms step_avg:34.61ms step:277/1600 train_time:9582ms step_avg:34.59ms step:278/1600 train_time:9619ms step_avg:34.60ms step:279/1600 train_time:9649ms step_avg:34.59ms step:280/1600 train_time:9686ms step_avg:34.59ms step:281/1600 train_time:9717ms step_avg:34.58ms step:282/1600 train_time:9754ms step_avg:34.59ms step:283/1600 train_time:9784ms step_avg:34.57ms step:284/1600 train_time:9822ms step_avg:34.58ms step:285/1600 train_time:9852ms step_avg:34.57ms step:286/1600 train_time:9890ms step_avg:34.58ms step:287/1600 train_time:9921ms step_avg:34.57ms step:288/1600 train_time:9958ms step_avg:34.58ms step:289/1600 train_time:9989ms step_avg:34.57ms step:290/1600 train_time:10026ms step_avg:34.57ms step:291/1600 train_time:10057ms step_avg:34.56ms step:292/1600 train_time:10094ms step_avg:34.57ms step:293/1600 train_time:10125ms step_avg:34.56ms step:294/1600 train_time:10162ms step_avg:34.56ms step:295/1600 train_time:10192ms step_avg:34.55ms step:296/1600 train_time:10230ms step_avg:34.56ms step:297/1600 train_time:10260ms step_avg:34.55ms step:298/1600 train_time:10297ms step_avg:34.56ms step:299/1600 train_time:10328ms step_avg:34.54ms step:300/1600 train_time:10365ms step_avg:34.55ms step:301/1600 train_time:10396ms step_avg:34.54ms step:302/1600 train_time:10433ms step_avg:34.55ms step:303/1600 train_time:10464ms step_avg:34.53ms step:304/1600 train_time:10500ms step_avg:34.54ms step:305/1600 train_time:10531ms step_avg:34.53ms step:306/1600 train_time:10568ms step_avg:34.54ms step:307/1600 train_time:10598ms step_avg:34.52ms step:308/1600 train_time:10635ms step_avg:34.53ms step:309/1600 train_time:10666ms step_avg:34.52ms step:310/1600 train_time:10703ms step_avg:34.53ms step:311/1600 train_time:10734ms step_avg:34.51ms step:312/1600 train_time:10771ms step_avg:34.52ms step:313/1600 train_time:10801ms step_avg:34.51ms step:314/1600 train_time:10839ms step_avg:34.52ms step:315/1600 train_time:10869ms step_avg:34.51ms step:316/1600 train_time:10906ms step_avg:34.51ms step:317/1600 train_time:10937ms step_avg:34.50ms step:318/1600 train_time:10974ms step_avg:34.51ms step:319/1600 train_time:11005ms step_avg:34.50ms step:320/1600 train_time:11043ms step_avg:34.51ms step:321/1600 train_time:11073ms step_avg:34.50ms step:322/1600 train_time:11111ms step_avg:34.51ms step:323/1600 train_time:11141ms step_avg:34.49ms step:324/1600 train_time:11179ms step_avg:34.50ms step:325/1600 train_time:11210ms step_avg:34.49ms step:326/1600 train_time:11247ms step_avg:34.50ms step:327/1600 train_time:11278ms step_avg:34.49ms step:328/1600 train_time:11315ms step_avg:34.50ms step:329/1600 train_time:11346ms step_avg:34.48ms step:330/1600 train_time:11382ms step_avg:34.49ms step:331/1600 train_time:11413ms step_avg:34.48ms step:332/1600 train_time:11450ms step_avg:34.49ms step:333/1600 train_time:11481ms step_avg:34.48ms step:334/1600 train_time:11519ms step_avg:34.49ms step:335/1600 train_time:11549ms step_avg:34.47ms step:336/1600 train_time:11586ms step_avg:34.48ms step:337/1600 train_time:11616ms step_avg:34.47ms step:338/1600 train_time:11653ms step_avg:34.48ms step:339/1600 train_time:11684ms step_avg:34.46ms step:340/1600 train_time:11720ms step_avg:34.47ms step:341/1600 train_time:11751ms step_avg:34.46ms step:342/1600 train_time:11788ms step_avg:34.47ms step:343/1600 train_time:11819ms step_avg:34.46ms step:344/1600 train_time:11856ms step_avg:34.46ms step:345/1600 train_time:11886ms step_avg:34.45ms step:346/1600 train_time:11923ms step_avg:34.46ms step:347/1600 train_time:11954ms step_avg:34.45ms step:348/1600 train_time:11991ms step_avg:34.46ms step:349/1600 train_time:12022ms step_avg:34.45ms step:350/1600 train_time:12059ms step_avg:34.45ms step:351/1600 train_time:12089ms step_avg:34.44ms step:352/1600 train_time:12126ms step_avg:34.45ms step:353/1600 train_time:12157ms step_avg:34.44ms step:354/1600 train_time:12194ms step_avg:34.45ms step:355/1600 train_time:12225ms step_avg:34.44ms step:356/1600 train_time:12261ms step_avg:34.44ms step:357/1600 train_time:12292ms step_avg:34.43ms step:358/1600 train_time:12329ms step_avg:34.44ms step:359/1600 train_time:12360ms step_avg:34.43ms step:360/1600 train_time:12397ms step_avg:34.44ms step:361/1600 train_time:12427ms step_avg:34.42ms step:362/1600 train_time:12464ms step_avg:34.43ms step:363/1600 train_time:12495ms step_avg:34.42ms step:364/1600 train_time:12532ms step_avg:34.43ms step:365/1600 train_time:12562ms step_avg:34.42ms step:366/1600 train_time:12599ms step_avg:34.42ms step:367/1600 train_time:12630ms step_avg:34.41ms step:368/1600 train_time:12667ms step_avg:34.42ms step:369/1600 train_time:12698ms step_avg:34.41ms step:370/1600 train_time:12735ms step_avg:34.42ms step:371/1600 train_time:12765ms step_avg:34.41ms step:372/1600 train_time:12802ms step_avg:34.41ms step:373/1600 train_time:12833ms step_avg:34.40ms step:374/1600 train_time:12870ms step_avg:34.41ms step:375/1600 train_time:12901ms step_avg:34.40ms step:376/1600 train_time:12938ms step_avg:34.41ms step:377/1600 train_time:12969ms step_avg:34.40ms step:378/1600 train_time:13006ms step_avg:34.41ms step:379/1600 train_time:13036ms step_avg:34.40ms step:380/1600 train_time:13074ms step_avg:34.40ms step:381/1600 train_time:13104ms step_avg:34.39ms step:382/1600 train_time:13141ms step_avg:34.40ms step:383/1600 train_time:13172ms step_avg:34.39ms step:384/1600 train_time:13209ms step_avg:34.40ms step:385/1600 train_time:13239ms step_avg:34.39ms step:386/1600 train_time:13277ms step_avg:34.40ms step:387/1600 train_time:13307ms step_avg:34.39ms step:388/1600 train_time:13344ms step_avg:34.39ms step:389/1600 train_time:13375ms step_avg:34.38ms step:390/1600 train_time:13412ms step_avg:34.39ms step:391/1600 train_time:13443ms step_avg:34.38ms step:392/1600 train_time:13480ms step_avg:34.39ms step:393/1600 train_time:13511ms step_avg:34.38ms step:394/1600 train_time:13548ms step_avg:34.39ms step:395/1600 train_time:13579ms step_avg:34.38ms step:396/1600 train_time:13616ms step_avg:34.38ms step:397/1600 train_time:13647ms step_avg:34.37ms step:398/1600 train_time:13684ms step_avg:34.38ms step:399/1600 train_time:13714ms step_avg:34.37ms step:400/1600 train_time:13752ms step_avg:34.38ms step:401/1600 train_time:13782ms step_avg:34.37ms step:402/1600 train_time:13819ms step_avg:34.38ms step:403/1600 train_time:13850ms step_avg:34.37ms step:404/1600 train_time:13887ms step_avg:34.37ms step:405/1600 train_time:13918ms step_avg:34.37ms step:406/1600 train_time:13955ms step_avg:34.37ms step:407/1600 train_time:13985ms step_avg:34.36ms step:408/1600 train_time:14022ms step_avg:34.37ms step:409/1600 train_time:14053ms step_avg:34.36ms step:410/1600 train_time:14091ms step_avg:34.37ms step:411/1600 train_time:14121ms step_avg:34.36ms step:412/1600 train_time:14158ms step_avg:34.36ms step:413/1600 train_time:14189ms step_avg:34.36ms step:414/1600 train_time:14226ms step_avg:34.36ms step:415/1600 train_time:14257ms step_avg:34.35ms step:416/1600 train_time:14294ms step_avg:34.36ms step:417/1600 train_time:14325ms step_avg:34.35ms step:418/1600 train_time:14361ms step_avg:34.36ms step:419/1600 train_time:14392ms step_avg:34.35ms step:420/1600 train_time:14429ms step_avg:34.36ms step:421/1600 train_time:14460ms step_avg:34.35ms step:422/1600 train_time:14497ms step_avg:34.35ms step:423/1600 train_time:14528ms step_avg:34.34ms step:424/1600 train_time:14564ms step_avg:34.35ms step:425/1600 train_time:14595ms step_avg:34.34ms step:426/1600 train_time:14632ms step_avg:34.35ms step:427/1600 train_time:14663ms step_avg:34.34ms step:428/1600 train_time:14700ms step_avg:34.35ms step:429/1600 train_time:14731ms step_avg:34.34ms step:430/1600 train_time:14769ms step_avg:34.35ms step:431/1600 train_time:14799ms step_avg:34.34ms step:432/1600 train_time:14836ms step_avg:34.34ms step:433/1600 train_time:14867ms step_avg:34.33ms step:434/1600 train_time:14903ms step_avg:34.34ms step:435/1600 train_time:14934ms step_avg:34.33ms step:436/1600 train_time:14972ms step_avg:34.34ms step:437/1600 train_time:15002ms step_avg:34.33ms step:438/1600 train_time:15039ms step_avg:34.34ms step:439/1600 train_time:15070ms step_avg:34.33ms step:440/1600 train_time:15108ms step_avg:34.34ms step:441/1600 train_time:15139ms step_avg:34.33ms step:442/1600 train_time:15176ms step_avg:34.33ms step:443/1600 train_time:15206ms step_avg:34.33ms step:444/1600 train_time:15243ms step_avg:34.33ms step:445/1600 train_time:15274ms step_avg:34.32ms step:446/1600 train_time:15311ms step_avg:34.33ms step:447/1600 train_time:15342ms step_avg:34.32ms step:448/1600 train_time:15379ms step_avg:34.33ms step:449/1600 train_time:15410ms step_avg:34.32ms step:450/1600 train_time:15447ms step_avg:34.33ms step:451/1600 train_time:15478ms step_avg:34.32ms step:452/1600 train_time:15515ms step_avg:34.33ms step:453/1600 train_time:15546ms step_avg:34.32ms step:454/1600 train_time:15583ms step_avg:34.32ms step:455/1600 train_time:15613ms step_avg:34.31ms step:456/1600 train_time:15650ms step_avg:34.32ms step:457/1600 train_time:15681ms step_avg:34.31ms step:458/1600 train_time:15718ms step_avg:34.32ms step:459/1600 train_time:15748ms step_avg:34.31ms step:460/1600 train_time:15785ms step_avg:34.32ms step:461/1600 train_time:15816ms step_avg:34.31ms step:462/1600 train_time:15853ms step_avg:34.31ms step:463/1600 train_time:15884ms step_avg:34.31ms step:464/1600 train_time:15921ms step_avg:34.31ms step:465/1600 train_time:15952ms step_avg:34.30ms step:466/1600 train_time:15988ms step_avg:34.31ms step:467/1600 train_time:16020ms step_avg:34.30ms step:468/1600 train_time:16057ms step_avg:34.31ms step:469/1600 train_time:16087ms step_avg:34.30ms step:470/1600 train_time:16125ms step_avg:34.31ms step:471/1600 train_time:16155ms step_avg:34.30ms step:472/1600 train_time:16192ms step_avg:34.31ms step:473/1600 train_time:16223ms step_avg:34.30ms step:474/1600 train_time:16260ms step_avg:34.30ms step:475/1600 train_time:16291ms step_avg:34.30ms step:476/1600 train_time:16328ms step_avg:34.30ms step:477/1600 train_time:16358ms step_avg:34.29ms step:478/1600 train_time:16396ms step_avg:34.30ms step:479/1600 train_time:16426ms step_avg:34.29ms step:480/1600 train_time:16463ms step_avg:34.30ms step:481/1600 train_time:16494ms step_avg:34.29ms step:482/1600 train_time:16532ms step_avg:34.30ms step:483/1600 train_time:16563ms step_avg:34.29ms step:484/1600 train_time:16599ms step_avg:34.30ms step:485/1600 train_time:16630ms step_avg:34.29ms step:486/1600 train_time:16667ms step_avg:34.30ms step:487/1600 train_time:16698ms step_avg:34.29ms step:488/1600 train_time:16736ms step_avg:34.29ms step:489/1600 train_time:16766ms step_avg:34.29ms step:490/1600 train_time:16803ms step_avg:34.29ms step:491/1600 train_time:16833ms step_avg:34.28ms step:492/1600 train_time:16871ms step_avg:34.29ms step:493/1600 train_time:16902ms step_avg:34.28ms step:494/1600 train_time:16939ms step_avg:34.29ms step:495/1600 train_time:16970ms step_avg:34.28ms step:496/1600 train_time:17007ms step_avg:34.29ms step:497/1600 train_time:17038ms step_avg:34.28ms step:498/1600 train_time:17075ms step_avg:34.29ms step:499/1600 train_time:17105ms step_avg:34.28ms step:500/1600 train_time:17142ms step_avg:34.28ms step:500/1600 val_loss:4.2321 train_time:17190ms step_avg:34.38ms step:501/1600 train_time:17208ms step_avg:34.35ms step:502/1600 train_time:17227ms step_avg:34.32ms step:503/1600 train_time:17243ms step_avg:34.28ms step:504/1600 train_time:17280ms step_avg:34.29ms step:505/1600 train_time:17312ms step_avg:34.28ms step:506/1600 train_time:17349ms step_avg:34.29ms step:507/1600 train_time:17380ms step_avg:34.28ms step:508/1600 train_time:17417ms step_avg:34.29ms step:509/1600 train_time:17448ms step_avg:34.28ms step:510/1600 train_time:17486ms step_avg:34.29ms step:511/1600 train_time:17516ms step_avg:34.28ms step:512/1600 train_time:17553ms step_avg:34.28ms step:513/1600 train_time:17584ms step_avg:34.28ms step:514/1600 train_time:17620ms step_avg:34.28ms step:515/1600 train_time:17651ms step_avg:34.27ms step:516/1600 train_time:17688ms step_avg:34.28ms step:517/1600 train_time:17718ms step_avg:34.27ms step:518/1600 train_time:17756ms step_avg:34.28ms step:519/1600 train_time:17786ms step_avg:34.27ms step:520/1600 train_time:17823ms step_avg:34.27ms step:521/1600 train_time:17893ms step_avg:34.34ms step:522/1600 train_time:17950ms step_avg:34.39ms step:523/1600 train_time:18011ms step_avg:34.44ms step:524/1600 train_time:18069ms step_avg:34.48ms step:525/1600 train_time:18131ms step_avg:34.54ms step:526/1600 train_time:18191ms step_avg:34.58ms step:527/1600 train_time:18255ms step_avg:34.64ms step:528/1600 train_time:18315ms step_avg:34.69ms step:529/1600 train_time:18378ms step_avg:34.74ms step:530/1600 train_time:18437ms step_avg:34.79ms step:531/1600 train_time:18500ms step_avg:34.84ms step:532/1600 train_time:18558ms step_avg:34.88ms step:533/1600 train_time:18621ms step_avg:34.94ms step:534/1600 train_time:18680ms step_avg:34.98ms step:535/1600 train_time:18743ms step_avg:35.03ms step:536/1600 train_time:18802ms step_avg:35.08ms step:537/1600 train_time:18864ms step_avg:35.13ms step:538/1600 train_time:18923ms step_avg:35.17ms step:539/1600 train_time:18985ms step_avg:35.22ms step:540/1600 train_time:19044ms step_avg:35.27ms step:541/1600 train_time:19106ms step_avg:35.32ms step:542/1600 train_time:19165ms step_avg:35.36ms step:543/1600 train_time:19227ms step_avg:35.41ms step:544/1600 train_time:19287ms step_avg:35.45ms step:545/1600 train_time:19349ms step_avg:35.50ms step:546/1600 train_time:19408ms step_avg:35.55ms step:547/1600 train_time:19470ms step_avg:35.59ms step:548/1600 train_time:19530ms step_avg:35.64ms step:549/1600 train_time:19592ms step_avg:35.69ms step:550/1600 train_time:19651ms step_avg:35.73ms step:551/1600 train_time:19714ms step_avg:35.78ms step:552/1600 train_time:19774ms step_avg:35.82ms step:553/1600 train_time:19835ms step_avg:35.87ms step:554/1600 train_time:19895ms step_avg:35.91ms step:555/1600 train_time:19957ms step_avg:35.96ms step:556/1600 train_time:20017ms step_avg:36.00ms step:557/1600 train_time:20079ms step_avg:36.05ms step:558/1600 train_time:20138ms step_avg:36.09ms step:559/1600 train_time:20201ms step_avg:36.14ms step:560/1600 train_time:20260ms step_avg:36.18ms step:561/1600 train_time:20323ms step_avg:36.23ms step:562/1600 train_time:20386ms step_avg:36.27ms step:563/1600 train_time:20447ms step_avg:36.32ms step:564/1600 train_time:20507ms step_avg:36.36ms step:565/1600 train_time:20566ms step_avg:36.40ms step:566/1600 train_time:20625ms step_avg:36.44ms step:567/1600 train_time:20687ms step_avg:36.49ms step:568/1600 train_time:20746ms step_avg:36.52ms step:569/1600 train_time:20808ms step_avg:36.57ms step:570/1600 train_time:20866ms step_avg:36.61ms step:571/1600 train_time:20929ms step_avg:36.65ms step:572/1600 train_time:20988ms step_avg:36.69ms step:573/1600 train_time:21054ms step_avg:36.74ms step:574/1600 train_time:21112ms step_avg:36.78ms step:575/1600 train_time:21174ms step_avg:36.82ms step:576/1600 train_time:21233ms step_avg:36.86ms step:577/1600 train_time:21295ms step_avg:36.91ms step:578/1600 train_time:21354ms step_avg:36.95ms step:579/1600 train_time:21418ms step_avg:36.99ms step:580/1600 train_time:21476ms step_avg:37.03ms step:581/1600 train_time:21538ms step_avg:37.07ms step:582/1600 train_time:21597ms step_avg:37.11ms step:583/1600 train_time:21659ms step_avg:37.15ms step:584/1600 train_time:21719ms step_avg:37.19ms step:585/1600 train_time:21781ms step_avg:37.23ms step:586/1600 train_time:21841ms step_avg:37.27ms step:587/1600 train_time:21906ms step_avg:37.32ms step:588/1600 train_time:21964ms step_avg:37.35ms step:589/1600 train_time:22026ms step_avg:37.39ms step:590/1600 train_time:22084ms step_avg:37.43ms step:591/1600 train_time:22147ms step_avg:37.47ms step:592/1600 train_time:22206ms step_avg:37.51ms step:593/1600 train_time:22268ms step_avg:37.55ms step:594/1600 train_time:22327ms step_avg:37.59ms step:595/1600 train_time:22389ms step_avg:37.63ms step:596/1600 train_time:22448ms step_avg:37.66ms step:597/1600 train_time:22510ms step_avg:37.71ms step:598/1600 train_time:22569ms step_avg:37.74ms step:599/1600 train_time:22633ms step_avg:37.78ms step:600/1600 train_time:22691ms step_avg:37.82ms step:601/1600 train_time:22754ms step_avg:37.86ms step:602/1600 train_time:22813ms step_avg:37.90ms step:603/1600 train_time:22876ms step_avg:37.94ms step:604/1600 train_time:22935ms step_avg:37.97ms step:605/1600 train_time:22998ms step_avg:38.01ms step:606/1600 train_time:23057ms step_avg:38.05ms step:607/1600 train_time:23119ms step_avg:38.09ms step:608/1600 train_time:23178ms step_avg:38.12ms step:609/1600 train_time:23241ms step_avg:38.16ms step:610/1600 train_time:23301ms step_avg:38.20ms step:611/1600 train_time:23363ms step_avg:38.24ms step:612/1600 train_time:23424ms step_avg:38.27ms step:613/1600 train_time:23485ms step_avg:38.31ms step:614/1600 train_time:23544ms step_avg:38.35ms step:615/1600 train_time:23607ms step_avg:38.38ms step:616/1600 train_time:23665ms step_avg:38.42ms step:617/1600 train_time:23726ms step_avg:38.45ms step:618/1600 train_time:23785ms step_avg:38.49ms step:619/1600 train_time:23848ms step_avg:38.53ms step:620/1600 train_time:23906ms step_avg:38.56ms step:621/1600 train_time:23969ms step_avg:38.60ms step:622/1600 train_time:24029ms step_avg:38.63ms step:623/1600 train_time:24092ms step_avg:38.67ms step:624/1600 train_time:24151ms step_avg:38.70ms step:625/1600 train_time:24213ms step_avg:38.74ms step:626/1600 train_time:24273ms step_avg:38.77ms step:627/1600 train_time:24336ms step_avg:38.81ms step:628/1600 train_time:24395ms step_avg:38.84ms step:629/1600 train_time:24458ms step_avg:38.88ms step:630/1600 train_time:24517ms step_avg:38.92ms step:631/1600 train_time:24579ms step_avg:38.95ms step:632/1600 train_time:24638ms step_avg:38.98ms step:633/1600 train_time:24701ms step_avg:39.02ms step:634/1600 train_time:24760ms step_avg:39.05ms step:635/1600 train_time:24823ms step_avg:39.09ms step:636/1600 train_time:24882ms step_avg:39.12ms step:637/1600 train_time:24944ms step_avg:39.16ms step:638/1600 train_time:25004ms step_avg:39.19ms step:639/1600 train_time:25066ms step_avg:39.23ms step:640/1600 train_time:25125ms step_avg:39.26ms step:641/1600 train_time:25188ms step_avg:39.29ms step:642/1600 train_time:25246ms step_avg:39.32ms step:643/1600 train_time:25308ms step_avg:39.36ms step:644/1600 train_time:25367ms step_avg:39.39ms step:645/1600 train_time:25430ms step_avg:39.43ms step:646/1600 train_time:25489ms step_avg:39.46ms step:647/1600 train_time:25550ms step_avg:39.49ms step:648/1600 train_time:25609ms step_avg:39.52ms step:649/1600 train_time:25672ms step_avg:39.56ms step:650/1600 train_time:25734ms step_avg:39.59ms step:651/1600 train_time:25798ms step_avg:39.63ms step:652/1600 train_time:25854ms step_avg:39.65ms step:653/1600 train_time:25916ms step_avg:39.69ms step:654/1600 train_time:25975ms step_avg:39.72ms step:655/1600 train_time:26038ms step_avg:39.75ms step:656/1600 train_time:26097ms step_avg:39.78ms step:657/1600 train_time:26159ms step_avg:39.82ms step:658/1600 train_time:26219ms step_avg:39.85ms step:659/1600 train_time:26281ms step_avg:39.88ms step:660/1600 train_time:26340ms step_avg:39.91ms step:661/1600 train_time:26403ms step_avg:39.94ms step:662/1600 train_time:26462ms step_avg:39.97ms step:663/1600 train_time:26525ms step_avg:40.01ms step:664/1600 train_time:26584ms step_avg:40.04ms step:665/1600 train_time:26646ms step_avg:40.07ms step:666/1600 train_time:26706ms step_avg:40.10ms step:667/1600 train_time:26768ms step_avg:40.13ms step:668/1600 train_time:26827ms step_avg:40.16ms step:669/1600 train_time:26889ms step_avg:40.19ms step:670/1600 train_time:26947ms step_avg:40.22ms step:671/1600 train_time:27010ms step_avg:40.25ms step:672/1600 train_time:27071ms step_avg:40.28ms step:673/1600 train_time:27132ms step_avg:40.31ms step:674/1600 train_time:27192ms step_avg:40.34ms step:675/1600 train_time:27254ms step_avg:40.38ms step:676/1600 train_time:27313ms step_avg:40.40ms step:677/1600 train_time:27375ms step_avg:40.44ms step:678/1600 train_time:27434ms step_avg:40.46ms step:679/1600 train_time:27496ms step_avg:40.50ms step:680/1600 train_time:27556ms step_avg:40.52ms step:681/1600 train_time:27619ms step_avg:40.56ms step:682/1600 train_time:27677ms step_avg:40.58ms step:683/1600 train_time:27740ms step_avg:40.61ms step:684/1600 train_time:27799ms step_avg:40.64ms step:685/1600 train_time:27861ms step_avg:40.67ms step:686/1600 train_time:27921ms step_avg:40.70ms step:687/1600 train_time:27983ms step_avg:40.73ms step:688/1600 train_time:28043ms step_avg:40.76ms step:689/1600 train_time:28105ms step_avg:40.79ms step:690/1600 train_time:28164ms step_avg:40.82ms step:691/1600 train_time:28227ms step_avg:40.85ms step:692/1600 train_time:28286ms step_avg:40.88ms step:693/1600 train_time:28348ms step_avg:40.91ms step:694/1600 train_time:28407ms step_avg:40.93ms step:695/1600 train_time:28469ms step_avg:40.96ms step:696/1600 train_time:28528ms step_avg:40.99ms step:697/1600 train_time:28590ms step_avg:41.02ms step:698/1600 train_time:28649ms step_avg:41.05ms step:699/1600 train_time:28712ms step_avg:41.08ms step:700/1600 train_time:28771ms step_avg:41.10ms step:701/1600 train_time:28833ms step_avg:41.13ms step:702/1600 train_time:28892ms step_avg:41.16ms step:703/1600 train_time:28956ms step_avg:41.19ms step:704/1600 train_time:29014ms step_avg:41.21ms step:705/1600 train_time:29076ms step_avg:41.24ms step:706/1600 train_time:29135ms step_avg:41.27ms step:707/1600 train_time:29198ms step_avg:41.30ms step:708/1600 train_time:29257ms step_avg:41.32ms step:709/1600 train_time:29320ms step_avg:41.35ms step:710/1600 train_time:29378ms step_avg:41.38ms step:711/1600 train_time:29441ms step_avg:41.41ms step:712/1600 train_time:29500ms step_avg:41.43ms step:713/1600 train_time:29563ms step_avg:41.46ms step:714/1600 train_time:29622ms step_avg:41.49ms step:715/1600 train_time:29685ms step_avg:41.52ms step:716/1600 train_time:29745ms step_avg:41.54ms step:717/1600 train_time:29807ms step_avg:41.57ms step:718/1600 train_time:29866ms step_avg:41.60ms step:719/1600 train_time:29928ms step_avg:41.62ms step:720/1600 train_time:29986ms step_avg:41.65ms step:721/1600 train_time:30048ms step_avg:41.68ms step:722/1600 train_time:30107ms step_avg:41.70ms step:723/1600 train_time:30169ms step_avg:41.73ms step:724/1600 train_time:30228ms step_avg:41.75ms step:725/1600 train_time:30291ms step_avg:41.78ms step:726/1600 train_time:30351ms step_avg:41.81ms step:727/1600 train_time:30413ms step_avg:41.83ms step:728/1600 train_time:30472ms step_avg:41.86ms step:729/1600 train_time:30534ms step_avg:41.89ms step:730/1600 train_time:30594ms step_avg:41.91ms step:731/1600 train_time:30656ms step_avg:41.94ms step:732/1600 train_time:30716ms step_avg:41.96ms step:733/1600 train_time:30778ms step_avg:41.99ms step:734/1600 train_time:30837ms step_avg:42.01ms step:735/1600 train_time:30900ms step_avg:42.04ms step:736/1600 train_time:30959ms step_avg:42.06ms step:737/1600 train_time:31021ms step_avg:42.09ms step:738/1600 train_time:31081ms step_avg:42.11ms step:739/1600 train_time:31143ms step_avg:42.14ms step:740/1600 train_time:31202ms step_avg:42.17ms step:741/1600 train_time:31265ms step_avg:42.19ms step:742/1600 train_time:31324ms step_avg:42.22ms step:743/1600 train_time:31387ms step_avg:42.24ms step:744/1600 train_time:31446ms step_avg:42.27ms step:745/1600 train_time:31508ms step_avg:42.29ms step:746/1600 train_time:31566ms step_avg:42.31ms step:747/1600 train_time:31628ms step_avg:42.34ms step:748/1600 train_time:31687ms step_avg:42.36ms step:749/1600 train_time:31749ms step_avg:42.39ms step:750/1600 train_time:31808ms step_avg:42.41ms step:750/1600 val_loss:3.8944 train_time:31856ms step_avg:42.47ms step:751/1600 train_time:31876ms step_avg:42.44ms step:752/1600 train_time:31934ms step_avg:42.47ms step:753/1600 train_time:31999ms step_avg:42.49ms step:754/1600 train_time:32060ms step_avg:42.52ms step:755/1600 train_time:32122ms step_avg:42.55ms step:756/1600 train_time:32182ms step_avg:42.57ms step:757/1600 train_time:32244ms step_avg:42.59ms step:758/1600 train_time:32303ms step_avg:42.62ms step:759/1600 train_time:32364ms step_avg:42.64ms step:760/1600 train_time:32423ms step_avg:42.66ms step:761/1600 train_time:32487ms step_avg:42.69ms step:762/1600 train_time:32545ms step_avg:42.71ms step:763/1600 train_time:32606ms step_avg:42.73ms step:764/1600 train_time:32664ms step_avg:42.75ms step:765/1600 train_time:32725ms step_avg:42.78ms step:766/1600 train_time:32784ms step_avg:42.80ms step:767/1600 train_time:32848ms step_avg:42.83ms step:768/1600 train_time:32907ms step_avg:42.85ms step:769/1600 train_time:32972ms step_avg:42.88ms step:770/1600 train_time:33031ms step_avg:42.90ms step:771/1600 train_time:33094ms step_avg:42.92ms step:772/1600 train_time:33154ms step_avg:42.95ms step:773/1600 train_time:33216ms step_avg:42.97ms step:774/1600 train_time:33275ms step_avg:42.99ms step:775/1600 train_time:33337ms step_avg:43.02ms step:776/1600 train_time:33396ms step_avg:43.04ms step:777/1600 train_time:33458ms step_avg:43.06ms step:778/1600 train_time:33516ms step_avg:43.08ms step:779/1600 train_time:33578ms step_avg:43.10ms step:780/1600 train_time:33636ms step_avg:43.12ms step:781/1600 train_time:33698ms step_avg:43.15ms step:782/1600 train_time:33757ms step_avg:43.17ms step:783/1600 train_time:33819ms step_avg:43.19ms step:784/1600 train_time:33880ms step_avg:43.21ms step:785/1600 train_time:33944ms step_avg:43.24ms step:786/1600 train_time:34004ms step_avg:43.26ms step:787/1600 train_time:34067ms step_avg:43.29ms step:788/1600 train_time:34126ms step_avg:43.31ms step:789/1600 train_time:34188ms step_avg:43.33ms step:790/1600 train_time:34247ms step_avg:43.35ms step:791/1600 train_time:34310ms step_avg:43.38ms step:792/1600 train_time:34370ms step_avg:43.40ms step:793/1600 train_time:34432ms step_avg:43.42ms step:794/1600 train_time:34492ms step_avg:43.44ms step:795/1600 train_time:34556ms step_avg:43.47ms step:796/1600 train_time:34615ms step_avg:43.49ms step:797/1600 train_time:34676ms step_avg:43.51ms step:798/1600 train_time:34735ms step_avg:43.53ms step:799/1600 train_time:34797ms step_avg:43.55ms step:800/1600 train_time:34856ms step_avg:43.57ms step:801/1600 train_time:34920ms step_avg:43.59ms step:802/1600 train_time:34980ms step_avg:43.62ms step:803/1600 train_time:35040ms step_avg:43.64ms step:804/1600 train_time:35100ms step_avg:43.66ms step:805/1600 train_time:35162ms step_avg:43.68ms step:806/1600 train_time:35221ms step_avg:43.70ms step:807/1600 train_time:35283ms step_avg:43.72ms step:808/1600 train_time:35343ms step_avg:43.74ms step:809/1600 train_time:35405ms step_avg:43.76ms step:810/1600 train_time:35464ms step_avg:43.78ms step:811/1600 train_time:35527ms step_avg:43.81ms step:812/1600 train_time:35586ms step_avg:43.83ms step:813/1600 train_time:35649ms step_avg:43.85ms step:814/1600 train_time:35708ms step_avg:43.87ms step:815/1600 train_time:35770ms step_avg:43.89ms step:816/1600 train_time:35830ms step_avg:43.91ms step:817/1600 train_time:35892ms step_avg:43.93ms step:818/1600 train_time:35954ms step_avg:43.95ms step:819/1600 train_time:36016ms step_avg:43.98ms step:820/1600 train_time:36076ms step_avg:44.00ms step:821/1600 train_time:36139ms step_avg:44.02ms step:822/1600 train_time:36196ms step_avg:44.03ms step:823/1600 train_time:36257ms step_avg:44.06ms step:824/1600 train_time:36316ms step_avg:44.07ms step:825/1600 train_time:36379ms step_avg:44.10ms step:826/1600 train_time:36438ms step_avg:44.11ms step:827/1600 train_time:36500ms step_avg:44.14ms step:828/1600 train_time:36560ms step_avg:44.15ms step:829/1600 train_time:36622ms step_avg:44.18ms step:830/1600 train_time:36681ms step_avg:44.19ms step:831/1600 train_time:36744ms step_avg:44.22ms step:832/1600 train_time:36803ms step_avg:44.23ms step:833/1600 train_time:36867ms step_avg:44.26ms step:834/1600 train_time:36926ms step_avg:44.28ms step:835/1600 train_time:36989ms step_avg:44.30ms step:836/1600 train_time:37047ms step_avg:44.32ms step:837/1600 train_time:37111ms step_avg:44.34ms step:838/1600 train_time:37170ms step_avg:44.36ms step:839/1600 train_time:37233ms step_avg:44.38ms step:840/1600 train_time:37292ms step_avg:44.40ms step:841/1600 train_time:37355ms step_avg:44.42ms step:842/1600 train_time:37414ms step_avg:44.43ms step:843/1600 train_time:37477ms step_avg:44.46ms step:844/1600 train_time:37536ms step_avg:44.47ms step:845/1600 train_time:37598ms step_avg:44.49ms step:846/1600 train_time:37657ms step_avg:44.51ms step:847/1600 train_time:37719ms step_avg:44.53ms step:848/1600 train_time:37778ms step_avg:44.55ms step:849/1600 train_time:37841ms step_avg:44.57ms step:850/1600 train_time:37900ms step_avg:44.59ms step:851/1600 train_time:37963ms step_avg:44.61ms step:852/1600 train_time:38022ms step_avg:44.63ms step:853/1600 train_time:38085ms step_avg:44.65ms step:854/1600 train_time:38144ms step_avg:44.67ms step:855/1600 train_time:38208ms step_avg:44.69ms step:856/1600 train_time:38267ms step_avg:44.70ms step:857/1600 train_time:38330ms step_avg:44.73ms step:858/1600 train_time:38388ms step_avg:44.74ms step:859/1600 train_time:38451ms step_avg:44.76ms step:860/1600 train_time:38509ms step_avg:44.78ms step:861/1600 train_time:38572ms step_avg:44.80ms step:862/1600 train_time:38631ms step_avg:44.82ms step:863/1600 train_time:38694ms step_avg:44.84ms step:864/1600 train_time:38752ms step_avg:44.85ms step:865/1600 train_time:38815ms step_avg:44.87ms step:866/1600 train_time:38874ms step_avg:44.89ms step:867/1600 train_time:38938ms step_avg:44.91ms step:868/1600 train_time:38997ms step_avg:44.93ms step:869/1600 train_time:39059ms step_avg:44.95ms step:870/1600 train_time:39117ms step_avg:44.96ms step:871/1600 train_time:39180ms step_avg:44.98ms step:872/1600 train_time:39239ms step_avg:45.00ms step:873/1600 train_time:39302ms step_avg:45.02ms step:874/1600 train_time:39361ms step_avg:45.04ms step:875/1600 train_time:39424ms step_avg:45.06ms step:876/1600 train_time:39484ms step_avg:45.07ms step:877/1600 train_time:39546ms step_avg:45.09ms step:878/1600 train_time:39605ms step_avg:45.11ms step:879/1600 train_time:39667ms step_avg:45.13ms step:880/1600 train_time:39726ms step_avg:45.14ms step:881/1600 train_time:39788ms step_avg:45.16ms step:882/1600 train_time:39848ms step_avg:45.18ms step:883/1600 train_time:39910ms step_avg:45.20ms step:884/1600 train_time:39969ms step_avg:45.21ms step:885/1600 train_time:40034ms step_avg:45.24ms step:886/1600 train_time:40094ms step_avg:45.25ms step:887/1600 train_time:40154ms step_avg:45.27ms step:888/1600 train_time:40213ms step_avg:45.28ms step:889/1600 train_time:40276ms step_avg:45.30ms step:890/1600 train_time:40341ms step_avg:45.33ms step:891/1600 train_time:40399ms step_avg:45.34ms step:892/1600 train_time:40457ms step_avg:45.36ms step:893/1600 train_time:40520ms step_avg:45.37ms step:894/1600 train_time:40577ms step_avg:45.39ms step:895/1600 train_time:40639ms step_avg:45.41ms step:896/1600 train_time:40698ms step_avg:45.42ms step:897/1600 train_time:40760ms step_avg:45.44ms step:898/1600 train_time:40819ms step_avg:45.46ms step:899/1600 train_time:40882ms step_avg:45.48ms step:900/1600 train_time:40944ms step_avg:45.49ms step:901/1600 train_time:41004ms step_avg:45.51ms step:902/1600 train_time:41064ms step_avg:45.53ms step:903/1600 train_time:41126ms step_avg:45.54ms step:904/1600 train_time:41185ms step_avg:45.56ms step:905/1600 train_time:41248ms step_avg:45.58ms step:906/1600 train_time:41307ms step_avg:45.59ms step:907/1600 train_time:41369ms step_avg:45.61ms step:908/1600 train_time:41428ms step_avg:45.63ms step:909/1600 train_time:41491ms step_avg:45.64ms step:910/1600 train_time:41551ms step_avg:45.66ms step:911/1600 train_time:41613ms step_avg:45.68ms step:912/1600 train_time:41672ms step_avg:45.69ms step:913/1600 train_time:41735ms step_avg:45.71ms step:914/1600 train_time:41793ms step_avg:45.73ms step:915/1600 train_time:41856ms step_avg:45.74ms step:916/1600 train_time:41915ms step_avg:45.76ms step:917/1600 train_time:41977ms step_avg:45.78ms step:918/1600 train_time:42036ms step_avg:45.79ms step:919/1600 train_time:42099ms step_avg:45.81ms step:920/1600 train_time:42157ms step_avg:45.82ms step:921/1600 train_time:42219ms step_avg:45.84ms step:922/1600 train_time:42278ms step_avg:45.85ms step:923/1600 train_time:42341ms step_avg:45.87ms step:924/1600 train_time:42401ms step_avg:45.89ms step:925/1600 train_time:42463ms step_avg:45.91ms step:926/1600 train_time:42522ms step_avg:45.92ms step:927/1600 train_time:42584ms step_avg:45.94ms step:928/1600 train_time:42643ms step_avg:45.95ms step:929/1600 train_time:42707ms step_avg:45.97ms step:930/1600 train_time:42766ms step_avg:45.98ms step:931/1600 train_time:42828ms step_avg:46.00ms step:932/1600 train_time:42887ms step_avg:46.02ms step:933/1600 train_time:42950ms step_avg:46.03ms step:934/1600 train_time:43009ms step_avg:46.05ms step:935/1600 train_time:43071ms step_avg:46.06ms step:936/1600 train_time:43129ms step_avg:46.08ms step:937/1600 train_time:43193ms step_avg:46.10ms step:938/1600 train_time:43250ms step_avg:46.11ms step:939/1600 train_time:43313ms step_avg:46.13ms step:940/1600 train_time:43372ms step_avg:46.14ms step:941/1600 train_time:43435ms step_avg:46.16ms step:942/1600 train_time:43494ms step_avg:46.17ms step:943/1600 train_time:43555ms step_avg:46.19ms step:944/1600 train_time:43614ms step_avg:46.20ms step:945/1600 train_time:43677ms step_avg:46.22ms step:946/1600 train_time:43736ms step_avg:46.23ms step:947/1600 train_time:43798ms step_avg:46.25ms step:948/1600 train_time:43857ms step_avg:46.26ms step:949/1600 train_time:43920ms step_avg:46.28ms step:950/1600 train_time:43979ms step_avg:46.29ms step:951/1600 train_time:44042ms step_avg:46.31ms step:952/1600 train_time:44101ms step_avg:46.32ms step:953/1600 train_time:44163ms step_avg:46.34ms step:954/1600 train_time:44222ms step_avg:46.35ms step:955/1600 train_time:44284ms step_avg:46.37ms step:956/1600 train_time:44343ms step_avg:46.38ms step:957/1600 train_time:44406ms step_avg:46.40ms step:958/1600 train_time:44465ms step_avg:46.41ms step:959/1600 train_time:44527ms step_avg:46.43ms step:960/1600 train_time:44586ms step_avg:46.44ms step:961/1600 train_time:44648ms step_avg:46.46ms step:962/1600 train_time:44708ms step_avg:46.47ms step:963/1600 train_time:44770ms step_avg:46.49ms step:964/1600 train_time:44829ms step_avg:46.50ms step:965/1600 train_time:44892ms step_avg:46.52ms step:966/1600 train_time:44952ms step_avg:46.53ms step:967/1600 train_time:45015ms step_avg:46.55ms step:968/1600 train_time:45074ms step_avg:46.56ms step:969/1600 train_time:45139ms step_avg:46.58ms step:970/1600 train_time:45197ms step_avg:46.59ms step:971/1600 train_time:45256ms step_avg:46.61ms step:972/1600 train_time:45315ms step_avg:46.62ms step:973/1600 train_time:45378ms step_avg:46.64ms step:974/1600 train_time:45437ms step_avg:46.65ms step:975/1600 train_time:45500ms step_avg:46.67ms step:976/1600 train_time:45559ms step_avg:46.68ms step:977/1600 train_time:45621ms step_avg:46.70ms step:978/1600 train_time:45680ms step_avg:46.71ms step:979/1600 train_time:45742ms step_avg:46.72ms step:980/1600 train_time:45802ms step_avg:46.74ms step:981/1600 train_time:45865ms step_avg:46.75ms step:982/1600 train_time:45924ms step_avg:46.77ms step:983/1600 train_time:45987ms step_avg:46.78ms step:984/1600 train_time:46045ms step_avg:46.79ms step:985/1600 train_time:46108ms step_avg:46.81ms step:986/1600 train_time:46167ms step_avg:46.82ms step:987/1600 train_time:46229ms step_avg:46.84ms step:988/1600 train_time:46288ms step_avg:46.85ms step:989/1600 train_time:46351ms step_avg:46.87ms step:990/1600 train_time:46411ms step_avg:46.88ms step:991/1600 train_time:46474ms step_avg:46.90ms step:992/1600 train_time:46534ms step_avg:46.91ms step:993/1600 train_time:46596ms step_avg:46.92ms step:994/1600 train_time:46654ms step_avg:46.94ms step:995/1600 train_time:46717ms step_avg:46.95ms step:996/1600 train_time:46775ms step_avg:46.96ms step:997/1600 train_time:46837ms step_avg:46.98ms step:998/1600 train_time:46897ms step_avg:46.99ms step:999/1600 train_time:46960ms step_avg:47.01ms step:1000/1600 train_time:47019ms step_avg:47.02ms step:1000/1600 val_loss:3.5962 train_time:47066ms step_avg:47.07ms step:1001/1600 train_time:47085ms step_avg:47.04ms step:1002/1600 train_time:47144ms step_avg:47.05ms step:1003/1600 train_time:47208ms step_avg:47.07ms step:1004/1600 train_time:47270ms step_avg:47.08ms step:1005/1600 train_time:47332ms step_avg:47.10ms step:1006/1600 train_time:47390ms step_avg:47.11ms step:1007/1600 train_time:47452ms step_avg:47.12ms step:1008/1600 train_time:47513ms step_avg:47.14ms step:1009/1600 train_time:47574ms step_avg:47.15ms step:1010/1600 train_time:47632ms step_avg:47.16ms step:1011/1600 train_time:47694ms step_avg:47.18ms step:1012/1600 train_time:47752ms step_avg:47.19ms step:1013/1600 train_time:47814ms step_avg:47.20ms step:1014/1600 train_time:47874ms step_avg:47.21ms step:1015/1600 train_time:47936ms step_avg:47.23ms step:1016/1600 train_time:47995ms step_avg:47.24ms step:1017/1600 train_time:48058ms step_avg:47.25ms step:1018/1600 train_time:48120ms step_avg:47.27ms step:1019/1600 train_time:48183ms step_avg:47.28ms step:1020/1600 train_time:48243ms step_avg:47.30ms step:1021/1600 train_time:48306ms step_avg:47.31ms step:1022/1600 train_time:48365ms step_avg:47.32ms step:1023/1600 train_time:48428ms step_avg:47.34ms step:1024/1600 train_time:48486ms step_avg:47.35ms step:1025/1600 train_time:48548ms step_avg:47.36ms step:1026/1600 train_time:48606ms step_avg:47.37ms step:1027/1600 train_time:48667ms step_avg:47.39ms step:1028/1600 train_time:48726ms step_avg:47.40ms step:1029/1600 train_time:48788ms step_avg:47.41ms step:1030/1600 train_time:48846ms step_avg:47.42ms step:1031/1600 train_time:48908ms step_avg:47.44ms step:1032/1600 train_time:48966ms step_avg:47.45ms step:1033/1600 train_time:49030ms step_avg:47.46ms step:1034/1600 train_time:49090ms step_avg:47.48ms step:1035/1600 train_time:49153ms step_avg:47.49ms step:1036/1600 train_time:49214ms step_avg:47.50ms step:1037/1600 train_time:49277ms step_avg:47.52ms step:1038/1600 train_time:49336ms step_avg:47.53ms step:1039/1600 train_time:49401ms step_avg:47.55ms step:1040/1600 train_time:49458ms step_avg:47.56ms step:1041/1600 train_time:49528ms step_avg:47.58ms step:1042/1600 train_time:49612ms step_avg:47.61ms step:1043/1600 train_time:49700ms step_avg:47.65ms step:1044/1600 train_time:49785ms step_avg:47.69ms step:1045/1600 train_time:49873ms step_avg:47.73ms step:1046/1600 train_time:49958ms step_avg:47.76ms step:1047/1600 train_time:50047ms step_avg:47.80ms step:1048/1600 train_time:50134ms step_avg:47.84ms step:1049/1600 train_time:50226ms step_avg:47.88ms step:1050/1600 train_time:50308ms step_avg:47.91ms step:1051/1600 train_time:50396ms step_avg:47.95ms step:1052/1600 train_time:50482ms step_avg:47.99ms step:1053/1600 train_time:50571ms step_avg:48.03ms step:1054/1600 train_time:50655ms step_avg:48.06ms step:1055/1600 train_time:50743ms step_avg:48.10ms step:1056/1600 train_time:50828ms step_avg:48.13ms step:1057/1600 train_time:50917ms step_avg:48.17ms step:1058/1600 train_time:51001ms step_avg:48.20ms step:1059/1600 train_time:51089ms step_avg:48.24ms step:1060/1600 train_time:51175ms step_avg:48.28ms step:1061/1600 train_time:51263ms step_avg:48.32ms step:1062/1600 train_time:51348ms step_avg:48.35ms step:1063/1600 train_time:51436ms step_avg:48.39ms step:1064/1600 train_time:51521ms step_avg:48.42ms step:1065/1600 train_time:51609ms step_avg:48.46ms step:1066/1600 train_time:51694ms step_avg:48.49ms step:1067/1600 train_time:51781ms step_avg:48.53ms step:1068/1600 train_time:51866ms step_avg:48.56ms step:1069/1600 train_time:51955ms step_avg:48.60ms step:1070/1600 train_time:52038ms step_avg:48.63ms step:1071/1600 train_time:52129ms step_avg:48.67ms step:1072/1600 train_time:52212ms step_avg:48.71ms step:1073/1600 train_time:52301ms step_avg:48.74ms step:1074/1600 train_time:52387ms step_avg:48.78ms step:1075/1600 train_time:52475ms step_avg:48.81ms step:1076/1600 train_time:52559ms step_avg:48.85ms step:1077/1600 train_time:52648ms step_avg:48.88ms step:1078/1600 train_time:52733ms step_avg:48.92ms step:1079/1600 train_time:52820ms step_avg:48.95ms step:1080/1600 train_time:52906ms step_avg:48.99ms step:1081/1600 train_time:52995ms step_avg:49.02ms step:1082/1600 train_time:53079ms step_avg:49.06ms step:1083/1600 train_time:53171ms step_avg:49.10ms step:1084/1600 train_time:53254ms step_avg:49.13ms step:1085/1600 train_time:53342ms step_avg:49.16ms step:1086/1600 train_time:53428ms step_avg:49.20ms step:1087/1600 train_time:53516ms step_avg:49.23ms step:1088/1600 train_time:53600ms step_avg:49.26ms step:1089/1600 train_time:53689ms step_avg:49.30ms step:1090/1600 train_time:53774ms step_avg:49.33ms step:1091/1600 train_time:53862ms step_avg:49.37ms step:1092/1600 train_time:53948ms step_avg:49.40ms step:1093/1600 train_time:54036ms step_avg:49.44ms step:1094/1600 train_time:54121ms step_avg:49.47ms step:1095/1600 train_time:54209ms step_avg:49.51ms step:1096/1600 train_time:54295ms step_avg:49.54ms step:1097/1600 train_time:54385ms step_avg:49.58ms step:1098/1600 train_time:54471ms step_avg:49.61ms step:1099/1600 train_time:54560ms step_avg:49.65ms step:1100/1600 train_time:54643ms step_avg:49.68ms step:1101/1600 train_time:54731ms step_avg:49.71ms step:1102/1600 train_time:54817ms step_avg:49.74ms step:1103/1600 train_time:54905ms step_avg:49.78ms step:1104/1600 train_time:54991ms step_avg:49.81ms step:1105/1600 train_time:55079ms step_avg:49.84ms step:1106/1600 train_time:55165ms step_avg:49.88ms step:1107/1600 train_time:55253ms step_avg:49.91ms step:1108/1600 train_time:55338ms step_avg:49.94ms step:1109/1600 train_time:55428ms step_avg:49.98ms step:1110/1600 train_time:55512ms step_avg:50.01ms step:1111/1600 train_time:55599ms step_avg:50.04ms step:1112/1600 train_time:55686ms step_avg:50.08ms step:1113/1600 train_time:55774ms step_avg:50.11ms step:1114/1600 train_time:55859ms step_avg:50.14ms step:1115/1600 train_time:55948ms step_avg:50.18ms step:1116/1600 train_time:56033ms step_avg:50.21ms step:1117/1600 train_time:56121ms step_avg:50.24ms step:1118/1600 train_time:56207ms step_avg:50.27ms step:1119/1600 train_time:56295ms step_avg:50.31ms step:1120/1600 train_time:56380ms step_avg:50.34ms step:1121/1600 train_time:56467ms step_avg:50.37ms step:1122/1600 train_time:56553ms step_avg:50.40ms step:1123/1600 train_time:56640ms step_avg:50.44ms step:1124/1600 train_time:56726ms step_avg:50.47ms step:1125/1600 train_time:56813ms step_avg:50.50ms step:1126/1600 train_time:56898ms step_avg:50.53ms step:1127/1600 train_time:56987ms step_avg:50.57ms step:1128/1600 train_time:57073ms step_avg:50.60ms step:1129/1600 train_time:57161ms step_avg:50.63ms step:1130/1600 train_time:57246ms step_avg:50.66ms step:1131/1600 train_time:57334ms step_avg:50.69ms step:1132/1600 train_time:57420ms step_avg:50.72ms step:1133/1600 train_time:57507ms step_avg:50.76ms step:1134/1600 train_time:57593ms step_avg:50.79ms step:1135/1600 train_time:57680ms step_avg:50.82ms step:1136/1600 train_time:57766ms step_avg:50.85ms step:1137/1600 train_time:57854ms step_avg:50.88ms step:1138/1600 train_time:57942ms step_avg:50.92ms step:1139/1600 train_time:58032ms step_avg:50.95ms step:1140/1600 train_time:58114ms step_avg:50.98ms step:1141/1600 train_time:58201ms step_avg:51.01ms step:1142/1600 train_time:58286ms step_avg:51.04ms step:1143/1600 train_time:58373ms step_avg:51.07ms step:1144/1600 train_time:58458ms step_avg:51.10ms step:1145/1600 train_time:58546ms step_avg:51.13ms step:1146/1600 train_time:58631ms step_avg:51.16ms step:1147/1600 train_time:58719ms step_avg:51.19ms step:1148/1600 train_time:58803ms step_avg:51.22ms step:1149/1600 train_time:58891ms step_avg:51.25ms step:1150/1600 train_time:58976ms step_avg:51.28ms step:1151/1600 train_time:59065ms step_avg:51.32ms step:1152/1600 train_time:59149ms step_avg:51.34ms step:1153/1600 train_time:59238ms step_avg:51.38ms step:1154/1600 train_time:59327ms step_avg:51.41ms step:1155/1600 train_time:59412ms step_avg:51.44ms step:1156/1600 train_time:59497ms step_avg:51.47ms step:1157/1600 train_time:59585ms step_avg:51.50ms step:1158/1600 train_time:59670ms step_avg:51.53ms step:1159/1600 train_time:59758ms step_avg:51.56ms step:1160/1600 train_time:59843ms step_avg:51.59ms step:1161/1600 train_time:59932ms step_avg:51.62ms step:1162/1600 train_time:60017ms step_avg:51.65ms step:1163/1600 train_time:60105ms step_avg:51.68ms step:1164/1600 train_time:60191ms step_avg:51.71ms step:1165/1600 train_time:60279ms step_avg:51.74ms step:1166/1600 train_time:60364ms step_avg:51.77ms step:1167/1600 train_time:60453ms step_avg:51.80ms step:1168/1600 train_time:60537ms step_avg:51.83ms step:1169/1600 train_time:60625ms step_avg:51.86ms step:1170/1600 train_time:60709ms step_avg:51.89ms step:1171/1600 train_time:60798ms step_avg:51.92ms step:1172/1600 train_time:60883ms step_avg:51.95ms step:1173/1600 train_time:60971ms step_avg:51.98ms step:1174/1600 train_time:61056ms step_avg:52.01ms step:1175/1600 train_time:61145ms step_avg:52.04ms step:1176/1600 train_time:61231ms step_avg:52.07ms step:1177/1600 train_time:61319ms step_avg:52.10ms step:1178/1600 train_time:61404ms step_avg:52.13ms step:1179/1600 train_time:61492ms step_avg:52.16ms step:1180/1600 train_time:61576ms step_avg:52.18ms step:1181/1600 train_time:61665ms step_avg:52.21ms step:1182/1600 train_time:61750ms step_avg:52.24ms step:1183/1600 train_time:61838ms step_avg:52.27ms step:1184/1600 train_time:61924ms step_avg:52.30ms step:1185/1600 train_time:62014ms step_avg:52.33ms step:1186/1600 train_time:62099ms step_avg:52.36ms step:1187/1600 train_time:62185ms step_avg:52.39ms step:1188/1600 train_time:62270ms step_avg:52.42ms step:1189/1600 train_time:62358ms step_avg:52.45ms step:1190/1600 train_time:62443ms step_avg:52.47ms step:1191/1600 train_time:62531ms step_avg:52.50ms step:1192/1600 train_time:62617ms step_avg:52.53ms step:1193/1600 train_time:62704ms step_avg:52.56ms step:1194/1600 train_time:62790ms step_avg:52.59ms step:1195/1600 train_time:62878ms step_avg:52.62ms step:1196/1600 train_time:62964ms step_avg:52.65ms step:1197/1600 train_time:63052ms step_avg:52.68ms step:1198/1600 train_time:63137ms step_avg:52.70ms step:1199/1600 train_time:63224ms step_avg:52.73ms step:1200/1600 train_time:63309ms step_avg:52.76ms step:1201/1600 train_time:63398ms step_avg:52.79ms step:1202/1600 train_time:63483ms step_avg:52.81ms step:1203/1600 train_time:63572ms step_avg:52.84ms step:1204/1600 train_time:63656ms step_avg:52.87ms step:1205/1600 train_time:63743ms step_avg:52.90ms step:1206/1600 train_time:63829ms step_avg:52.93ms step:1207/1600 train_time:63916ms step_avg:52.95ms step:1208/1600 train_time:64001ms step_avg:52.98ms step:1209/1600 train_time:64090ms step_avg:53.01ms step:1210/1600 train_time:64175ms step_avg:53.04ms step:1211/1600 train_time:64263ms step_avg:53.07ms step:1212/1600 train_time:64348ms step_avg:53.09ms step:1213/1600 train_time:64436ms step_avg:53.12ms step:1214/1600 train_time:64521ms step_avg:53.15ms step:1215/1600 train_time:64610ms step_avg:53.18ms step:1216/1600 train_time:64695ms step_avg:53.20ms step:1217/1600 train_time:64784ms step_avg:53.23ms step:1218/1600 train_time:64869ms step_avg:53.26ms step:1219/1600 train_time:64958ms step_avg:53.29ms step:1220/1600 train_time:65042ms step_avg:53.31ms step:1221/1600 train_time:65131ms step_avg:53.34ms step:1222/1600 train_time:65217ms step_avg:53.37ms step:1223/1600 train_time:65305ms step_avg:53.40ms step:1224/1600 train_time:65390ms step_avg:53.42ms step:1225/1600 train_time:65477ms step_avg:53.45ms step:1226/1600 train_time:65562ms step_avg:53.48ms step:1227/1600 train_time:65650ms step_avg:53.50ms step:1228/1600 train_time:65735ms step_avg:53.53ms step:1229/1600 train_time:65823ms step_avg:53.56ms step:1230/1600 train_time:65908ms step_avg:53.58ms step:1231/1600 train_time:65997ms step_avg:53.61ms step:1232/1600 train_time:66082ms step_avg:53.64ms step:1233/1600 train_time:66171ms step_avg:53.67ms step:1234/1600 train_time:66256ms step_avg:53.69ms step:1235/1600 train_time:66343ms step_avg:53.72ms step:1236/1600 train_time:66430ms step_avg:53.75ms step:1237/1600 train_time:66517ms step_avg:53.77ms step:1238/1600 train_time:66602ms step_avg:53.80ms step:1239/1600 train_time:66691ms step_avg:53.83ms step:1240/1600 train_time:66775ms step_avg:53.85ms step:1241/1600 train_time:66863ms step_avg:53.88ms step:1242/1600 train_time:66948ms step_avg:53.90ms step:1243/1600 train_time:67036ms step_avg:53.93ms step:1244/1600 train_time:67121ms step_avg:53.96ms step:1245/1600 train_time:67211ms step_avg:53.98ms step:1246/1600 train_time:67296ms step_avg:54.01ms step:1247/1600 train_time:67384ms step_avg:54.04ms step:1248/1600 train_time:67469ms step_avg:54.06ms step:1249/1600 train_time:67557ms step_avg:54.09ms step:1250/1600 train_time:67642ms step_avg:54.11ms step:1250/1600 val_loss:3.4146 train_time:67715ms step_avg:54.17ms step:1251/1600 train_time:67734ms step_avg:54.14ms step:1252/1600 train_time:67819ms step_avg:54.17ms step:1253/1600 train_time:67908ms step_avg:54.20ms step:1254/1600 train_time:67993ms step_avg:54.22ms step:1255/1600 train_time:68082ms step_avg:54.25ms step:1256/1600 train_time:68166ms step_avg:54.27ms step:1257/1600 train_time:68254ms step_avg:54.30ms step:1258/1600 train_time:68339ms step_avg:54.32ms step:1259/1600 train_time:68426ms step_avg:54.35ms step:1260/1600 train_time:68510ms step_avg:54.37ms step:1261/1600 train_time:68599ms step_avg:54.40ms step:1262/1600 train_time:68685ms step_avg:54.43ms step:1263/1600 train_time:68776ms step_avg:54.45ms step:1264/1600 train_time:68862ms step_avg:54.48ms step:1265/1600 train_time:68952ms step_avg:54.51ms step:1266/1600 train_time:69037ms step_avg:54.53ms step:1267/1600 train_time:69126ms step_avg:54.56ms step:1268/1600 train_time:69210ms step_avg:54.58ms step:1269/1600 train_time:69296ms step_avg:54.61ms step:1270/1600 train_time:69381ms step_avg:54.63ms step:1271/1600 train_time:69468ms step_avg:54.66ms step:1272/1600 train_time:69553ms step_avg:54.68ms step:1273/1600 train_time:69642ms step_avg:54.71ms step:1274/1600 train_time:69729ms step_avg:54.73ms step:1275/1600 train_time:69817ms step_avg:54.76ms step:1276/1600 train_time:69903ms step_avg:54.78ms step:1277/1600 train_time:69992ms step_avg:54.81ms step:1278/1600 train_time:70076ms step_avg:54.83ms step:1279/1600 train_time:70164ms step_avg:54.86ms step:1280/1600 train_time:70249ms step_avg:54.88ms step:1281/1600 train_time:70337ms step_avg:54.91ms step:1282/1600 train_time:70421ms step_avg:54.93ms step:1283/1600 train_time:70509ms step_avg:54.96ms step:1284/1600 train_time:70594ms step_avg:54.98ms step:1285/1600 train_time:70684ms step_avg:55.01ms step:1286/1600 train_time:70770ms step_avg:55.03ms step:1287/1600 train_time:70859ms step_avg:55.06ms step:1288/1600 train_time:70945ms step_avg:55.08ms step:1289/1600 train_time:71032ms step_avg:55.11ms step:1290/1600 train_time:71117ms step_avg:55.13ms step:1291/1600 train_time:71205ms step_avg:55.16ms step:1292/1600 train_time:71290ms step_avg:55.18ms step:1293/1600 train_time:71379ms step_avg:55.20ms step:1294/1600 train_time:71463ms step_avg:55.23ms step:1295/1600 train_time:71551ms step_avg:55.25ms step:1296/1600 train_time:71637ms step_avg:55.28ms step:1297/1600 train_time:71725ms step_avg:55.30ms step:1298/1600 train_time:71811ms step_avg:55.32ms step:1299/1600 train_time:71900ms step_avg:55.35ms step:1300/1600 train_time:71985ms step_avg:55.37ms step:1301/1600 train_time:72072ms step_avg:55.40ms step:1302/1600 train_time:72157ms step_avg:55.42ms step:1303/1600 train_time:72245ms step_avg:55.45ms step:1304/1600 train_time:72329ms step_avg:55.47ms step:1305/1600 train_time:72417ms step_avg:55.49ms step:1306/1600 train_time:72503ms step_avg:55.51ms step:1307/1600 train_time:72591ms step_avg:55.54ms step:1308/1600 train_time:72677ms step_avg:55.56ms step:1309/1600 train_time:72765ms step_avg:55.59ms step:1310/1600 train_time:72851ms step_avg:55.61ms step:1311/1600 train_time:72940ms step_avg:55.64ms step:1312/1600 train_time:73025ms step_avg:55.66ms step:1313/1600 train_time:73113ms step_avg:55.68ms step:1314/1600 train_time:73198ms step_avg:55.71ms step:1315/1600 train_time:73286ms step_avg:55.73ms step:1316/1600 train_time:73370ms step_avg:55.75ms step:1317/1600 train_time:73457ms step_avg:55.78ms step:1318/1600 train_time:73543ms step_avg:55.80ms step:1319/1600 train_time:73632ms step_avg:55.82ms step:1320/1600 train_time:73719ms step_avg:55.85ms step:1321/1600 train_time:73807ms step_avg:55.87ms step:1322/1600 train_time:73892ms step_avg:55.89ms step:1323/1600 train_time:73981ms step_avg:55.92ms step:1324/1600 train_time:74067ms step_avg:55.94ms step:1325/1600 train_time:74155ms step_avg:55.97ms step:1326/1600 train_time:74240ms step_avg:55.99ms step:1327/1600 train_time:74328ms step_avg:56.01ms step:1328/1600 train_time:74412ms step_avg:56.03ms step:1329/1600 train_time:74500ms step_avg:56.06ms step:1330/1600 train_time:74585ms step_avg:56.08ms step:1331/1600 train_time:74673ms step_avg:56.10ms step:1332/1600 train_time:74759ms step_avg:56.13ms step:1333/1600 train_time:74848ms step_avg:56.15ms step:1334/1600 train_time:74934ms step_avg:56.17ms step:1335/1600 train_time:75022ms step_avg:56.20ms step:1336/1600 train_time:75107ms step_avg:56.22ms step:1337/1600 train_time:75195ms step_avg:56.24ms step:1338/1600 train_time:75280ms step_avg:56.26ms step:1339/1600 train_time:75368ms step_avg:56.29ms step:1340/1600 train_time:75453ms step_avg:56.31ms step:1341/1600 train_time:75541ms step_avg:56.33ms step:1342/1600 train_time:75626ms step_avg:56.35ms step:1343/1600 train_time:75715ms step_avg:56.38ms step:1344/1600 train_time:75800ms step_avg:56.40ms step:1345/1600 train_time:75888ms step_avg:56.42ms step:1346/1600 train_time:75973ms step_avg:56.44ms step:1347/1600 train_time:76062ms step_avg:56.47ms step:1348/1600 train_time:76146ms step_avg:56.49ms step:1349/1600 train_time:76234ms step_avg:56.51ms step:1350/1600 train_time:76319ms step_avg:56.53ms step:1351/1600 train_time:76408ms step_avg:56.56ms step:1352/1600 train_time:76492ms step_avg:56.58ms step:1353/1600 train_time:76580ms step_avg:56.60ms step:1354/1600 train_time:76666ms step_avg:56.62ms step:1355/1600 train_time:76754ms step_avg:56.64ms step:1356/1600 train_time:76841ms step_avg:56.67ms step:1357/1600 train_time:76928ms step_avg:56.69ms step:1358/1600 train_time:77013ms step_avg:56.71ms step:1359/1600 train_time:77102ms step_avg:56.73ms step:1360/1600 train_time:77186ms step_avg:56.75ms step:1361/1600 train_time:77275ms step_avg:56.78ms step:1362/1600 train_time:77360ms step_avg:56.80ms step:1363/1600 train_time:77448ms step_avg:56.82ms step:1364/1600 train_time:77532ms step_avg:56.84ms step:1365/1600 train_time:77622ms step_avg:56.87ms step:1366/1600 train_time:77712ms step_avg:56.89ms step:1367/1600 train_time:77798ms step_avg:56.91ms step:1368/1600 train_time:77883ms step_avg:56.93ms step:1369/1600 train_time:77970ms step_avg:56.95ms step:1370/1600 train_time:78056ms step_avg:56.98ms step:1371/1600 train_time:78144ms step_avg:57.00ms step:1372/1600 train_time:78229ms step_avg:57.02ms step:1373/1600 train_time:78318ms step_avg:57.04ms step:1374/1600 train_time:78403ms step_avg:57.06ms step:1375/1600 train_time:78490ms step_avg:57.08ms step:1376/1600 train_time:78576ms step_avg:57.10ms step:1377/1600 train_time:78665ms step_avg:57.13ms step:1378/1600 train_time:78750ms step_avg:57.15ms step:1379/1600 train_time:78838ms step_avg:57.17ms step:1380/1600 train_time:78923ms step_avg:57.19ms step:1381/1600 train_time:79011ms step_avg:57.21ms step:1382/1600 train_time:79097ms step_avg:57.23ms step:1383/1600 train_time:79186ms step_avg:57.26ms step:1384/1600 train_time:79271ms step_avg:57.28ms step:1385/1600 train_time:79360ms step_avg:57.30ms step:1386/1600 train_time:79445ms step_avg:57.32ms step:1387/1600 train_time:79532ms step_avg:57.34ms step:1388/1600 train_time:79617ms step_avg:57.36ms step:1389/1600 train_time:79706ms step_avg:57.38ms step:1390/1600 train_time:79791ms step_avg:57.40ms step:1391/1600 train_time:79879ms step_avg:57.43ms step:1392/1600 train_time:79964ms step_avg:57.45ms step:1393/1600 train_time:80052ms step_avg:57.47ms step:1394/1600 train_time:80138ms step_avg:57.49ms step:1395/1600 train_time:80227ms step_avg:57.51ms step:1396/1600 train_time:80313ms step_avg:57.53ms step:1397/1600 train_time:80401ms step_avg:57.55ms step:1398/1600 train_time:80486ms step_avg:57.57ms step:1399/1600 train_time:80575ms step_avg:57.59ms step:1400/1600 train_time:80660ms step_avg:57.61ms step:1401/1600 train_time:80748ms step_avg:57.64ms step:1402/1600 train_time:80833ms step_avg:57.66ms step:1403/1600 train_time:80922ms step_avg:57.68ms step:1404/1600 train_time:81007ms step_avg:57.70ms step:1405/1600 train_time:81094ms step_avg:57.72ms step:1406/1600 train_time:81180ms step_avg:57.74ms step:1407/1600 train_time:81268ms step_avg:57.76ms step:1408/1600 train_time:81353ms step_avg:57.78ms step:1409/1600 train_time:81441ms step_avg:57.80ms step:1410/1600 train_time:81527ms step_avg:57.82ms step:1411/1600 train_time:81615ms step_avg:57.84ms step:1412/1600 train_time:81701ms step_avg:57.86ms step:1413/1600 train_time:81789ms step_avg:57.88ms step:1414/1600 train_time:81874ms step_avg:57.90ms step:1415/1600 train_time:81962ms step_avg:57.92ms step:1416/1600 train_time:82047ms step_avg:57.94ms step:1417/1600 train_time:82135ms step_avg:57.96ms step:1418/1600 train_time:82220ms step_avg:57.98ms step:1419/1600 train_time:82309ms step_avg:58.00ms step:1420/1600 train_time:82394ms step_avg:58.02ms step:1421/1600 train_time:82482ms step_avg:58.04ms step:1422/1600 train_time:82567ms step_avg:58.06ms step:1423/1600 train_time:82656ms step_avg:58.09ms step:1424/1600 train_time:82741ms step_avg:58.10ms step:1425/1600 train_time:82829ms step_avg:58.13ms step:1426/1600 train_time:82915ms step_avg:58.14ms step:1427/1600 train_time:83003ms step_avg:58.17ms step:1428/1600 train_time:83087ms step_avg:58.18ms step:1429/1600 train_time:83176ms step_avg:58.21ms step:1430/1600 train_time:83262ms step_avg:58.22ms step:1431/1600 train_time:83350ms step_avg:58.25ms step:1432/1600 train_time:83435ms step_avg:58.26ms step:1433/1600 train_time:83524ms step_avg:58.29ms step:1434/1600 train_time:83608ms step_avg:58.30ms step:1435/1600 train_time:83696ms step_avg:58.32ms step:1436/1600 train_time:83781ms step_avg:58.34ms step:1437/1600 train_time:83869ms step_avg:58.36ms step:1438/1600 train_time:83954ms step_avg:58.38ms step:1439/1600 train_time:84042ms step_avg:58.40ms step:1440/1600 train_time:84128ms step_avg:58.42ms step:1441/1600 train_time:84216ms step_avg:58.44ms step:1442/1600 train_time:84301ms step_avg:58.46ms step:1443/1600 train_time:84389ms step_avg:58.48ms step:1444/1600 train_time:84474ms step_avg:58.50ms step:1445/1600 train_time:84563ms step_avg:58.52ms step:1446/1600 train_time:84650ms step_avg:58.54ms step:1447/1600 train_time:84738ms step_avg:58.56ms step:1448/1600 train_time:84823ms step_avg:58.58ms step:1449/1600 train_time:84910ms step_avg:58.60ms step:1450/1600 train_time:84997ms step_avg:58.62ms step:1451/1600 train_time:85084ms step_avg:58.64ms step:1452/1600 train_time:85169ms step_avg:58.66ms step:1453/1600 train_time:85258ms step_avg:58.68ms step:1454/1600 train_time:85343ms step_avg:58.70ms step:1455/1600 train_time:85431ms step_avg:58.72ms step:1456/1600 train_time:85516ms step_avg:58.73ms step:1457/1600 train_time:85604ms step_avg:58.75ms step:1458/1600 train_time:85689ms step_avg:58.77ms step:1459/1600 train_time:85778ms step_avg:58.79ms step:1460/1600 train_time:85864ms step_avg:58.81ms step:1461/1600 train_time:85952ms step_avg:58.83ms step:1462/1600 train_time:86037ms step_avg:58.85ms step:1463/1600 train_time:86125ms step_avg:58.87ms step:1464/1600 train_time:86211ms step_avg:58.89ms step:1465/1600 train_time:86300ms step_avg:58.91ms step:1466/1600 train_time:86385ms step_avg:58.93ms step:1467/1600 train_time:86473ms step_avg:58.95ms step:1468/1600 train_time:86559ms step_avg:58.96ms step:1469/1600 train_time:86647ms step_avg:58.98ms step:1470/1600 train_time:86733ms step_avg:59.00ms step:1471/1600 train_time:86821ms step_avg:59.02ms step:1472/1600 train_time:86907ms step_avg:59.04ms step:1473/1600 train_time:86993ms step_avg:59.06ms step:1474/1600 train_time:87079ms step_avg:59.08ms step:1475/1600 train_time:87168ms step_avg:59.10ms step:1476/1600 train_time:87253ms step_avg:59.11ms step:1477/1600 train_time:87342ms step_avg:59.13ms step:1478/1600 train_time:87428ms step_avg:59.15ms step:1479/1600 train_time:87515ms step_avg:59.17ms step:1480/1600 train_time:87600ms step_avg:59.19ms step:1481/1600 train_time:87689ms step_avg:59.21ms step:1482/1600 train_time:87773ms step_avg:59.23ms step:1483/1600 train_time:87861ms step_avg:59.25ms step:1484/1600 train_time:87946ms step_avg:59.26ms step:1485/1600 train_time:88035ms step_avg:59.28ms step:1486/1600 train_time:88120ms step_avg:59.30ms step:1487/1600 train_time:88209ms step_avg:59.32ms step:1488/1600 train_time:88294ms step_avg:59.34ms step:1489/1600 train_time:88383ms step_avg:59.36ms step:1490/1600 train_time:88468ms step_avg:59.37ms step:1491/1600 train_time:88555ms step_avg:59.39ms step:1492/1600 train_time:88640ms step_avg:59.41ms step:1493/1600 train_time:88728ms step_avg:59.43ms step:1494/1600 train_time:88813ms step_avg:59.45ms step:1495/1600 train_time:88902ms step_avg:59.47ms step:1496/1600 train_time:88988ms step_avg:59.48ms step:1497/1600 train_time:89076ms step_avg:59.50ms step:1498/1600 train_time:89161ms step_avg:59.52ms step:1499/1600 train_time:89249ms step_avg:59.54ms step:1500/1600 train_time:89334ms step_avg:59.56ms step:1500/1600 val_loss:3.3058 train_time:89408ms step_avg:59.61ms step:1501/1600 train_time:89427ms step_avg:59.58ms step:1502/1600 train_time:89512ms step_avg:59.60ms step:1503/1600 train_time:89607ms step_avg:59.62ms step:1504/1600 train_time:89694ms step_avg:59.64ms step:1505/1600 train_time:89782ms step_avg:59.66ms step:1506/1600 train_time:89866ms step_avg:59.67ms step:1507/1600 train_time:89953ms step_avg:59.69ms step:1508/1600 train_time:90036ms step_avg:59.71ms step:1509/1600 train_time:90123ms step_avg:59.72ms step:1510/1600 train_time:90207ms step_avg:59.74ms step:1511/1600 train_time:90295ms step_avg:59.76ms step:1512/1600 train_time:90381ms step_avg:59.78ms step:1513/1600 train_time:90472ms step_avg:59.80ms step:1514/1600 train_time:90558ms step_avg:59.81ms step:1515/1600 train_time:90648ms step_avg:59.83ms step:1516/1600 train_time:90734ms step_avg:59.85ms step:1517/1600 train_time:90823ms step_avg:59.87ms step:1518/1600 train_time:90907ms step_avg:59.89ms step:1519/1600 train_time:90993ms step_avg:59.90ms step:1520/1600 train_time:91077ms step_avg:59.92ms step:1521/1600 train_time:91165ms step_avg:59.94ms step:1522/1600 train_time:91249ms step_avg:59.95ms step:1523/1600 train_time:91337ms step_avg:59.97ms step:1524/1600 train_time:91423ms step_avg:59.99ms step:1525/1600 train_time:91512ms step_avg:60.01ms step:1526/1600 train_time:91599ms step_avg:60.03ms step:1527/1600 train_time:91688ms step_avg:60.04ms step:1528/1600 train_time:91775ms step_avg:60.06ms step:1529/1600 train_time:91864ms step_avg:60.08ms step:1530/1600 train_time:91948ms step_avg:60.10ms step:1531/1600 train_time:92035ms step_avg:60.11ms step:1532/1600 train_time:92120ms step_avg:60.13ms step:1533/1600 train_time:92208ms step_avg:60.15ms step:1534/1600 train_time:92293ms step_avg:60.16ms step:1535/1600 train_time:92382ms step_avg:60.18ms step:1536/1600 train_time:92466ms step_avg:60.20ms step:1537/1600 train_time:92555ms step_avg:60.22ms step:1538/1600 train_time:92641ms step_avg:60.23ms step:1539/1600 train_time:92730ms step_avg:60.25ms step:1540/1600 train_time:92816ms step_avg:60.27ms step:1541/1600 train_time:92904ms step_avg:60.29ms step:1542/1600 train_time:92989ms step_avg:60.30ms step:1543/1600 train_time:93076ms step_avg:60.32ms step:1544/1600 train_time:93161ms step_avg:60.34ms step:1545/1600 train_time:93248ms step_avg:60.35ms step:1546/1600 train_time:93333ms step_avg:60.37ms step:1547/1600 train_time:93422ms step_avg:60.39ms step:1548/1600 train_time:93508ms step_avg:60.41ms step:1549/1600 train_time:93594ms step_avg:60.42ms step:1550/1600 train_time:93681ms step_avg:60.44ms step:1551/1600 train_time:93769ms step_avg:60.46ms step:1552/1600 train_time:93854ms step_avg:60.47ms step:1553/1600 train_time:93942ms step_avg:60.49ms step:1554/1600 train_time:94028ms step_avg:60.51ms step:1555/1600 train_time:94116ms step_avg:60.52ms step:1556/1600 train_time:94201ms step_avg:60.54ms step:1557/1600 train_time:94289ms step_avg:60.56ms step:1558/1600 train_time:94375ms step_avg:60.57ms step:1559/1600 train_time:94463ms step_avg:60.59ms step:1560/1600 train_time:94548ms step_avg:60.61ms step:1561/1600 train_time:94644ms step_avg:60.63ms step:1562/1600 train_time:94727ms step_avg:60.64ms step:1563/1600 train_time:94814ms step_avg:60.66ms step:1564/1600 train_time:94900ms step_avg:60.68ms step:1565/1600 train_time:94988ms step_avg:60.70ms step:1566/1600 train_time:95075ms step_avg:60.71ms step:1567/1600 train_time:95164ms step_avg:60.73ms step:1568/1600 train_time:95249ms step_avg:60.75ms step:1569/1600 train_time:95337ms step_avg:60.76ms step:1570/1600 train_time:95422ms step_avg:60.78ms step:1571/1600 train_time:95511ms step_avg:60.80ms step:1572/1600 train_time:95598ms step_avg:60.81ms step:1573/1600 train_time:95686ms step_avg:60.83ms step:1574/1600 train_time:95773ms step_avg:60.85ms step:1575/1600 train_time:95862ms step_avg:60.86ms step:1576/1600 train_time:95947ms step_avg:60.88ms step:1577/1600 train_time:96040ms step_avg:60.90ms step:1578/1600 train_time:96123ms step_avg:60.91ms step:1579/1600 train_time:96210ms step_avg:60.93ms step:1580/1600 train_time:96295ms step_avg:60.95ms step:1581/1600 train_time:96383ms step_avg:60.96ms step:1582/1600 train_time:96469ms step_avg:60.98ms step:1583/1600 train_time:96558ms step_avg:61.00ms step:1584/1600 train_time:96644ms step_avg:61.01ms step:1585/1600 train_time:96733ms step_avg:61.03ms step:1586/1600 train_time:96818ms step_avg:61.05ms step:1587/1600 train_time:96907ms step_avg:61.06ms step:1588/1600 train_time:96993ms step_avg:61.08ms step:1589/1600 train_time:97081ms step_avg:61.10ms step:1590/1600 train_time:97166ms step_avg:61.11ms step:1591/1600 train_time:97255ms step_avg:61.13ms step:1592/1600 train_time:97340ms step_avg:61.14ms step:1593/1600 train_time:97429ms step_avg:61.16ms step:1594/1600 train_time:97514ms step_avg:61.18ms step:1595/1600 train_time:97603ms step_avg:61.19ms step:1596/1600 train_time:97688ms step_avg:61.21ms step:1597/1600 train_time:97777ms step_avg:61.23ms step:1598/1600 train_time:97863ms step_avg:61.24ms step:1599/1600 train_time:97951ms step_avg:61.26ms step:1600/1600 train_time:98037ms step_avg:61.27ms step:1600/1600 val_loss:3.2763 train_time:98110ms step_avg:61.32ms peak memory allocated: 30264 MiB reserved: 46338 MiB