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:23:19 2026 +-----------------------------------------------------------------------------------------+ | NVIDIA-SMI 570.148.08 Driver Version: 570.148.08 CUDA Version: 12.8 | |-----------------------------------------+------------------------+----------------------+ | GPU Name Persistence-M | Bus-Id Disp.A | Volatile Uncorr. ECC | | Fan Temp Perf Pwr:Usage/Cap | Memory-Usage | GPU-Util Compute M. | | | | MIG M. | |=========================================+========================+======================| | 0 NVIDIA H100 80GB HBM3 On | 00000000:61:00.0 Off | 0 | | N/A 34C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:62:00.0 Off | 0 | | N/A 38C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 40C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 35C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:6A:00.0 Off | 0 | | N/A 34C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 41C P0 129W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 39C P0 125W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 35C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 303355 C /usr/bin/python3 1510MiB | | 1 N/A N/A 303356 C /usr/bin/python3 1510MiB | | 2 N/A N/A 303357 C /usr/bin/python3 1510MiB | | 3 N/A N/A 303358 C /usr/bin/python3 1510MiB | | 4 N/A N/A 303359 C /usr/bin/python3 1510MiB | | 5 N/A N/A 303360 C /usr/bin/python3 1510MiB | | 6 N/A N/A 303361 C /usr/bin/python3 1510MiB | | 7 N/A N/A 303362 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.8285 train_time:0ms step_avg:0.06ms step:1/1600 train_time:79ms step_avg:78.68ms step:2/1600 train_time:104ms step_avg:52.04ms step:3/1600 train_time:123ms step_avg:40.87ms step:4/1600 train_time:145ms step_avg:36.33ms step:5/1600 train_time:176ms step_avg:35.15ms step:6/1600 train_time:286ms step_avg:47.73ms step:7/1600 train_time:307ms step_avg:43.88ms step:8/1600 train_time:326ms step_avg:40.77ms step:9/1600 train_time:352ms step_avg:39.06ms step:10/1600 train_time:389ms step_avg:38.86ms step:11/1600 train_time:419ms step_avg:38.08ms step:12/1600 train_time:456ms step_avg:38.01ms step:13/1600 train_time:487ms step_avg:37.44ms step:14/1600 train_time:524ms step_avg:37.42ms step:15/1600 train_time:554ms step_avg:36.97ms step:16/1600 train_time:592ms step_avg:37.01ms step:17/1600 train_time:624ms step_avg:36.68ms step:18/1600 train_time:661ms step_avg:36.73ms step:19/1600 train_time:692ms step_avg:36.40ms step:20/1600 train_time:729ms step_avg:36.44ms step:21/1600 train_time:759ms step_avg:36.16ms step:22/1600 train_time:796ms step_avg:36.20ms step:23/1600 train_time:827ms step_avg:35.96ms step:24/1600 train_time:864ms step_avg:36.02ms step:25/1600 train_time:895ms step_avg:35.80ms step:26/1600 train_time:933ms step_avg:35.87ms step:27/1600 train_time:963ms step_avg:35.68ms step:28/1600 train_time:1001ms step_avg:35.74ms step:29/1600 train_time:1031ms step_avg:35.56ms step:30/1600 train_time:1068ms step_avg:35.61ms step:31/1600 train_time:1099ms step_avg:35.46ms step:32/1600 train_time:1136ms step_avg:35.51ms step:33/1600 train_time:1167ms step_avg:35.37ms step:34/1600 train_time:1204ms step_avg:35.42ms step:35/1600 train_time:1235ms step_avg:35.30ms step:36/1600 train_time:1274ms step_avg:35.38ms step:37/1600 train_time:1305ms step_avg:35.26ms step:38/1600 train_time:1342ms step_avg:35.31ms step:39/1600 train_time:1373ms step_avg:35.20ms step:40/1600 train_time:1410ms step_avg:35.26ms step:41/1600 train_time:1441ms step_avg:35.15ms step:42/1600 train_time:1478ms step_avg:35.20ms step:43/1600 train_time:1509ms step_avg:35.09ms step:44/1600 train_time:1546ms step_avg:35.14ms step:45/1600 train_time:1577ms step_avg:35.05ms step:46/1600 train_time:1614ms step_avg:35.09ms step:47/1600 train_time:1645ms step_avg:35.01ms step:48/1600 train_time:1682ms step_avg:35.05ms step:49/1600 train_time:1713ms step_avg:34.96ms step:50/1600 train_time:1750ms step_avg:35.01ms step:51/1600 train_time:1781ms step_avg:34.93ms step:52/1600 train_time:1818ms step_avg:34.97ms step:53/1600 train_time:1849ms step_avg:34.89ms step:54/1600 train_time:1887ms step_avg:34.94ms step:55/1600 train_time:1917ms step_avg:34.86ms step:56/1600 train_time:1954ms step_avg:34.90ms step:57/1600 train_time:1986ms step_avg:34.84ms step:58/1600 train_time:2022ms step_avg:34.87ms step:59/1600 train_time:2053ms step_avg:34.80ms step:60/1600 train_time:2091ms step_avg:34.85ms step:61/1600 train_time:2122ms step_avg:34.78ms step:62/1600 train_time:2159ms step_avg:34.82ms step:63/1600 train_time:2190ms step_avg:34.76ms step:64/1600 train_time:2227ms step_avg:34.79ms step:65/1600 train_time:2258ms step_avg:34.73ms step:66/1600 train_time:2295ms step_avg:34.77ms step:67/1600 train_time:2326ms step_avg:34.71ms step:68/1600 train_time:2363ms step_avg:34.75ms step:69/1600 train_time:2394ms step_avg:34.69ms step:70/1600 train_time:2431ms step_avg:34.73ms step:71/1600 train_time:2462ms step_avg:34.68ms step:72/1600 train_time:2500ms step_avg:34.72ms step:73/1600 train_time:2531ms step_avg:34.67ms step:74/1600 train_time:2568ms step_avg:34.70ms step:75/1600 train_time:2599ms step_avg:34.65ms step:76/1600 train_time:2636ms step_avg:34.69ms step:77/1600 train_time:2667ms step_avg:34.64ms step:78/1600 train_time:2704ms step_avg:34.66ms step:79/1600 train_time:2734ms step_avg:34.61ms step:80/1600 train_time:2772ms step_avg:34.65ms step:81/1600 train_time:2802ms step_avg:34.60ms step:82/1600 train_time:2839ms step_avg:34.62ms step:83/1600 train_time:2870ms step_avg:34.58ms step:84/1600 train_time:2907ms step_avg:34.61ms step:85/1600 train_time:2938ms step_avg:34.56ms step:86/1600 train_time:2975ms step_avg:34.59ms step:87/1600 train_time:3005ms step_avg:34.55ms step:88/1600 train_time:3043ms step_avg:34.58ms step:89/1600 train_time:3074ms step_avg:34.54ms step:90/1600 train_time:3111ms step_avg:34.57ms step:91/1600 train_time:3142ms step_avg:34.53ms step:92/1600 train_time:3180ms step_avg:34.56ms step:93/1600 train_time:3210ms step_avg:34.52ms step:94/1600 train_time:3248ms step_avg:34.55ms step:95/1600 train_time:3279ms step_avg:34.51ms step:96/1600 train_time:3316ms step_avg:34.54ms step:97/1600 train_time:3346ms step_avg:34.50ms step:98/1600 train_time:3384ms step_avg:34.53ms step:99/1600 train_time:3414ms step_avg:34.49ms step:100/1600 train_time:3452ms step_avg:34.52ms step:101/1600 train_time:3483ms step_avg:34.48ms step:102/1600 train_time:3520ms step_avg:34.51ms step:103/1600 train_time:3551ms step_avg:34.47ms step:104/1600 train_time:3588ms step_avg:34.50ms step:105/1600 train_time:3618ms step_avg:34.46ms step:106/1600 train_time:3655ms step_avg:34.48ms step:107/1600 train_time:3686ms step_avg:34.45ms step:108/1600 train_time:3723ms step_avg:34.47ms step:109/1600 train_time:3753ms step_avg:34.44ms step:110/1600 train_time:3791ms step_avg:34.47ms step:111/1600 train_time:3822ms step_avg:34.43ms step:112/1600 train_time:3859ms step_avg:34.46ms step:113/1600 train_time:3890ms step_avg:34.42ms step:114/1600 train_time:3927ms step_avg:34.45ms step:115/1600 train_time:3958ms step_avg:34.41ms step:116/1600 train_time:3995ms step_avg:34.44ms step:117/1600 train_time:4026ms step_avg:34.41ms step:118/1600 train_time:4063ms step_avg:34.43ms step:119/1600 train_time:4094ms step_avg:34.40ms step:120/1600 train_time:4131ms step_avg:34.43ms step:121/1600 train_time:4162ms step_avg:34.40ms step:122/1600 train_time:4199ms step_avg:34.42ms step:123/1600 train_time:4230ms step_avg:34.39ms step:124/1600 train_time:4267ms step_avg:34.41ms step:125/1600 train_time:4298ms step_avg:34.38ms step:126/1600 train_time:4335ms step_avg:34.41ms step:127/1600 train_time:4366ms step_avg:34.37ms step:128/1600 train_time:4403ms step_avg:34.40ms step:129/1600 train_time:4434ms step_avg:34.37ms step:130/1600 train_time:4471ms step_avg:34.39ms step:131/1600 train_time:4502ms step_avg:34.37ms step:132/1600 train_time:4540ms step_avg:34.39ms step:133/1600 train_time:4570ms step_avg:34.36ms step:134/1600 train_time:4607ms step_avg:34.38ms step:135/1600 train_time:4638ms step_avg:34.35ms step:136/1600 train_time:4676ms step_avg:34.38ms step:137/1600 train_time:4706ms step_avg:34.35ms step:138/1600 train_time:4743ms step_avg:34.37ms step:139/1600 train_time:4774ms step_avg:34.34ms step:140/1600 train_time:4811ms step_avg:34.36ms step:141/1600 train_time:4842ms step_avg:34.34ms step:142/1600 train_time:4880ms step_avg:34.36ms step:143/1600 train_time:4910ms step_avg:34.34ms step:144/1600 train_time:4947ms step_avg:34.35ms step:145/1600 train_time:4978ms step_avg:34.33ms step:146/1600 train_time:5015ms step_avg:34.35ms step:147/1600 train_time:5046ms step_avg:34.32ms step:148/1600 train_time:5083ms step_avg:34.35ms step:149/1600 train_time:5114ms step_avg:34.32ms step:150/1600 train_time:5151ms step_avg:34.34ms step:151/1600 train_time:5183ms step_avg:34.32ms step:152/1600 train_time:5220ms step_avg:34.34ms step:153/1600 train_time:5251ms step_avg:34.32ms step:154/1600 train_time:5289ms step_avg:34.34ms step:155/1600 train_time:5319ms step_avg:34.32ms step:156/1600 train_time:5356ms step_avg:34.33ms step:157/1600 train_time:5387ms step_avg:34.31ms step:158/1600 train_time:5423ms step_avg:34.32ms step:159/1600 train_time:5455ms step_avg:34.31ms step:160/1600 train_time:5492ms step_avg:34.33ms step:161/1600 train_time:5523ms step_avg:34.31ms step:162/1600 train_time:5560ms step_avg:34.32ms step:163/1600 train_time:5591ms step_avg:34.30ms step:164/1600 train_time:5628ms step_avg:34.32ms step:165/1600 train_time:5659ms step_avg:34.30ms step:166/1600 train_time:5697ms step_avg:34.32ms step:167/1600 train_time:5727ms step_avg:34.29ms step:168/1600 train_time:5764ms step_avg:34.31ms step:169/1600 train_time:5795ms step_avg:34.29ms step:170/1600 train_time:5832ms step_avg:34.31ms step:171/1600 train_time:5863ms step_avg:34.29ms step:172/1600 train_time:5900ms step_avg:34.30ms step:173/1600 train_time:5931ms step_avg:34.28ms step:174/1600 train_time:5968ms step_avg:34.30ms step:175/1600 train_time:5999ms step_avg:34.28ms step:176/1600 train_time:6036ms step_avg:34.29ms step:177/1600 train_time:6066ms step_avg:34.27ms step:178/1600 train_time:6103ms step_avg:34.29ms step:179/1600 train_time:6134ms step_avg:34.27ms step:180/1600 train_time:6171ms step_avg:34.29ms step:181/1600 train_time:6202ms step_avg:34.27ms step:182/1600 train_time:6239ms step_avg:34.28ms step:183/1600 train_time:6270ms step_avg:34.26ms step:184/1600 train_time:6307ms step_avg:34.28ms step:185/1600 train_time:6337ms step_avg:34.26ms step:186/1600 train_time:6375ms step_avg:34.27ms step:187/1600 train_time:6405ms step_avg:34.25ms step:188/1600 train_time:6442ms step_avg:34.27ms step:189/1600 train_time:6473ms step_avg:34.25ms step:190/1600 train_time:6510ms step_avg:34.26ms step:191/1600 train_time:6541ms step_avg:34.25ms step:192/1600 train_time:6579ms step_avg:34.26ms step:193/1600 train_time:6609ms step_avg:34.24ms step:194/1600 train_time:6646ms step_avg:34.26ms step:195/1600 train_time:6677ms step_avg:34.24ms step:196/1600 train_time:6714ms step_avg:34.26ms step:197/1600 train_time:6745ms step_avg:34.24ms step:198/1600 train_time:6782ms step_avg:34.25ms step:199/1600 train_time:6813ms step_avg:34.24ms step:200/1600 train_time:6850ms step_avg:34.25ms step:201/1600 train_time:6882ms step_avg:34.24ms step:202/1600 train_time:6919ms step_avg:34.25ms step:203/1600 train_time:6950ms step_avg:34.24ms step:204/1600 train_time:6987ms step_avg:34.25ms step:205/1600 train_time:7017ms step_avg:34.23ms step:206/1600 train_time:7054ms step_avg:34.24ms step:207/1600 train_time:7085ms step_avg:34.23ms step:208/1600 train_time:7122ms step_avg:34.24ms step:209/1600 train_time:7153ms step_avg:34.23ms step:210/1600 train_time:7190ms step_avg:34.24ms step:211/1600 train_time:7222ms step_avg:34.23ms step:212/1600 train_time:7258ms step_avg:34.24ms step:213/1600 train_time:7289ms step_avg:34.22ms step:214/1600 train_time:7327ms step_avg:34.24ms step:215/1600 train_time:7357ms step_avg:34.22ms step:216/1600 train_time:7394ms step_avg:34.23ms step:217/1600 train_time:7424ms step_avg:34.21ms step:218/1600 train_time:7462ms step_avg:34.23ms step:219/1600 train_time:7492ms step_avg:34.21ms step:220/1600 train_time:7529ms step_avg:34.22ms step:221/1600 train_time:7560ms step_avg:34.21ms step:222/1600 train_time:7598ms step_avg:34.22ms step:223/1600 train_time:7629ms step_avg:34.21ms step:224/1600 train_time:7665ms step_avg:34.22ms step:225/1600 train_time:7696ms step_avg:34.21ms step:226/1600 train_time:7734ms step_avg:34.22ms step:227/1600 train_time:7764ms step_avg:34.20ms step:228/1600 train_time:7801ms step_avg:34.22ms step:229/1600 train_time:7832ms step_avg:34.20ms step:230/1600 train_time:7870ms step_avg:34.22ms step:231/1600 train_time:7901ms step_avg:34.20ms step:232/1600 train_time:7938ms step_avg:34.21ms step:233/1600 train_time:7968ms step_avg:34.20ms step:234/1600 train_time:8006ms step_avg:34.21ms step:235/1600 train_time:8036ms step_avg:34.20ms step:236/1600 train_time:8073ms step_avg:34.21ms step:237/1600 train_time:8104ms step_avg:34.19ms step:238/1600 train_time:8141ms step_avg:34.21ms step:239/1600 train_time:8172ms step_avg:34.19ms step:240/1600 train_time:8209ms step_avg:34.20ms step:241/1600 train_time:8240ms step_avg:34.19ms step:242/1600 train_time:8277ms step_avg:34.20ms step:243/1600 train_time:8308ms step_avg:34.19ms step:244/1600 train_time:8345ms step_avg:34.20ms step:245/1600 train_time:8375ms step_avg:34.19ms step:246/1600 train_time:8412ms step_avg:34.20ms step:247/1600 train_time:8444ms step_avg:34.18ms step:248/1600 train_time:8480ms step_avg:34.20ms step:249/1600 train_time:8511ms step_avg:34.18ms step:250/1600 train_time:8548ms step_avg:34.19ms step:250/1600 val_loss:4.5834 train_time:8596ms step_avg:34.38ms step:251/1600 train_time:8615ms step_avg:34.32ms step:252/1600 train_time:8633ms step_avg:34.26ms step:253/1600 train_time:8650ms step_avg:34.19ms step:254/1600 train_time:8688ms step_avg:34.20ms step:255/1600 train_time:8721ms step_avg:34.20ms step:256/1600 train_time:8760ms step_avg:34.22ms step:257/1600 train_time:8791ms step_avg:34.21ms step:258/1600 train_time:8828ms step_avg:34.22ms step:259/1600 train_time:8859ms step_avg:34.20ms step:260/1600 train_time:8897ms step_avg:34.22ms step:261/1600 train_time:8928ms step_avg:34.21ms step:262/1600 train_time:8966ms step_avg:34.22ms step:263/1600 train_time:8996ms step_avg:34.21ms step:264/1600 train_time:9034ms step_avg:34.22ms step:265/1600 train_time:9064ms step_avg:34.20ms step:266/1600 train_time:9101ms step_avg:34.21ms step:267/1600 train_time:9132ms step_avg:34.20ms step:268/1600 train_time:9168ms step_avg:34.21ms step:269/1600 train_time:9199ms step_avg:34.20ms step:270/1600 train_time:9237ms step_avg:34.21ms step:271/1600 train_time:9267ms step_avg:34.19ms step:272/1600 train_time:9304ms step_avg:34.20ms step:273/1600 train_time:9334ms step_avg:34.19ms step:274/1600 train_time:9371ms step_avg:34.20ms step:275/1600 train_time:9401ms step_avg:34.19ms step:276/1600 train_time:9438ms step_avg:34.20ms step:277/1600 train_time:9469ms step_avg:34.18ms step:278/1600 train_time:9506ms step_avg:34.19ms step:279/1600 train_time:9536ms step_avg:34.18ms step:280/1600 train_time:9573ms step_avg:34.19ms step:281/1600 train_time:9604ms step_avg:34.18ms step:282/1600 train_time:9641ms step_avg:34.19ms step:283/1600 train_time:9672ms step_avg:34.18ms step:284/1600 train_time:9709ms step_avg:34.19ms step:285/1600 train_time:9740ms step_avg:34.18ms step:286/1600 train_time:9777ms step_avg:34.19ms step:287/1600 train_time:9808ms step_avg:34.17ms step:288/1600 train_time:9845ms step_avg:34.18ms step:289/1600 train_time:9876ms step_avg:34.17ms step:290/1600 train_time:9913ms step_avg:34.18ms step:291/1600 train_time:9944ms step_avg:34.17ms step:292/1600 train_time:9981ms step_avg:34.18ms step:293/1600 train_time:10012ms step_avg:34.17ms step:294/1600 train_time:10049ms step_avg:34.18ms step:295/1600 train_time:10080ms step_avg:34.17ms step:296/1600 train_time:10117ms step_avg:34.18ms step:297/1600 train_time:10148ms step_avg:34.17ms step:298/1600 train_time:10185ms step_avg:34.18ms step:299/1600 train_time:10215ms step_avg:34.16ms step:300/1600 train_time:10252ms step_avg:34.17ms step:301/1600 train_time:10282ms step_avg:34.16ms step:302/1600 train_time:10320ms step_avg:34.17ms step:303/1600 train_time:10350ms step_avg:34.16ms step:304/1600 train_time:10387ms step_avg:34.17ms step:305/1600 train_time:10418ms step_avg:34.16ms step:306/1600 train_time:10455ms step_avg:34.17ms step:307/1600 train_time:10486ms step_avg:34.16ms step:308/1600 train_time:10523ms step_avg:34.16ms step:309/1600 train_time:10553ms step_avg:34.15ms step:310/1600 train_time:10590ms step_avg:34.16ms step:311/1600 train_time:10621ms step_avg:34.15ms step:312/1600 train_time:10658ms step_avg:34.16ms step:313/1600 train_time:10689ms step_avg:34.15ms step:314/1600 train_time:10726ms step_avg:34.16ms step:315/1600 train_time:10756ms step_avg:34.15ms step:316/1600 train_time:10793ms step_avg:34.15ms step:317/1600 train_time:10824ms step_avg:34.14ms step:318/1600 train_time:10861ms step_avg:34.15ms step:319/1600 train_time:10892ms step_avg:34.14ms step:320/1600 train_time:10929ms step_avg:34.15ms step:321/1600 train_time:10960ms step_avg:34.14ms step:322/1600 train_time:10997ms step_avg:34.15ms step:323/1600 train_time:11028ms step_avg:34.14ms step:324/1600 train_time:11065ms step_avg:34.15ms step:325/1600 train_time:11096ms step_avg:34.14ms step:326/1600 train_time:11133ms step_avg:34.15ms step:327/1600 train_time:11164ms step_avg:34.14ms step:328/1600 train_time:11201ms step_avg:34.15ms step:329/1600 train_time:11232ms step_avg:34.14ms step:330/1600 train_time:11269ms step_avg:34.15ms step:331/1600 train_time:11300ms step_avg:34.14ms step:332/1600 train_time:11337ms step_avg:34.15ms step:333/1600 train_time:11368ms step_avg:34.14ms step:334/1600 train_time:11405ms step_avg:34.15ms step:335/1600 train_time:11436ms step_avg:34.14ms step:336/1600 train_time:11473ms step_avg:34.15ms step:337/1600 train_time:11503ms step_avg:34.13ms step:338/1600 train_time:11540ms step_avg:34.14ms step:339/1600 train_time:11571ms step_avg:34.13ms step:340/1600 train_time:11608ms step_avg:34.14ms step:341/1600 train_time:11639ms step_avg:34.13ms step:342/1600 train_time:11676ms step_avg:34.14ms step:343/1600 train_time:11707ms step_avg:34.13ms step:344/1600 train_time:11744ms step_avg:34.14ms step:345/1600 train_time:11775ms step_avg:34.13ms step:346/1600 train_time:11812ms step_avg:34.14ms step:347/1600 train_time:11842ms step_avg:34.13ms step:348/1600 train_time:11879ms step_avg:34.14ms step:349/1600 train_time:11910ms step_avg:34.13ms step:350/1600 train_time:11947ms step_avg:34.13ms step:351/1600 train_time:11978ms step_avg:34.12ms step:352/1600 train_time:12015ms step_avg:34.13ms step:353/1600 train_time:12046ms step_avg:34.12ms step:354/1600 train_time:12083ms step_avg:34.13ms step:355/1600 train_time:12114ms step_avg:34.12ms step:356/1600 train_time:12151ms step_avg:34.13ms step:357/1600 train_time:12181ms step_avg:34.12ms step:358/1600 train_time:12219ms step_avg:34.13ms step:359/1600 train_time:12249ms step_avg:34.12ms step:360/1600 train_time:12286ms step_avg:34.13ms step:361/1600 train_time:12317ms step_avg:34.12ms step:362/1600 train_time:12355ms step_avg:34.13ms step:363/1600 train_time:12385ms step_avg:34.12ms step:364/1600 train_time:12423ms step_avg:34.13ms step:365/1600 train_time:12453ms step_avg:34.12ms step:366/1600 train_time:12490ms step_avg:34.13ms step:367/1600 train_time:12521ms step_avg:34.12ms step:368/1600 train_time:12558ms step_avg:34.12ms step:369/1600 train_time:12588ms step_avg:34.11ms step:370/1600 train_time:12625ms step_avg:34.12ms step:371/1600 train_time:12656ms step_avg:34.11ms step:372/1600 train_time:12693ms step_avg:34.12ms step:373/1600 train_time:12724ms step_avg:34.11ms step:374/1600 train_time:12761ms step_avg:34.12ms step:375/1600 train_time:12791ms step_avg:34.11ms step:376/1600 train_time:12828ms step_avg:34.12ms step:377/1600 train_time:12859ms step_avg:34.11ms step:378/1600 train_time:12896ms step_avg:34.12ms step:379/1600 train_time:12927ms step_avg:34.11ms step:380/1600 train_time:12964ms step_avg:34.11ms step:381/1600 train_time:12994ms step_avg:34.10ms step:382/1600 train_time:13031ms step_avg:34.11ms step:383/1600 train_time:13061ms step_avg:34.10ms step:384/1600 train_time:13098ms step_avg:34.11ms step:385/1600 train_time:13129ms step_avg:34.10ms step:386/1600 train_time:13166ms step_avg:34.11ms step:387/1600 train_time:13197ms step_avg:34.10ms step:388/1600 train_time:13233ms step_avg:34.11ms step:389/1600 train_time:13264ms step_avg:34.10ms step:390/1600 train_time:13301ms step_avg:34.11ms step:391/1600 train_time:13332ms step_avg:34.10ms step:392/1600 train_time:13368ms step_avg:34.10ms step:393/1600 train_time:13399ms step_avg:34.09ms step:394/1600 train_time:13437ms step_avg:34.10ms step:395/1600 train_time:13468ms step_avg:34.10ms step:396/1600 train_time:13505ms step_avg:34.10ms step:397/1600 train_time:13535ms step_avg:34.09ms step:398/1600 train_time:13572ms step_avg:34.10ms step:399/1600 train_time:13602ms step_avg:34.09ms step:400/1600 train_time:13639ms step_avg:34.10ms step:401/1600 train_time:13670ms step_avg:34.09ms step:402/1600 train_time:13707ms step_avg:34.10ms step:403/1600 train_time:13738ms step_avg:34.09ms step:404/1600 train_time:13776ms step_avg:34.10ms step:405/1600 train_time:13806ms step_avg:34.09ms step:406/1600 train_time:13844ms step_avg:34.10ms step:407/1600 train_time:13875ms step_avg:34.09ms step:408/1600 train_time:13912ms step_avg:34.10ms step:409/1600 train_time:13942ms step_avg:34.09ms step:410/1600 train_time:13979ms step_avg:34.10ms step:411/1600 train_time:14010ms step_avg:34.09ms step:412/1600 train_time:14047ms step_avg:34.09ms step:413/1600 train_time:14078ms step_avg:34.09ms step:414/1600 train_time:14115ms step_avg:34.09ms step:415/1600 train_time:14146ms step_avg:34.09ms step:416/1600 train_time:14183ms step_avg:34.09ms step:417/1600 train_time:14214ms step_avg:34.09ms step:418/1600 train_time:14250ms step_avg:34.09ms step:419/1600 train_time:14281ms step_avg:34.08ms step:420/1600 train_time:14318ms step_avg:34.09ms step:421/1600 train_time:14349ms step_avg:34.08ms step:422/1600 train_time:14386ms step_avg:34.09ms step:423/1600 train_time:14416ms step_avg:34.08ms step:424/1600 train_time:14455ms step_avg:34.09ms step:425/1600 train_time:14485ms step_avg:34.08ms step:426/1600 train_time:14522ms step_avg:34.09ms step:427/1600 train_time:14552ms step_avg:34.08ms step:428/1600 train_time:14590ms step_avg:34.09ms step:429/1600 train_time:14620ms step_avg:34.08ms step:430/1600 train_time:14658ms step_avg:34.09ms step:431/1600 train_time:14688ms step_avg:34.08ms step:432/1600 train_time:14726ms step_avg:34.09ms step:433/1600 train_time:14756ms step_avg:34.08ms step:434/1600 train_time:14794ms step_avg:34.09ms step:435/1600 train_time:14825ms step_avg:34.08ms step:436/1600 train_time:14862ms step_avg:34.09ms step:437/1600 train_time:14893ms step_avg:34.08ms step:438/1600 train_time:14930ms step_avg:34.09ms step:439/1600 train_time:14960ms step_avg:34.08ms step:440/1600 train_time:14997ms step_avg:34.08ms step:441/1600 train_time:15029ms step_avg:34.08ms step:442/1600 train_time:15066ms step_avg:34.09ms step:443/1600 train_time:15096ms step_avg:34.08ms step:444/1600 train_time:15133ms step_avg:34.08ms step:445/1600 train_time:15164ms step_avg:34.08ms step:446/1600 train_time:15201ms step_avg:34.08ms step:447/1600 train_time:15232ms step_avg:34.08ms step:448/1600 train_time:15269ms step_avg:34.08ms step:449/1600 train_time:15300ms step_avg:34.07ms step:450/1600 train_time:15337ms step_avg:34.08ms step:451/1600 train_time:15367ms step_avg:34.07ms step:452/1600 train_time:15404ms step_avg:34.08ms step:453/1600 train_time:15435ms step_avg:34.07ms step:454/1600 train_time:15472ms step_avg:34.08ms step:455/1600 train_time:15503ms step_avg:34.07ms step:456/1600 train_time:15540ms step_avg:34.08ms step:457/1600 train_time:15571ms step_avg:34.07ms step:458/1600 train_time:15608ms step_avg:34.08ms step:459/1600 train_time:15638ms step_avg:34.07ms step:460/1600 train_time:15676ms step_avg:34.08ms step:461/1600 train_time:15706ms step_avg:34.07ms step:462/1600 train_time:15743ms step_avg:34.08ms step:463/1600 train_time:15774ms step_avg:34.07ms step:464/1600 train_time:15811ms step_avg:34.07ms step:465/1600 train_time:15841ms step_avg:34.07ms step:466/1600 train_time:15878ms step_avg:34.07ms step:467/1600 train_time:15909ms step_avg:34.07ms step:468/1600 train_time:15946ms step_avg:34.07ms step:469/1600 train_time:15977ms step_avg:34.07ms step:470/1600 train_time:16014ms step_avg:34.07ms step:471/1600 train_time:16044ms step_avg:34.06ms step:472/1600 train_time:16082ms step_avg:34.07ms step:473/1600 train_time:16112ms step_avg:34.06ms step:474/1600 train_time:16149ms step_avg:34.07ms step:475/1600 train_time:16180ms step_avg:34.06ms step:476/1600 train_time:16217ms step_avg:34.07ms step:477/1600 train_time:16248ms step_avg:34.06ms step:478/1600 train_time:16285ms step_avg:34.07ms step:479/1600 train_time:16316ms step_avg:34.06ms step:480/1600 train_time:16353ms step_avg:34.07ms step:481/1600 train_time:16384ms step_avg:34.06ms step:482/1600 train_time:16421ms step_avg:34.07ms step:483/1600 train_time:16452ms step_avg:34.06ms step:484/1600 train_time:16488ms step_avg:34.07ms step:485/1600 train_time:16519ms step_avg:34.06ms step:486/1600 train_time:16556ms step_avg:34.07ms step:487/1600 train_time:16587ms step_avg:34.06ms step:488/1600 train_time:16623ms step_avg:34.06ms step:489/1600 train_time:16654ms step_avg:34.06ms step:490/1600 train_time:16691ms step_avg:34.06ms step:491/1600 train_time:16721ms step_avg:34.06ms step:492/1600 train_time:16759ms step_avg:34.06ms step:493/1600 train_time:16790ms step_avg:34.06ms step:494/1600 train_time:16827ms step_avg:34.06ms step:495/1600 train_time:16857ms step_avg:34.06ms step:496/1600 train_time:16895ms step_avg:34.06ms step:497/1600 train_time:16926ms step_avg:34.06ms step:498/1600 train_time:16964ms step_avg:34.06ms step:499/1600 train_time:16994ms step_avg:34.06ms step:500/1600 train_time:17031ms step_avg:34.06ms step:500/1600 val_loss:4.2447 train_time:17078ms step_avg:34.16ms step:501/1600 train_time:17097ms step_avg:34.13ms step:502/1600 train_time:17115ms step_avg:34.09ms step:503/1600 train_time:17132ms step_avg:34.06ms step:504/1600 train_time:17169ms step_avg:34.07ms step:505/1600 train_time:17202ms step_avg:34.06ms step:506/1600 train_time:17240ms step_avg:34.07ms step:507/1600 train_time:17272ms step_avg:34.07ms step:508/1600 train_time:17310ms step_avg:34.07ms step:509/1600 train_time:17341ms step_avg:34.07ms step:510/1600 train_time:17378ms step_avg:34.07ms step:511/1600 train_time:17409ms step_avg:34.07ms step:512/1600 train_time:17446ms step_avg:34.07ms step:513/1600 train_time:17476ms step_avg:34.07ms step:514/1600 train_time:17513ms step_avg:34.07ms step:515/1600 train_time:17544ms step_avg:34.07ms step:516/1600 train_time:17581ms step_avg:34.07ms step:517/1600 train_time:17611ms step_avg:34.06ms step:518/1600 train_time:17648ms step_avg:34.07ms step:519/1600 train_time:17678ms step_avg:34.06ms step:520/1600 train_time:17715ms step_avg:34.07ms step:521/1600 train_time:17787ms step_avg:34.14ms step:522/1600 train_time:17842ms step_avg:34.18ms step:523/1600 train_time:17903ms step_avg:34.23ms step:524/1600 train_time:17961ms step_avg:34.28ms step:525/1600 train_time:18022ms step_avg:34.33ms step:526/1600 train_time:18081ms step_avg:34.37ms step:527/1600 train_time:18145ms step_avg:34.43ms step:528/1600 train_time:18206ms step_avg:34.48ms step:529/1600 train_time:18270ms step_avg:34.54ms step:530/1600 train_time:18329ms step_avg:34.58ms step:531/1600 train_time:18391ms step_avg:34.64ms step:532/1600 train_time:18450ms step_avg:34.68ms step:533/1600 train_time:18512ms step_avg:34.73ms step:534/1600 train_time:18572ms step_avg:34.78ms step:535/1600 train_time:18635ms step_avg:34.83ms step:536/1600 train_time:18694ms step_avg:34.88ms step:537/1600 train_time:18757ms step_avg:34.93ms step:538/1600 train_time:18816ms step_avg:34.97ms step:539/1600 train_time:18878ms step_avg:35.02ms step:540/1600 train_time:18937ms step_avg:35.07ms step:541/1600 train_time:18999ms step_avg:35.12ms step:542/1600 train_time:19057ms step_avg:35.16ms step:543/1600 train_time:19120ms step_avg:35.21ms step:544/1600 train_time:19178ms step_avg:35.25ms step:545/1600 train_time:19241ms step_avg:35.30ms step:546/1600 train_time:19300ms step_avg:35.35ms step:547/1600 train_time:19362ms step_avg:35.40ms step:548/1600 train_time:19422ms step_avg:35.44ms step:549/1600 train_time:19485ms step_avg:35.49ms step:550/1600 train_time:19545ms step_avg:35.54ms step:551/1600 train_time:19607ms step_avg:35.58ms step:552/1600 train_time:19667ms step_avg:35.63ms step:553/1600 train_time:19730ms step_avg:35.68ms step:554/1600 train_time:19788ms step_avg:35.72ms step:555/1600 train_time:19850ms step_avg:35.77ms step:556/1600 train_time:19909ms step_avg:35.81ms step:557/1600 train_time:19971ms step_avg:35.85ms step:558/1600 train_time:20031ms step_avg:35.90ms step:559/1600 train_time:20093ms step_avg:35.94ms step:560/1600 train_time:20153ms step_avg:35.99ms step:561/1600 train_time:20216ms step_avg:36.04ms step:562/1600 train_time:20275ms step_avg:36.08ms step:563/1600 train_time:20338ms step_avg:36.12ms step:564/1600 train_time:20397ms step_avg:36.16ms step:565/1600 train_time:20460ms step_avg:36.21ms step:566/1600 train_time:20519ms step_avg:36.25ms step:567/1600 train_time:20581ms step_avg:36.30ms step:568/1600 train_time:20640ms step_avg:36.34ms step:569/1600 train_time:20702ms step_avg:36.38ms step:570/1600 train_time:20762ms step_avg:36.42ms step:571/1600 train_time:20823ms step_avg:36.47ms step:572/1600 train_time:20882ms step_avg:36.51ms step:573/1600 train_time:20949ms step_avg:36.56ms step:574/1600 train_time:21006ms step_avg:36.60ms step:575/1600 train_time:21072ms step_avg:36.65ms step:576/1600 train_time:21127ms step_avg:36.68ms step:577/1600 train_time:21189ms step_avg:36.72ms step:578/1600 train_time:21248ms step_avg:36.76ms step:579/1600 train_time:21310ms step_avg:36.81ms step:580/1600 train_time:21369ms step_avg:36.84ms step:581/1600 train_time:21432ms step_avg:36.89ms step:582/1600 train_time:21492ms step_avg:36.93ms step:583/1600 train_time:21554ms step_avg:36.97ms step:584/1600 train_time:21613ms step_avg:37.01ms step:585/1600 train_time:21676ms step_avg:37.05ms step:586/1600 train_time:21736ms step_avg:37.09ms step:587/1600 train_time:21798ms step_avg:37.14ms step:588/1600 train_time:21857ms step_avg:37.17ms step:589/1600 train_time:21920ms step_avg:37.22ms step:590/1600 train_time:21978ms step_avg:37.25ms step:591/1600 train_time:22041ms step_avg:37.29ms step:592/1600 train_time:22099ms step_avg:37.33ms step:593/1600 train_time:22161ms step_avg:37.37ms step:594/1600 train_time:22220ms step_avg:37.41ms step:595/1600 train_time:22283ms step_avg:37.45ms step:596/1600 train_time:22342ms step_avg:37.49ms step:597/1600 train_time:22404ms step_avg:37.53ms step:598/1600 train_time:22465ms step_avg:37.57ms step:599/1600 train_time:22529ms step_avg:37.61ms step:600/1600 train_time:22586ms step_avg:37.64ms step:601/1600 train_time:22649ms step_avg:37.69ms step:602/1600 train_time:22707ms step_avg:37.72ms step:603/1600 train_time:22769ms step_avg:37.76ms step:604/1600 train_time:22828ms step_avg:37.79ms step:605/1600 train_time:22890ms step_avg:37.84ms step:606/1600 train_time:22949ms step_avg:37.87ms step:607/1600 train_time:23012ms step_avg:37.91ms step:608/1600 train_time:23071ms step_avg:37.95ms step:609/1600 train_time:23134ms step_avg:37.99ms step:610/1600 train_time:23194ms step_avg:38.02ms step:611/1600 train_time:23257ms step_avg:38.06ms step:612/1600 train_time:23316ms step_avg:38.10ms step:613/1600 train_time:23378ms step_avg:38.14ms step:614/1600 train_time:23437ms step_avg:38.17ms step:615/1600 train_time:23500ms step_avg:38.21ms step:616/1600 train_time:23558ms step_avg:38.24ms step:617/1600 train_time:23620ms step_avg:38.28ms step:618/1600 train_time:23679ms step_avg:38.32ms step:619/1600 train_time:23742ms step_avg:38.35ms step:620/1600 train_time:23800ms step_avg:38.39ms step:621/1600 train_time:23862ms step_avg:38.43ms step:622/1600 train_time:23921ms step_avg:38.46ms step:623/1600 train_time:23984ms step_avg:38.50ms step:624/1600 train_time:24043ms step_avg:38.53ms step:625/1600 train_time:24105ms step_avg:38.57ms step:626/1600 train_time:24166ms step_avg:38.60ms step:627/1600 train_time:24229ms step_avg:38.64ms step:628/1600 train_time:24288ms step_avg:38.67ms step:629/1600 train_time:24350ms step_avg:38.71ms step:630/1600 train_time:24409ms step_avg:38.74ms step:631/1600 train_time:24471ms step_avg:38.78ms step:632/1600 train_time:24530ms step_avg:38.81ms step:633/1600 train_time:24592ms step_avg:38.85ms step:634/1600 train_time:24652ms step_avg:38.88ms step:635/1600 train_time:24715ms step_avg:38.92ms step:636/1600 train_time:24774ms step_avg:38.95ms step:637/1600 train_time:24836ms step_avg:38.99ms step:638/1600 train_time:24895ms step_avg:39.02ms step:639/1600 train_time:24958ms step_avg:39.06ms step:640/1600 train_time:25018ms step_avg:39.09ms step:641/1600 train_time:25080ms step_avg:39.13ms step:642/1600 train_time:25138ms step_avg:39.16ms step:643/1600 train_time:25201ms step_avg:39.19ms step:644/1600 train_time:25260ms step_avg:39.22ms step:645/1600 train_time:25322ms step_avg:39.26ms step:646/1600 train_time:25381ms step_avg:39.29ms step:647/1600 train_time:25443ms step_avg:39.32ms step:648/1600 train_time:25502ms step_avg:39.35ms step:649/1600 train_time:25564ms step_avg:39.39ms step:650/1600 train_time:25623ms step_avg:39.42ms step:651/1600 train_time:25685ms step_avg:39.46ms step:652/1600 train_time:25744ms step_avg:39.48ms step:653/1600 train_time:25809ms step_avg:39.52ms step:654/1600 train_time:25866ms step_avg:39.55ms step:655/1600 train_time:25928ms step_avg:39.58ms step:656/1600 train_time:25987ms step_avg:39.61ms step:657/1600 train_time:26049ms step_avg:39.65ms step:658/1600 train_time:26108ms step_avg:39.68ms step:659/1600 train_time:26173ms step_avg:39.72ms step:660/1600 train_time:26231ms step_avg:39.74ms step:661/1600 train_time:26293ms step_avg:39.78ms step:662/1600 train_time:26353ms step_avg:39.81ms step:663/1600 train_time:26416ms step_avg:39.84ms step:664/1600 train_time:26475ms step_avg:39.87ms step:665/1600 train_time:26538ms step_avg:39.91ms step:666/1600 train_time:26597ms step_avg:39.93ms step:667/1600 train_time:26659ms step_avg:39.97ms step:668/1600 train_time:26717ms step_avg:40.00ms step:669/1600 train_time:26780ms step_avg:40.03ms step:670/1600 train_time:26839ms step_avg:40.06ms step:671/1600 train_time:26900ms step_avg:40.09ms step:672/1600 train_time:26959ms step_avg:40.12ms step:673/1600 train_time:27021ms step_avg:40.15ms step:674/1600 train_time:27080ms step_avg:40.18ms step:675/1600 train_time:27142ms step_avg:40.21ms step:676/1600 train_time:27202ms step_avg:40.24ms step:677/1600 train_time:27265ms step_avg:40.27ms step:678/1600 train_time:27325ms step_avg:40.30ms step:679/1600 train_time:27388ms step_avg:40.34ms step:680/1600 train_time:27448ms step_avg:40.36ms step:681/1600 train_time:27511ms step_avg:40.40ms step:682/1600 train_time:27569ms step_avg:40.42ms step:683/1600 train_time:27632ms step_avg:40.46ms step:684/1600 train_time:27692ms step_avg:40.48ms step:685/1600 train_time:27754ms step_avg:40.52ms step:686/1600 train_time:27813ms step_avg:40.54ms step:687/1600 train_time:27875ms step_avg:40.58ms step:688/1600 train_time:27935ms step_avg:40.60ms step:689/1600 train_time:27997ms step_avg:40.63ms step:690/1600 train_time:28056ms step_avg:40.66ms step:691/1600 train_time:28118ms step_avg:40.69ms step:692/1600 train_time:28178ms step_avg:40.72ms step:693/1600 train_time:28240ms step_avg:40.75ms step:694/1600 train_time:28299ms step_avg:40.78ms step:695/1600 train_time:28362ms step_avg:40.81ms step:696/1600 train_time:28420ms step_avg:40.83ms step:697/1600 train_time:28482ms step_avg:40.86ms step:698/1600 train_time:28541ms step_avg:40.89ms step:699/1600 train_time:28604ms step_avg:40.92ms step:700/1600 train_time:28663ms step_avg:40.95ms step:701/1600 train_time:28726ms step_avg:40.98ms step:702/1600 train_time:28786ms step_avg:41.01ms step:703/1600 train_time:28848ms step_avg:41.04ms step:704/1600 train_time:28907ms step_avg:41.06ms step:705/1600 train_time:28969ms step_avg:41.09ms step:706/1600 train_time:29029ms step_avg:41.12ms step:707/1600 train_time:29092ms step_avg:41.15ms step:708/1600 train_time:29151ms step_avg:41.17ms step:709/1600 train_time:29213ms step_avg:41.20ms step:710/1600 train_time:29272ms step_avg:41.23ms step:711/1600 train_time:29335ms step_avg:41.26ms step:712/1600 train_time:29394ms step_avg:41.28ms step:713/1600 train_time:29457ms step_avg:41.31ms step:714/1600 train_time:29516ms step_avg:41.34ms step:715/1600 train_time:29578ms step_avg:41.37ms step:716/1600 train_time:29638ms step_avg:41.39ms step:717/1600 train_time:29700ms step_avg:41.42ms step:718/1600 train_time:29759ms step_avg:41.45ms step:719/1600 train_time:29821ms step_avg:41.48ms step:720/1600 train_time:29880ms step_avg:41.50ms step:721/1600 train_time:29942ms step_avg:41.53ms step:722/1600 train_time:30001ms step_avg:41.55ms step:723/1600 train_time:30065ms step_avg:41.58ms step:724/1600 train_time:30124ms step_avg:41.61ms step:725/1600 train_time:30186ms step_avg:41.64ms step:726/1600 train_time:30246ms step_avg:41.66ms step:727/1600 train_time:30308ms step_avg:41.69ms step:728/1600 train_time:30368ms step_avg:41.71ms step:729/1600 train_time:30434ms step_avg:41.75ms step:730/1600 train_time:30491ms step_avg:41.77ms step:731/1600 train_time:30553ms step_avg:41.80ms step:732/1600 train_time:30612ms step_avg:41.82ms step:733/1600 train_time:30677ms step_avg:41.85ms step:734/1600 train_time:30736ms step_avg:41.87ms step:735/1600 train_time:30798ms step_avg:41.90ms step:736/1600 train_time:30857ms step_avg:41.92ms step:737/1600 train_time:30919ms step_avg:41.95ms step:738/1600 train_time:30978ms step_avg:41.98ms step:739/1600 train_time:31041ms step_avg:42.00ms step:740/1600 train_time:31100ms step_avg:42.03ms step:741/1600 train_time:31162ms step_avg:42.05ms step:742/1600 train_time:31221ms step_avg:42.08ms step:743/1600 train_time:31284ms step_avg:42.10ms step:744/1600 train_time:31343ms step_avg:42.13ms step:745/1600 train_time:31406ms step_avg:42.16ms step:746/1600 train_time:31466ms step_avg:42.18ms step:747/1600 train_time:31528ms step_avg:42.21ms step:748/1600 train_time:31588ms step_avg:42.23ms step:749/1600 train_time:31650ms step_avg:42.26ms step:750/1600 train_time:31709ms step_avg:42.28ms step:750/1600 val_loss:3.8933 train_time:31756ms step_avg:42.34ms step:751/1600 train_time:31778ms step_avg:42.31ms step:752/1600 train_time:31834ms step_avg:42.33ms step:753/1600 train_time:31897ms step_avg:42.36ms step:754/1600 train_time:31958ms step_avg:42.38ms step:755/1600 train_time:32020ms step_avg:42.41ms step:756/1600 train_time:32080ms step_avg:42.43ms step:757/1600 train_time:32142ms step_avg:42.46ms step:758/1600 train_time:32201ms step_avg:42.48ms step:759/1600 train_time:32265ms step_avg:42.51ms step:760/1600 train_time:32323ms step_avg:42.53ms step:761/1600 train_time:32386ms step_avg:42.56ms step:762/1600 train_time:32445ms step_avg:42.58ms step:763/1600 train_time:32505ms step_avg:42.60ms step:764/1600 train_time:32564ms step_avg:42.62ms step:765/1600 train_time:32625ms step_avg:42.65ms step:766/1600 train_time:32684ms step_avg:42.67ms step:767/1600 train_time:32747ms step_avg:42.69ms step:768/1600 train_time:32808ms step_avg:42.72ms step:769/1600 train_time:32870ms step_avg:42.74ms step:770/1600 train_time:32930ms step_avg:42.77ms step:771/1600 train_time:32994ms step_avg:42.79ms step:772/1600 train_time:33053ms step_avg:42.82ms step:773/1600 train_time:33116ms step_avg:42.84ms step:774/1600 train_time:33174ms step_avg:42.86ms step:775/1600 train_time:33237ms step_avg:42.89ms step:776/1600 train_time:33295ms step_avg:42.91ms step:777/1600 train_time:33357ms step_avg:42.93ms step:778/1600 train_time:33416ms step_avg:42.95ms step:779/1600 train_time:33479ms step_avg:42.98ms step:780/1600 train_time:33537ms step_avg:43.00ms step:781/1600 train_time:33599ms step_avg:43.02ms step:782/1600 train_time:33658ms step_avg:43.04ms step:783/1600 train_time:33721ms step_avg:43.07ms step:784/1600 train_time:33780ms step_avg:43.09ms step:785/1600 train_time:33843ms step_avg:43.11ms step:786/1600 train_time:33903ms step_avg:43.13ms step:787/1600 train_time:33966ms step_avg:43.16ms step:788/1600 train_time:34025ms step_avg:43.18ms step:789/1600 train_time:34089ms step_avg:43.21ms step:790/1600 train_time:34146ms step_avg:43.22ms step:791/1600 train_time:34208ms step_avg:43.25ms step:792/1600 train_time:34267ms step_avg:43.27ms step:793/1600 train_time:34329ms step_avg:43.29ms step:794/1600 train_time:34388ms step_avg:43.31ms step:795/1600 train_time:34450ms step_avg:43.33ms step:796/1600 train_time:34509ms step_avg:43.35ms step:797/1600 train_time:34572ms step_avg:43.38ms step:798/1600 train_time:34631ms step_avg:43.40ms step:799/1600 train_time:34693ms step_avg:43.42ms step:800/1600 train_time:34752ms step_avg:43.44ms step:801/1600 train_time:34815ms step_avg:43.46ms step:802/1600 train_time:34874ms step_avg:43.48ms step:803/1600 train_time:34938ms step_avg:43.51ms step:804/1600 train_time:34996ms step_avg:43.53ms step:805/1600 train_time:35059ms step_avg:43.55ms step:806/1600 train_time:35118ms step_avg:43.57ms step:807/1600 train_time:35181ms step_avg:43.60ms step:808/1600 train_time:35242ms step_avg:43.62ms step:809/1600 train_time:35304ms step_avg:43.64ms step:810/1600 train_time:35363ms step_avg:43.66ms step:811/1600 train_time:35425ms step_avg:43.68ms step:812/1600 train_time:35484ms step_avg:43.70ms step:813/1600 train_time:35546ms step_avg:43.72ms step:814/1600 train_time:35605ms step_avg:43.74ms step:815/1600 train_time:35667ms step_avg:43.76ms step:816/1600 train_time:35726ms step_avg:43.78ms step:817/1600 train_time:35788ms step_avg:43.80ms step:818/1600 train_time:35847ms step_avg:43.82ms step:819/1600 train_time:35910ms step_avg:43.85ms step:820/1600 train_time:35969ms step_avg:43.86ms step:821/1600 train_time:36032ms step_avg:43.89ms step:822/1600 train_time:36092ms step_avg:43.91ms step:823/1600 train_time:36154ms step_avg:43.93ms step:824/1600 train_time:36213ms step_avg:43.95ms step:825/1600 train_time:36276ms step_avg:43.97ms step:826/1600 train_time:36334ms step_avg:43.99ms step:827/1600 train_time:36396ms step_avg:44.01ms step:828/1600 train_time:36454ms step_avg:44.03ms step:829/1600 train_time:36517ms step_avg:44.05ms step:830/1600 train_time:36576ms step_avg:44.07ms step:831/1600 train_time:36642ms step_avg:44.09ms step:832/1600 train_time:36699ms step_avg:44.11ms step:833/1600 train_time:36762ms step_avg:44.13ms step:834/1600 train_time:36821ms step_avg:44.15ms step:835/1600 train_time:36883ms step_avg:44.17ms step:836/1600 train_time:36943ms step_avg:44.19ms step:837/1600 train_time:37006ms step_avg:44.21ms step:838/1600 train_time:37065ms step_avg:44.23ms step:839/1600 train_time:37127ms step_avg:44.25ms step:840/1600 train_time:37186ms step_avg:44.27ms step:841/1600 train_time:37248ms step_avg:44.29ms step:842/1600 train_time:37307ms step_avg:44.31ms step:843/1600 train_time:37369ms step_avg:44.33ms step:844/1600 train_time:37428ms step_avg:44.35ms step:845/1600 train_time:37490ms step_avg:44.37ms step:846/1600 train_time:37549ms step_avg:44.38ms step:847/1600 train_time:37611ms step_avg:44.41ms step:848/1600 train_time:37672ms step_avg:44.42ms step:849/1600 train_time:37734ms step_avg:44.45ms step:850/1600 train_time:37794ms step_avg:44.46ms step:851/1600 train_time:37856ms step_avg:44.48ms step:852/1600 train_time:37915ms step_avg:44.50ms step:853/1600 train_time:37978ms step_avg:44.52ms step:854/1600 train_time:38037ms step_avg:44.54ms step:855/1600 train_time:38100ms step_avg:44.56ms step:856/1600 train_time:38160ms step_avg:44.58ms step:857/1600 train_time:38222ms step_avg:44.60ms step:858/1600 train_time:38281ms step_avg:44.62ms step:859/1600 train_time:38344ms step_avg:44.64ms step:860/1600 train_time:38403ms step_avg:44.65ms step:861/1600 train_time:38465ms step_avg:44.67ms step:862/1600 train_time:38524ms step_avg:44.69ms step:863/1600 train_time:38586ms step_avg:44.71ms step:864/1600 train_time:38645ms step_avg:44.73ms step:865/1600 train_time:38708ms step_avg:44.75ms step:866/1600 train_time:38766ms step_avg:44.76ms step:867/1600 train_time:38828ms step_avg:44.78ms step:868/1600 train_time:38887ms step_avg:44.80ms step:869/1600 train_time:38950ms step_avg:44.82ms step:870/1600 train_time:39009ms step_avg:44.84ms step:871/1600 train_time:39072ms step_avg:44.86ms step:872/1600 train_time:39132ms step_avg:44.88ms step:873/1600 train_time:39194ms step_avg:44.90ms step:874/1600 train_time:39253ms step_avg:44.91ms step:875/1600 train_time:39315ms step_avg:44.93ms step:876/1600 train_time:39374ms step_avg:44.95ms step:877/1600 train_time:39436ms step_avg:44.97ms step:878/1600 train_time:39495ms step_avg:44.98ms step:879/1600 train_time:39557ms step_avg:45.00ms step:880/1600 train_time:39617ms step_avg:45.02ms step:881/1600 train_time:39680ms step_avg:45.04ms step:882/1600 train_time:39739ms step_avg:45.06ms step:883/1600 train_time:39802ms step_avg:45.08ms step:884/1600 train_time:39861ms step_avg:45.09ms step:885/1600 train_time:39925ms step_avg:45.11ms step:886/1600 train_time:39983ms step_avg:45.13ms step:887/1600 train_time:40046ms step_avg:45.15ms step:888/1600 train_time:40105ms step_avg:45.16ms step:889/1600 train_time:40167ms step_avg:45.18ms step:890/1600 train_time:40231ms step_avg:45.20ms step:891/1600 train_time:40289ms step_avg:45.22ms step:892/1600 train_time:40348ms step_avg:45.23ms step:893/1600 train_time:40412ms step_avg:45.25ms step:894/1600 train_time:40469ms step_avg:45.27ms step:895/1600 train_time:40532ms step_avg:45.29ms step:896/1600 train_time:40591ms step_avg:45.30ms step:897/1600 train_time:40653ms step_avg:45.32ms step:898/1600 train_time:40712ms step_avg:45.34ms step:899/1600 train_time:40774ms step_avg:45.36ms step:900/1600 train_time:40833ms step_avg:45.37ms step:901/1600 train_time:40897ms step_avg:45.39ms step:902/1600 train_time:40955ms step_avg:45.40ms step:903/1600 train_time:41018ms step_avg:45.42ms step:904/1600 train_time:41077ms step_avg:45.44ms step:905/1600 train_time:41139ms step_avg:45.46ms step:906/1600 train_time:41198ms step_avg:45.47ms step:907/1600 train_time:41261ms step_avg:45.49ms step:908/1600 train_time:41320ms step_avg:45.51ms step:909/1600 train_time:41383ms step_avg:45.53ms step:910/1600 train_time:41442ms step_avg:45.54ms step:911/1600 train_time:41504ms step_avg:45.56ms step:912/1600 train_time:41563ms step_avg:45.57ms step:913/1600 train_time:41626ms step_avg:45.59ms step:914/1600 train_time:41685ms step_avg:45.61ms step:915/1600 train_time:41747ms step_avg:45.63ms step:916/1600 train_time:41806ms step_avg:45.64ms step:917/1600 train_time:41868ms step_avg:45.66ms step:918/1600 train_time:41927ms step_avg:45.67ms step:919/1600 train_time:41990ms step_avg:45.69ms step:920/1600 train_time:42049ms step_avg:45.71ms step:921/1600 train_time:42111ms step_avg:45.72ms step:922/1600 train_time:42171ms step_avg:45.74ms step:923/1600 train_time:42233ms step_avg:45.76ms step:924/1600 train_time:42293ms step_avg:45.77ms step:925/1600 train_time:42356ms step_avg:45.79ms step:926/1600 train_time:42414ms step_avg:45.80ms step:927/1600 train_time:42476ms step_avg:45.82ms step:928/1600 train_time:42536ms step_avg:45.84ms step:929/1600 train_time:42598ms step_avg:45.85ms step:930/1600 train_time:42657ms step_avg:45.87ms step:931/1600 train_time:42719ms step_avg:45.89ms step:932/1600 train_time:42778ms step_avg:45.90ms step:933/1600 train_time:42841ms step_avg:45.92ms step:934/1600 train_time:42900ms step_avg:45.93ms step:935/1600 train_time:42963ms step_avg:45.95ms step:936/1600 train_time:43023ms step_avg:45.96ms step:937/1600 train_time:43085ms step_avg:45.98ms step:938/1600 train_time:43144ms step_avg:46.00ms step:939/1600 train_time:43207ms step_avg:46.01ms step:940/1600 train_time:43267ms step_avg:46.03ms step:941/1600 train_time:43329ms step_avg:46.05ms step:942/1600 train_time:43387ms step_avg:46.06ms step:943/1600 train_time:43450ms step_avg:46.08ms step:944/1600 train_time:43510ms step_avg:46.09ms step:945/1600 train_time:43572ms step_avg:46.11ms step:946/1600 train_time:43631ms step_avg:46.12ms step:947/1600 train_time:43694ms step_avg:46.14ms step:948/1600 train_time:43752ms step_avg:46.15ms step:949/1600 train_time:43815ms step_avg:46.17ms step:950/1600 train_time:43874ms step_avg:46.18ms step:951/1600 train_time:43936ms step_avg:46.20ms step:952/1600 train_time:43995ms step_avg:46.21ms step:953/1600 train_time:44058ms step_avg:46.23ms step:954/1600 train_time:44117ms step_avg:46.24ms step:955/1600 train_time:44180ms step_avg:46.26ms step:956/1600 train_time:44240ms step_avg:46.28ms step:957/1600 train_time:44303ms step_avg:46.29ms step:958/1600 train_time:44362ms step_avg:46.31ms step:959/1600 train_time:44425ms step_avg:46.32ms step:960/1600 train_time:44484ms step_avg:46.34ms step:961/1600 train_time:44547ms step_avg:46.35ms step:962/1600 train_time:44606ms step_avg:46.37ms step:963/1600 train_time:44668ms step_avg:46.38ms step:964/1600 train_time:44727ms step_avg:46.40ms step:965/1600 train_time:44789ms step_avg:46.41ms step:966/1600 train_time:44847ms step_avg:46.43ms step:967/1600 train_time:44910ms step_avg:46.44ms step:968/1600 train_time:44970ms step_avg:46.46ms step:969/1600 train_time:45032ms step_avg:46.47ms step:970/1600 train_time:45092ms step_avg:46.49ms step:971/1600 train_time:45156ms step_avg:46.50ms step:972/1600 train_time:45215ms step_avg:46.52ms step:973/1600 train_time:45277ms step_avg:46.53ms step:974/1600 train_time:45336ms step_avg:46.55ms step:975/1600 train_time:45398ms step_avg:46.56ms step:976/1600 train_time:45457ms step_avg:46.57ms step:977/1600 train_time:45523ms step_avg:46.60ms step:978/1600 train_time:45582ms step_avg:46.61ms step:979/1600 train_time:45642ms step_avg:46.62ms step:980/1600 train_time:45701ms step_avg:46.63ms step:981/1600 train_time:45765ms step_avg:46.65ms step:982/1600 train_time:45824ms step_avg:46.66ms step:983/1600 train_time:45887ms step_avg:46.68ms step:984/1600 train_time:45946ms step_avg:46.69ms step:985/1600 train_time:46008ms step_avg:46.71ms step:986/1600 train_time:46067ms step_avg:46.72ms step:987/1600 train_time:46129ms step_avg:46.74ms step:988/1600 train_time:46189ms step_avg:46.75ms step:989/1600 train_time:46252ms step_avg:46.77ms step:990/1600 train_time:46312ms step_avg:46.78ms step:991/1600 train_time:46375ms step_avg:46.80ms step:992/1600 train_time:46434ms step_avg:46.81ms step:993/1600 train_time:46497ms step_avg:46.82ms step:994/1600 train_time:46555ms step_avg:46.84ms step:995/1600 train_time:46617ms step_avg:46.85ms step:996/1600 train_time:46677ms step_avg:46.86ms step:997/1600 train_time:46743ms step_avg:46.88ms step:998/1600 train_time:46800ms step_avg:46.89ms step:999/1600 train_time:46864ms step_avg:46.91ms step:1000/1600 train_time:46922ms step_avg:46.92ms step:1000/1600 val_loss:3.5972 train_time:46968ms step_avg:46.97ms step:1001/1600 train_time:46989ms step_avg:46.94ms step:1002/1600 train_time:47046ms step_avg:46.95ms step:1003/1600 train_time:47110ms step_avg:46.97ms step:1004/1600 train_time:47169ms step_avg:46.98ms step:1005/1600 train_time:47231ms step_avg:47.00ms step:1006/1600 train_time:47290ms step_avg:47.01ms step:1007/1600 train_time:47352ms step_avg:47.02ms step:1008/1600 train_time:47410ms step_avg:47.03ms step:1009/1600 train_time:47472ms step_avg:47.05ms step:1010/1600 train_time:47530ms step_avg:47.06ms step:1011/1600 train_time:47593ms step_avg:47.07ms step:1012/1600 train_time:47651ms step_avg:47.09ms step:1013/1600 train_time:47713ms step_avg:47.10ms step:1014/1600 train_time:47772ms step_avg:47.11ms step:1015/1600 train_time:47833ms step_avg:47.13ms step:1016/1600 train_time:47893ms step_avg:47.14ms step:1017/1600 train_time:47958ms step_avg:47.16ms step:1018/1600 train_time:48019ms step_avg:47.17ms step:1019/1600 train_time:48082ms step_avg:47.19ms step:1020/1600 train_time:48141ms step_avg:47.20ms step:1021/1600 train_time:48203ms step_avg:47.21ms step:1022/1600 train_time:48262ms step_avg:47.22ms step:1023/1600 train_time:48325ms step_avg:47.24ms step:1024/1600 train_time:48384ms step_avg:47.25ms step:1025/1600 train_time:48446ms step_avg:47.26ms step:1026/1600 train_time:48504ms step_avg:47.28ms step:1027/1600 train_time:48567ms step_avg:47.29ms step:1028/1600 train_time:48626ms step_avg:47.30ms step:1029/1600 train_time:48688ms step_avg:47.32ms step:1030/1600 train_time:48748ms step_avg:47.33ms step:1031/1600 train_time:48810ms step_avg:47.34ms step:1032/1600 train_time:48868ms step_avg:47.35ms step:1033/1600 train_time:48932ms step_avg:47.37ms step:1034/1600 train_time:48991ms step_avg:47.38ms step:1035/1600 train_time:49053ms step_avg:47.39ms step:1036/1600 train_time:49112ms step_avg:47.41ms step:1037/1600 train_time:49174ms step_avg:47.42ms step:1038/1600 train_time:49233ms step_avg:47.43ms step:1039/1600 train_time:49296ms step_avg:47.45ms step:1040/1600 train_time:49355ms step_avg:47.46ms step:1041/1600 train_time:49426ms step_avg:47.48ms step:1042/1600 train_time:49509ms step_avg:47.51ms step:1043/1600 train_time:49598ms step_avg:47.55ms step:1044/1600 train_time:49682ms step_avg:47.59ms step:1045/1600 train_time:49771ms step_avg:47.63ms step:1046/1600 train_time:49855ms step_avg:47.66ms step:1047/1600 train_time:49945ms step_avg:47.70ms step:1048/1600 train_time:50031ms step_avg:47.74ms step:1049/1600 train_time:50121ms step_avg:47.78ms step:1050/1600 train_time:50207ms step_avg:47.82ms step:1051/1600 train_time:50296ms step_avg:47.85ms step:1052/1600 train_time:50381ms step_avg:47.89ms step:1053/1600 train_time:50469ms step_avg:47.93ms step:1054/1600 train_time:50554ms step_avg:47.96ms step:1055/1600 train_time:50642ms step_avg:48.00ms step:1056/1600 train_time:50729ms step_avg:48.04ms step:1057/1600 train_time:50816ms step_avg:48.08ms step:1058/1600 train_time:50900ms step_avg:48.11ms step:1059/1600 train_time:50989ms step_avg:48.15ms step:1060/1600 train_time:51074ms step_avg:48.18ms step:1061/1600 train_time:51164ms step_avg:48.22ms step:1062/1600 train_time:51249ms step_avg:48.26ms step:1063/1600 train_time:51337ms step_avg:48.29ms step:1064/1600 train_time:51422ms step_avg:48.33ms step:1065/1600 train_time:51509ms step_avg:48.37ms step:1066/1600 train_time:51594ms step_avg:48.40ms step:1067/1600 train_time:51683ms step_avg:48.44ms step:1068/1600 train_time:51768ms step_avg:48.47ms step:1069/1600 train_time:51857ms step_avg:48.51ms step:1070/1600 train_time:51941ms step_avg:48.54ms step:1071/1600 train_time:52029ms step_avg:48.58ms step:1072/1600 train_time:52115ms step_avg:48.61ms step:1073/1600 train_time:52204ms step_avg:48.65ms step:1074/1600 train_time:52289ms step_avg:48.69ms step:1075/1600 train_time:52377ms step_avg:48.72ms step:1076/1600 train_time:52462ms step_avg:48.76ms step:1077/1600 train_time:52551ms step_avg:48.79ms step:1078/1600 train_time:52635ms step_avg:48.83ms step:1079/1600 train_time:52724ms step_avg:48.86ms step:1080/1600 train_time:52809ms step_avg:48.90ms step:1081/1600 train_time:52897ms step_avg:48.93ms step:1082/1600 train_time:52983ms step_avg:48.97ms step:1083/1600 train_time:53071ms step_avg:49.00ms step:1084/1600 train_time:53156ms step_avg:49.04ms step:1085/1600 train_time:53245ms step_avg:49.07ms step:1086/1600 train_time:53331ms step_avg:49.11ms step:1087/1600 train_time:53420ms step_avg:49.14ms step:1088/1600 train_time:53504ms step_avg:49.18ms step:1089/1600 train_time:53592ms step_avg:49.21ms step:1090/1600 train_time:53678ms step_avg:49.25ms step:1091/1600 train_time:53766ms step_avg:49.28ms step:1092/1600 train_time:53851ms step_avg:49.31ms step:1093/1600 train_time:53939ms step_avg:49.35ms step:1094/1600 train_time:54023ms step_avg:49.38ms step:1095/1600 train_time:54112ms step_avg:49.42ms step:1096/1600 train_time:54198ms step_avg:49.45ms step:1097/1600 train_time:54287ms step_avg:49.49ms step:1098/1600 train_time:54372ms step_avg:49.52ms step:1099/1600 train_time:54461ms step_avg:49.55ms step:1100/1600 train_time:54546ms step_avg:49.59ms step:1101/1600 train_time:54635ms step_avg:49.62ms step:1102/1600 train_time:54719ms step_avg:49.65ms step:1103/1600 train_time:54808ms step_avg:49.69ms step:1104/1600 train_time:54893ms step_avg:49.72ms step:1105/1600 train_time:54982ms step_avg:49.76ms step:1106/1600 train_time:55067ms step_avg:49.79ms step:1107/1600 train_time:55156ms step_avg:49.82ms step:1108/1600 train_time:55240ms step_avg:49.86ms step:1109/1600 train_time:55329ms step_avg:49.89ms step:1110/1600 train_time:55414ms step_avg:49.92ms step:1111/1600 train_time:55503ms step_avg:49.96ms step:1112/1600 train_time:55588ms step_avg:49.99ms step:1113/1600 train_time:55676ms step_avg:50.02ms step:1114/1600 train_time:55761ms step_avg:50.05ms step:1115/1600 train_time:55851ms step_avg:50.09ms step:1116/1600 train_time:55935ms step_avg:50.12ms step:1117/1600 train_time:56024ms step_avg:50.16ms step:1118/1600 train_time:56109ms step_avg:50.19ms step:1119/1600 train_time:56198ms step_avg:50.22ms step:1120/1600 train_time:56282ms step_avg:50.25ms step:1121/1600 train_time:56369ms step_avg:50.28ms step:1122/1600 train_time:56455ms step_avg:50.32ms step:1123/1600 train_time:56543ms step_avg:50.35ms step:1124/1600 train_time:56629ms step_avg:50.38ms step:1125/1600 train_time:56717ms step_avg:50.42ms step:1126/1600 train_time:56803ms step_avg:50.45ms step:1127/1600 train_time:56891ms step_avg:50.48ms step:1128/1600 train_time:56977ms step_avg:50.51ms step:1129/1600 train_time:57065ms step_avg:50.54ms step:1130/1600 train_time:57150ms step_avg:50.58ms step:1131/1600 train_time:57238ms step_avg:50.61ms step:1132/1600 train_time:57323ms step_avg:50.64ms step:1133/1600 train_time:57411ms step_avg:50.67ms step:1134/1600 train_time:57497ms step_avg:50.70ms step:1135/1600 train_time:57585ms step_avg:50.74ms step:1136/1600 train_time:57670ms step_avg:50.77ms step:1137/1600 train_time:57758ms step_avg:50.80ms step:1138/1600 train_time:57843ms step_avg:50.83ms step:1139/1600 train_time:57933ms step_avg:50.86ms step:1140/1600 train_time:58017ms step_avg:50.89ms step:1141/1600 train_time:58106ms step_avg:50.93ms step:1142/1600 train_time:58190ms step_avg:50.95ms step:1143/1600 train_time:58279ms step_avg:50.99ms step:1144/1600 train_time:58363ms step_avg:51.02ms step:1145/1600 train_time:58453ms step_avg:51.05ms step:1146/1600 train_time:58538ms step_avg:51.08ms step:1147/1600 train_time:58626ms step_avg:51.11ms step:1148/1600 train_time:58711ms step_avg:51.14ms step:1149/1600 train_time:58805ms step_avg:51.18ms step:1150/1600 train_time:58889ms step_avg:51.21ms step:1151/1600 train_time:58973ms step_avg:51.24ms step:1152/1600 train_time:59058ms step_avg:51.27ms step:1153/1600 train_time:59146ms step_avg:51.30ms step:1154/1600 train_time:59235ms step_avg:51.33ms step:1155/1600 train_time:59323ms step_avg:51.36ms step:1156/1600 train_time:59408ms step_avg:51.39ms step:1157/1600 train_time:59496ms step_avg:51.42ms step:1158/1600 train_time:59581ms step_avg:51.45ms step:1159/1600 train_time:59668ms step_avg:51.48ms step:1160/1600 train_time:59754ms step_avg:51.51ms step:1161/1600 train_time:59842ms step_avg:51.54ms step:1162/1600 train_time:59927ms step_avg:51.57ms step:1163/1600 train_time:60018ms step_avg:51.61ms step:1164/1600 train_time:60101ms step_avg:51.63ms step:1165/1600 train_time:60190ms step_avg:51.66ms step:1166/1600 train_time:60277ms step_avg:51.70ms step:1167/1600 train_time:60364ms step_avg:51.73ms step:1168/1600 train_time:60449ms step_avg:51.75ms step:1169/1600 train_time:60537ms step_avg:51.78ms step:1170/1600 train_time:60622ms step_avg:51.81ms step:1171/1600 train_time:60710ms step_avg:51.84ms step:1172/1600 train_time:60795ms step_avg:51.87ms step:1173/1600 train_time:60887ms step_avg:51.91ms step:1174/1600 train_time:60969ms step_avg:51.93ms step:1175/1600 train_time:61057ms step_avg:51.96ms step:1176/1600 train_time:61142ms step_avg:51.99ms step:1177/1600 train_time:61229ms step_avg:52.02ms step:1178/1600 train_time:61314ms step_avg:52.05ms step:1179/1600 train_time:61403ms step_avg:52.08ms step:1180/1600 train_time:61488ms step_avg:52.11ms step:1181/1600 train_time:61577ms step_avg:52.14ms step:1182/1600 train_time:61662ms step_avg:52.17ms step:1183/1600 train_time:61750ms step_avg:52.20ms step:1184/1600 train_time:61835ms step_avg:52.23ms step:1185/1600 train_time:61922ms step_avg:52.25ms step:1186/1600 train_time:62007ms step_avg:52.28ms step:1187/1600 train_time:62095ms step_avg:52.31ms step:1188/1600 train_time:62180ms step_avg:52.34ms step:1189/1600 train_time:62269ms step_avg:52.37ms step:1190/1600 train_time:62354ms step_avg:52.40ms step:1191/1600 train_time:62443ms step_avg:52.43ms step:1192/1600 train_time:62528ms step_avg:52.46ms step:1193/1600 train_time:62617ms step_avg:52.49ms step:1194/1600 train_time:62702ms step_avg:52.51ms step:1195/1600 train_time:62790ms step_avg:52.54ms step:1196/1600 train_time:62876ms step_avg:52.57ms step:1197/1600 train_time:62964ms step_avg:52.60ms step:1198/1600 train_time:63050ms step_avg:52.63ms step:1199/1600 train_time:63137ms step_avg:52.66ms step:1200/1600 train_time:63223ms step_avg:52.69ms step:1201/1600 train_time:63311ms step_avg:52.72ms step:1202/1600 train_time:63396ms step_avg:52.74ms step:1203/1600 train_time:63484ms step_avg:52.77ms step:1204/1600 train_time:63570ms step_avg:52.80ms step:1205/1600 train_time:63658ms step_avg:52.83ms step:1206/1600 train_time:63744ms step_avg:52.86ms step:1207/1600 train_time:63833ms step_avg:52.89ms step:1208/1600 train_time:63919ms step_avg:52.91ms step:1209/1600 train_time:64008ms step_avg:52.94ms step:1210/1600 train_time:64093ms step_avg:52.97ms step:1211/1600 train_time:64182ms step_avg:53.00ms step:1212/1600 train_time:64267ms step_avg:53.03ms step:1213/1600 train_time:64355ms step_avg:53.05ms step:1214/1600 train_time:64441ms step_avg:53.08ms step:1215/1600 train_time:64529ms step_avg:53.11ms step:1216/1600 train_time:64614ms step_avg:53.14ms step:1217/1600 train_time:64703ms step_avg:53.17ms step:1218/1600 train_time:64788ms step_avg:53.19ms step:1219/1600 train_time:64877ms step_avg:53.22ms step:1220/1600 train_time:64962ms step_avg:53.25ms step:1221/1600 train_time:65051ms step_avg:53.28ms step:1222/1600 train_time:65136ms step_avg:53.30ms step:1223/1600 train_time:65223ms step_avg:53.33ms step:1224/1600 train_time:65309ms step_avg:53.36ms step:1225/1600 train_time:65398ms step_avg:53.39ms step:1226/1600 train_time:65482ms step_avg:53.41ms step:1227/1600 train_time:65571ms step_avg:53.44ms step:1228/1600 train_time:65656ms step_avg:53.47ms step:1229/1600 train_time:65747ms step_avg:53.50ms step:1230/1600 train_time:65831ms step_avg:53.52ms step:1231/1600 train_time:65919ms step_avg:53.55ms step:1232/1600 train_time:66004ms step_avg:53.57ms step:1233/1600 train_time:66093ms step_avg:53.60ms step:1234/1600 train_time:66178ms step_avg:53.63ms step:1235/1600 train_time:66266ms step_avg:53.66ms step:1236/1600 train_time:66351ms step_avg:53.68ms step:1237/1600 train_time:66440ms step_avg:53.71ms step:1238/1600 train_time:66525ms step_avg:53.74ms step:1239/1600 train_time:66613ms step_avg:53.76ms step:1240/1600 train_time:66698ms step_avg:53.79ms step:1241/1600 train_time:66786ms step_avg:53.82ms step:1242/1600 train_time:66872ms step_avg:53.84ms step:1243/1600 train_time:66960ms step_avg:53.87ms step:1244/1600 train_time:67045ms step_avg:53.89ms step:1245/1600 train_time:67135ms step_avg:53.92ms step:1246/1600 train_time:67220ms step_avg:53.95ms step:1247/1600 train_time:67308ms step_avg:53.98ms step:1248/1600 train_time:67393ms step_avg:54.00ms step:1249/1600 train_time:67481ms step_avg:54.03ms step:1250/1600 train_time:67566ms step_avg:54.05ms step:1250/1600 val_loss:3.4174 train_time:67638ms step_avg:54.11ms step:1251/1600 train_time:67658ms step_avg:54.08ms step:1252/1600 train_time:67744ms step_avg:54.11ms step:1253/1600 train_time:67839ms step_avg:54.14ms step:1254/1600 train_time:67924ms step_avg:54.17ms step:1255/1600 train_time:68014ms step_avg:54.19ms step:1256/1600 train_time:68101ms step_avg:54.22ms step:1257/1600 train_time:68184ms step_avg:54.24ms step:1258/1600 train_time:68268ms step_avg:54.27ms step:1259/1600 train_time:68355ms step_avg:54.29ms step:1260/1600 train_time:68440ms step_avg:54.32ms step:1261/1600 train_time:68528ms step_avg:54.34ms step:1262/1600 train_time:68613ms step_avg:54.37ms step:1263/1600 train_time:68704ms step_avg:54.40ms step:1264/1600 train_time:68791ms step_avg:54.42ms step:1265/1600 train_time:68880ms step_avg:54.45ms step:1266/1600 train_time:68967ms step_avg:54.48ms step:1267/1600 train_time:69055ms step_avg:54.50ms step:1268/1600 train_time:69140ms step_avg:54.53ms step:1269/1600 train_time:69227ms step_avg:54.55ms step:1270/1600 train_time:69312ms step_avg:54.58ms step:1271/1600 train_time:69403ms step_avg:54.60ms step:1272/1600 train_time:69484ms step_avg:54.63ms step:1273/1600 train_time:69572ms step_avg:54.65ms step:1274/1600 train_time:69659ms step_avg:54.68ms step:1275/1600 train_time:69748ms step_avg:54.70ms step:1276/1600 train_time:69834ms step_avg:54.73ms step:1277/1600 train_time:69923ms step_avg:54.76ms step:1278/1600 train_time:70009ms step_avg:54.78ms step:1279/1600 train_time:70097ms step_avg:54.81ms step:1280/1600 train_time:70182ms step_avg:54.83ms step:1281/1600 train_time:70270ms step_avg:54.86ms step:1282/1600 train_time:70354ms step_avg:54.88ms step:1283/1600 train_time:70442ms step_avg:54.90ms step:1284/1600 train_time:70527ms step_avg:54.93ms step:1285/1600 train_time:70615ms step_avg:54.95ms step:1286/1600 train_time:70701ms step_avg:54.98ms step:1287/1600 train_time:70789ms step_avg:55.00ms step:1288/1600 train_time:70875ms step_avg:55.03ms step:1289/1600 train_time:70964ms step_avg:55.05ms step:1290/1600 train_time:71049ms step_avg:55.08ms step:1291/1600 train_time:71137ms step_avg:55.10ms step:1292/1600 train_time:71221ms step_avg:55.12ms step:1293/1600 train_time:71309ms step_avg:55.15ms step:1294/1600 train_time:71394ms step_avg:55.17ms step:1295/1600 train_time:71482ms step_avg:55.20ms step:1296/1600 train_time:71567ms step_avg:55.22ms step:1297/1600 train_time:71656ms step_avg:55.25ms step:1298/1600 train_time:71742ms step_avg:55.27ms step:1299/1600 train_time:71830ms step_avg:55.30ms step:1300/1600 train_time:71917ms step_avg:55.32ms step:1301/1600 train_time:72004ms step_avg:55.34ms step:1302/1600 train_time:72089ms step_avg:55.37ms step:1303/1600 train_time:72177ms step_avg:55.39ms step:1304/1600 train_time:72262ms step_avg:55.42ms step:1305/1600 train_time:72350ms step_avg:55.44ms step:1306/1600 train_time:72435ms step_avg:55.46ms step:1307/1600 train_time:72524ms step_avg:55.49ms step:1308/1600 train_time:72609ms step_avg:55.51ms step:1309/1600 train_time:72697ms step_avg:55.54ms step:1310/1600 train_time:72782ms step_avg:55.56ms step:1311/1600 train_time:72870ms step_avg:55.58ms step:1312/1600 train_time:72955ms step_avg:55.61ms step:1313/1600 train_time:73044ms step_avg:55.63ms step:1314/1600 train_time:73129ms step_avg:55.65ms step:1315/1600 train_time:73217ms step_avg:55.68ms step:1316/1600 train_time:73302ms step_avg:55.70ms step:1317/1600 train_time:73389ms step_avg:55.72ms step:1318/1600 train_time:73474ms step_avg:55.75ms step:1319/1600 train_time:73563ms step_avg:55.77ms step:1320/1600 train_time:73648ms step_avg:55.79ms step:1321/1600 train_time:73736ms step_avg:55.82ms step:1322/1600 train_time:73822ms step_avg:55.84ms step:1323/1600 train_time:73910ms step_avg:55.87ms step:1324/1600 train_time:73996ms step_avg:55.89ms step:1325/1600 train_time:74084ms step_avg:55.91ms step:1326/1600 train_time:74170ms step_avg:55.93ms step:1327/1600 train_time:74257ms step_avg:55.96ms step:1328/1600 train_time:74343ms step_avg:55.98ms step:1329/1600 train_time:74431ms step_avg:56.00ms step:1330/1600 train_time:74518ms step_avg:56.03ms step:1331/1600 train_time:74605ms step_avg:56.05ms step:1332/1600 train_time:74690ms step_avg:56.07ms step:1333/1600 train_time:74779ms step_avg:56.10ms step:1334/1600 train_time:74863ms step_avg:56.12ms step:1335/1600 train_time:74953ms step_avg:56.14ms step:1336/1600 train_time:75040ms step_avg:56.17ms step:1337/1600 train_time:75126ms step_avg:56.19ms step:1338/1600 train_time:75210ms step_avg:56.21ms step:1339/1600 train_time:75298ms step_avg:56.23ms step:1340/1600 train_time:75384ms step_avg:56.26ms step:1341/1600 train_time:75472ms step_avg:56.28ms step:1342/1600 train_time:75558ms step_avg:56.30ms step:1343/1600 train_time:75645ms step_avg:56.33ms step:1344/1600 train_time:75730ms step_avg:56.35ms step:1345/1600 train_time:75818ms step_avg:56.37ms step:1346/1600 train_time:75903ms step_avg:56.39ms step:1347/1600 train_time:75991ms step_avg:56.42ms step:1348/1600 train_time:76077ms step_avg:56.44ms step:1349/1600 train_time:76165ms step_avg:56.46ms step:1350/1600 train_time:76251ms step_avg:56.48ms step:1351/1600 train_time:76339ms step_avg:56.51ms step:1352/1600 train_time:76423ms step_avg:56.53ms step:1353/1600 train_time:76512ms step_avg:56.55ms step:1354/1600 train_time:76598ms step_avg:56.57ms step:1355/1600 train_time:76686ms step_avg:56.59ms step:1356/1600 train_time:76771ms step_avg:56.62ms step:1357/1600 train_time:76860ms step_avg:56.64ms step:1358/1600 train_time:76946ms step_avg:56.66ms step:1359/1600 train_time:77033ms step_avg:56.68ms step:1360/1600 train_time:77119ms step_avg:56.70ms step:1361/1600 train_time:77206ms step_avg:56.73ms step:1362/1600 train_time:77292ms step_avg:56.75ms step:1363/1600 train_time:77380ms step_avg:56.77ms step:1364/1600 train_time:77465ms step_avg:56.79ms step:1365/1600 train_time:77553ms step_avg:56.82ms step:1366/1600 train_time:77642ms step_avg:56.84ms step:1367/1600 train_time:77728ms step_avg:56.86ms step:1368/1600 train_time:77813ms step_avg:56.88ms step:1369/1600 train_time:77902ms step_avg:56.90ms step:1370/1600 train_time:77987ms step_avg:56.92ms step:1371/1600 train_time:78076ms step_avg:56.95ms step:1372/1600 train_time:78161ms step_avg:56.97ms step:1373/1600 train_time:78249ms step_avg:56.99ms step:1374/1600 train_time:78335ms step_avg:57.01ms step:1375/1600 train_time:78424ms step_avg:57.04ms step:1376/1600 train_time:78509ms step_avg:57.06ms step:1377/1600 train_time:78598ms step_avg:57.08ms step:1378/1600 train_time:78683ms step_avg:57.10ms step:1379/1600 train_time:78772ms step_avg:57.12ms step:1380/1600 train_time:78858ms step_avg:57.14ms step:1381/1600 train_time:78945ms step_avg:57.17ms step:1382/1600 train_time:79031ms step_avg:57.19ms step:1383/1600 train_time:79119ms step_avg:57.21ms step:1384/1600 train_time:79204ms step_avg:57.23ms step:1385/1600 train_time:79294ms step_avg:57.25ms step:1386/1600 train_time:79378ms step_avg:57.27ms step:1387/1600 train_time:79467ms step_avg:57.29ms step:1388/1600 train_time:79552ms step_avg:57.31ms step:1389/1600 train_time:79641ms step_avg:57.34ms step:1390/1600 train_time:79726ms step_avg:57.36ms step:1391/1600 train_time:79815ms step_avg:57.38ms step:1392/1600 train_time:79903ms step_avg:57.40ms step:1393/1600 train_time:79989ms step_avg:57.42ms step:1394/1600 train_time:80074ms step_avg:57.44ms step:1395/1600 train_time:80163ms step_avg:57.46ms step:1396/1600 train_time:80249ms step_avg:57.49ms step:1397/1600 train_time:80338ms step_avg:57.51ms step:1398/1600 train_time:80423ms step_avg:57.53ms step:1399/1600 train_time:80511ms step_avg:57.55ms step:1400/1600 train_time:80596ms step_avg:57.57ms step:1401/1600 train_time:80685ms step_avg:57.59ms step:1402/1600 train_time:80770ms step_avg:57.61ms step:1403/1600 train_time:80858ms step_avg:57.63ms step:1404/1600 train_time:80943ms step_avg:57.65ms step:1405/1600 train_time:81032ms step_avg:57.67ms step:1406/1600 train_time:81117ms step_avg:57.69ms step:1407/1600 train_time:81205ms step_avg:57.71ms step:1408/1600 train_time:81290ms step_avg:57.73ms step:1409/1600 train_time:81378ms step_avg:57.76ms step:1410/1600 train_time:81463ms step_avg:57.78ms step:1411/1600 train_time:81552ms step_avg:57.80ms step:1412/1600 train_time:81636ms step_avg:57.82ms step:1413/1600 train_time:81724ms step_avg:57.84ms step:1414/1600 train_time:81809ms step_avg:57.86ms step:1415/1600 train_time:81898ms step_avg:57.88ms step:1416/1600 train_time:81984ms step_avg:57.90ms step:1417/1600 train_time:82072ms step_avg:57.92ms step:1418/1600 train_time:82157ms step_avg:57.94ms step:1419/1600 train_time:82245ms step_avg:57.96ms step:1420/1600 train_time:82331ms step_avg:57.98ms step:1421/1600 train_time:82418ms step_avg:58.00ms step:1422/1600 train_time:82502ms step_avg:58.02ms step:1423/1600 train_time:82590ms step_avg:58.04ms step:1424/1600 train_time:82675ms step_avg:58.06ms step:1425/1600 train_time:82764ms step_avg:58.08ms step:1426/1600 train_time:82850ms step_avg:58.10ms step:1427/1600 train_time:82939ms step_avg:58.12ms step:1428/1600 train_time:83024ms step_avg:58.14ms step:1429/1600 train_time:83113ms step_avg:58.16ms step:1430/1600 train_time:83199ms step_avg:58.18ms step:1431/1600 train_time:83287ms step_avg:58.20ms step:1432/1600 train_time:83371ms step_avg:58.22ms step:1433/1600 train_time:83459ms step_avg:58.24ms step:1434/1600 train_time:83544ms step_avg:58.26ms step:1435/1600 train_time:83633ms step_avg:58.28ms step:1436/1600 train_time:83718ms step_avg:58.30ms step:1437/1600 train_time:83806ms step_avg:58.32ms step:1438/1600 train_time:83893ms step_avg:58.34ms step:1439/1600 train_time:83981ms step_avg:58.36ms step:1440/1600 train_time:84066ms step_avg:58.38ms step:1441/1600 train_time:84154ms step_avg:58.40ms step:1442/1600 train_time:84239ms step_avg:58.42ms step:1443/1600 train_time:84327ms step_avg:58.44ms step:1444/1600 train_time:84412ms step_avg:58.46ms step:1445/1600 train_time:84500ms step_avg:58.48ms step:1446/1600 train_time:84585ms step_avg:58.50ms step:1447/1600 train_time:84675ms step_avg:58.52ms step:1448/1600 train_time:84760ms step_avg:58.54ms step:1449/1600 train_time:84849ms step_avg:58.56ms step:1450/1600 train_time:84933ms step_avg:58.57ms step:1451/1600 train_time:85022ms step_avg:58.60ms step:1452/1600 train_time:85107ms step_avg:58.61ms step:1453/1600 train_time:85196ms step_avg:58.63ms step:1454/1600 train_time:85280ms step_avg:58.65ms step:1455/1600 train_time:85369ms step_avg:58.67ms step:1456/1600 train_time:85453ms step_avg:58.69ms step:1457/1600 train_time:85542ms step_avg:58.71ms step:1458/1600 train_time:85627ms step_avg:58.73ms step:1459/1600 train_time:85716ms step_avg:58.75ms step:1460/1600 train_time:85801ms step_avg:58.77ms step:1461/1600 train_time:85890ms step_avg:58.79ms step:1462/1600 train_time:85975ms step_avg:58.81ms step:1463/1600 train_time:86064ms step_avg:58.83ms step:1464/1600 train_time:86149ms step_avg:58.84ms step:1465/1600 train_time:86238ms step_avg:58.87ms step:1466/1600 train_time:86323ms step_avg:58.88ms step:1467/1600 train_time:86412ms step_avg:58.90ms step:1468/1600 train_time:86497ms step_avg:58.92ms step:1469/1600 train_time:86585ms step_avg:58.94ms step:1470/1600 train_time:86670ms step_avg:58.96ms step:1471/1600 train_time:86758ms step_avg:58.98ms step:1472/1600 train_time:86844ms step_avg:59.00ms step:1473/1600 train_time:86932ms step_avg:59.02ms step:1474/1600 train_time:87017ms step_avg:59.03ms step:1475/1600 train_time:87105ms step_avg:59.05ms step:1476/1600 train_time:87191ms step_avg:59.07ms step:1477/1600 train_time:87279ms step_avg:59.09ms step:1478/1600 train_time:87364ms step_avg:59.11ms step:1479/1600 train_time:87452ms step_avg:59.13ms step:1480/1600 train_time:87537ms step_avg:59.15ms step:1481/1600 train_time:87625ms step_avg:59.17ms step:1482/1600 train_time:87711ms step_avg:59.18ms step:1483/1600 train_time:87799ms step_avg:59.20ms step:1484/1600 train_time:87884ms step_avg:59.22ms step:1485/1600 train_time:87972ms step_avg:59.24ms step:1486/1600 train_time:88057ms step_avg:59.26ms step:1487/1600 train_time:88145ms step_avg:59.28ms step:1488/1600 train_time:88230ms step_avg:59.29ms step:1489/1600 train_time:88318ms step_avg:59.31ms step:1490/1600 train_time:88403ms step_avg:59.33ms step:1491/1600 train_time:88492ms step_avg:59.35ms step:1492/1600 train_time:88577ms step_avg:59.37ms step:1493/1600 train_time:88665ms step_avg:59.39ms step:1494/1600 train_time:88750ms step_avg:59.40ms step:1495/1600 train_time:88838ms step_avg:59.42ms step:1496/1600 train_time:88924ms step_avg:59.44ms step:1497/1600 train_time:89013ms step_avg:59.46ms step:1498/1600 train_time:89098ms step_avg:59.48ms step:1499/1600 train_time:89186ms step_avg:59.50ms step:1500/1600 train_time:89272ms step_avg:59.51ms step:1500/1600 val_loss:3.3075 train_time:89345ms step_avg:59.56ms step:1501/1600 train_time:89364ms step_avg:59.54ms step:1502/1600 train_time:89449ms step_avg:59.55ms step:1503/1600 train_time:89540ms step_avg:59.57ms step:1504/1600 train_time:89625ms step_avg:59.59ms step:1505/1600 train_time:89713ms step_avg:59.61ms step:1506/1600 train_time:89798ms step_avg:59.63ms step:1507/1600 train_time:89885ms step_avg:59.64ms step:1508/1600 train_time:89969ms step_avg:59.66ms step:1509/1600 train_time:90056ms step_avg:59.68ms step:1510/1600 train_time:90141ms step_avg:59.70ms step:1511/1600 train_time:90229ms step_avg:59.71ms step:1512/1600 train_time:90315ms step_avg:59.73ms step:1513/1600 train_time:90405ms step_avg:59.75ms step:1514/1600 train_time:90492ms step_avg:59.77ms step:1515/1600 train_time:90582ms step_avg:59.79ms step:1516/1600 train_time:90668ms step_avg:59.81ms step:1517/1600 train_time:90756ms step_avg:59.83ms step:1518/1600 train_time:90841ms step_avg:59.84ms step:1519/1600 train_time:90928ms step_avg:59.86ms step:1520/1600 train_time:91012ms step_avg:59.88ms step:1521/1600 train_time:91100ms step_avg:59.89ms step:1522/1600 train_time:91184ms step_avg:59.91ms step:1523/1600 train_time:91272ms step_avg:59.93ms step:1524/1600 train_time:91358ms step_avg:59.95ms step:1525/1600 train_time:91447ms step_avg:59.97ms step:1526/1600 train_time:91535ms step_avg:59.98ms step:1527/1600 train_time:91622ms step_avg:60.00ms step:1528/1600 train_time:91708ms step_avg:60.02ms step:1529/1600 train_time:91797ms step_avg:60.04ms step:1530/1600 train_time:91881ms step_avg:60.05ms step:1531/1600 train_time:91968ms step_avg:60.07ms step:1532/1600 train_time:92053ms step_avg:60.09ms step:1533/1600 train_time:92140ms step_avg:60.10ms step:1534/1600 train_time:92225ms step_avg:60.12ms step:1535/1600 train_time:92314ms step_avg:60.14ms step:1536/1600 train_time:92401ms step_avg:60.16ms step:1537/1600 train_time:92490ms step_avg:60.18ms step:1538/1600 train_time:92576ms step_avg:60.19ms step:1539/1600 train_time:92664ms step_avg:60.21ms step:1540/1600 train_time:92749ms step_avg:60.23ms step:1541/1600 train_time:92838ms step_avg:60.25ms step:1542/1600 train_time:92922ms step_avg:60.26ms step:1543/1600 train_time:93010ms step_avg:60.28ms step:1544/1600 train_time:93094ms step_avg:60.29ms step:1545/1600 train_time:93182ms step_avg:60.31ms step:1546/1600 train_time:93267ms step_avg:60.33ms step:1547/1600 train_time:93356ms step_avg:60.35ms step:1548/1600 train_time:93441ms step_avg:60.36ms step:1549/1600 train_time:93532ms step_avg:60.38ms step:1550/1600 train_time:93617ms step_avg:60.40ms step:1551/1600 train_time:93705ms step_avg:60.42ms step:1552/1600 train_time:93790ms step_avg:60.43ms step:1553/1600 train_time:93879ms step_avg:60.45ms step:1554/1600 train_time:93966ms step_avg:60.47ms step:1555/1600 train_time:94052ms step_avg:60.48ms step:1556/1600 train_time:94136ms step_avg:60.50ms step:1557/1600 train_time:94224ms step_avg:60.52ms step:1558/1600 train_time:94309ms step_avg:60.53ms step:1559/1600 train_time:94398ms step_avg:60.55ms step:1560/1600 train_time:94484ms step_avg:60.57ms step:1561/1600 train_time:94582ms step_avg:60.59ms step:1562/1600 train_time:94666ms step_avg:60.61ms step:1563/1600 train_time:94754ms step_avg:60.62ms step:1564/1600 train_time:94840ms step_avg:60.64ms step:1565/1600 train_time:94928ms step_avg:60.66ms step:1566/1600 train_time:95014ms step_avg:60.67ms step:1567/1600 train_time:95102ms step_avg:60.69ms step:1568/1600 train_time:95188ms step_avg:60.71ms step:1569/1600 train_time:95276ms step_avg:60.72ms step:1570/1600 train_time:95361ms step_avg:60.74ms step:1571/1600 train_time:95451ms step_avg:60.76ms step:1572/1600 train_time:95538ms step_avg:60.77ms step:1573/1600 train_time:95626ms step_avg:60.79ms step:1574/1600 train_time:95712ms step_avg:60.81ms step:1575/1600 train_time:95802ms step_avg:60.83ms step:1576/1600 train_time:95887ms step_avg:60.84ms step:1577/1600 train_time:95978ms step_avg:60.86ms step:1578/1600 train_time:96062ms step_avg:60.88ms step:1579/1600 train_time:96150ms step_avg:60.89ms step:1580/1600 train_time:96235ms step_avg:60.91ms step:1581/1600 train_time:96323ms step_avg:60.93ms step:1582/1600 train_time:96408ms step_avg:60.94ms step:1583/1600 train_time:96498ms step_avg:60.96ms step:1584/1600 train_time:96583ms step_avg:60.97ms step:1585/1600 train_time:96673ms step_avg:60.99ms step:1586/1600 train_time:96759ms step_avg:61.01ms step:1587/1600 train_time:96846ms step_avg:61.02ms step:1588/1600 train_time:96933ms step_avg:61.04ms step:1589/1600 train_time:97021ms step_avg:61.06ms step:1590/1600 train_time:97106ms step_avg:61.07ms step:1591/1600 train_time:97195ms step_avg:61.09ms step:1592/1600 train_time:97280ms step_avg:61.11ms step:1593/1600 train_time:97370ms step_avg:61.12ms step:1594/1600 train_time:97455ms step_avg:61.14ms step:1595/1600 train_time:97544ms step_avg:61.16ms step:1596/1600 train_time:97630ms step_avg:61.17ms step:1597/1600 train_time:97719ms step_avg:61.19ms step:1598/1600 train_time:97804ms step_avg:61.20ms step:1599/1600 train_time:97892ms step_avg:61.22ms step:1600/1600 train_time:97979ms step_avg:61.24ms step:1600/1600 val_loss:3.2785 train_time:98053ms step_avg:61.28ms peak memory allocated: 30264 MiB reserved: 46240 MiB