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:16:38 2026 +-----------------------------------------------------------------------------------------+ | NVIDIA-SMI 570.148.08 Driver Version: 570.148.08 CUDA Version: 12.8 | |-----------------------------------------+------------------------+----------------------+ | GPU Name Persistence-M | Bus-Id Disp.A | Volatile Uncorr. ECC | | Fan Temp Perf Pwr:Usage/Cap | Memory-Usage | GPU-Util Compute M. | | | | MIG M. | |=========================================+========================+======================| | 0 NVIDIA H100 80GB HBM3 On | 00000000:61:00.0 Off | 0 | | N/A 35C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:62:00.0 Off | 0 | | N/A 39C P0 122W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 41C P0 122W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 35C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:6A:00.0 Off | 0 | | N/A 35C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 42C P0 133W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 40C P0 125W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 36C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 295628 C /usr/bin/python3 1510MiB | | 1 N/A N/A 295629 C /usr/bin/python3 1510MiB | | 2 N/A N/A 295630 C /usr/bin/python3 1510MiB | | 3 N/A N/A 295631 C /usr/bin/python3 1510MiB | | 4 N/A N/A 295632 C /usr/bin/python3 1510MiB | | 5 N/A N/A 295633 C /usr/bin/python3 1510MiB | | 6 N/A N/A 295634 C /usr/bin/python3 1510MiB | | 7 N/A N/A 295635 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.8305 train_time:0ms step_avg:0.03ms step:1/1600 train_time:79ms step_avg:79.10ms step:2/1600 train_time:103ms step_avg:51.45ms step:3/1600 train_time:122ms step_avg:40.65ms step:4/1600 train_time:146ms step_avg:36.41ms step:5/1600 train_time:176ms step_avg:35.29ms step:6/1600 train_time:358ms step_avg:59.72ms step:7/1600 train_time:378ms step_avg:54.06ms step:8/1600 train_time:399ms step_avg:49.86ms step:9/1600 train_time:424ms step_avg:47.15ms step:10/1600 train_time:461ms step_avg:46.14ms step:11/1600 train_time:492ms step_avg:44.74ms step:12/1600 train_time:529ms step_avg:44.07ms step:13/1600 train_time:560ms step_avg:43.04ms step:14/1600 train_time:596ms step_avg:42.60ms step:15/1600 train_time:628ms step_avg:41.85ms step:16/1600 train_time:665ms step_avg:41.56ms step:17/1600 train_time:696ms step_avg:40.93ms step:18/1600 train_time:732ms step_avg:40.69ms step:19/1600 train_time:764ms step_avg:40.20ms step:20/1600 train_time:801ms step_avg:40.05ms step:21/1600 train_time:832ms step_avg:39.62ms step:22/1600 train_time:869ms step_avg:39.52ms step:23/1600 train_time:900ms step_avg:39.15ms step:24/1600 train_time:938ms step_avg:39.06ms step:25/1600 train_time:969ms step_avg:38.74ms step:26/1600 train_time:1006ms step_avg:38.68ms step:27/1600 train_time:1037ms step_avg:38.39ms step:28/1600 train_time:1073ms step_avg:38.33ms step:29/1600 train_time:1104ms step_avg:38.08ms step:30/1600 train_time:1141ms step_avg:38.04ms step:31/1600 train_time:1172ms step_avg:37.81ms step:32/1600 train_time:1209ms step_avg:37.78ms step:33/1600 train_time:1240ms step_avg:37.58ms step:34/1600 train_time:1278ms step_avg:37.59ms step:35/1600 train_time:1309ms step_avg:37.40ms step:36/1600 train_time:1346ms step_avg:37.40ms step:37/1600 train_time:1377ms step_avg:37.22ms step:38/1600 train_time:1414ms step_avg:37.21ms step:39/1600 train_time:1445ms step_avg:37.05ms step:40/1600 train_time:1482ms step_avg:37.06ms step:41/1600 train_time:1513ms step_avg:36.91ms step:42/1600 train_time:1550ms step_avg:36.91ms step:43/1600 train_time:1581ms step_avg:36.77ms step:44/1600 train_time:1618ms step_avg:36.78ms step:45/1600 train_time:1649ms step_avg:36.65ms step:46/1600 train_time:1687ms step_avg:36.66ms step:47/1600 train_time:1717ms step_avg:36.54ms step:48/1600 train_time:1754ms step_avg:36.55ms step:49/1600 train_time:1785ms step_avg:36.43ms step:50/1600 train_time:1822ms step_avg:36.45ms step:51/1600 train_time:1853ms step_avg:36.34ms step:52/1600 train_time:1890ms step_avg:36.35ms step:53/1600 train_time:1921ms step_avg:36.24ms step:54/1600 train_time:1958ms step_avg:36.25ms step:55/1600 train_time:1988ms step_avg:36.15ms step:56/1600 train_time:2025ms step_avg:36.17ms step:57/1600 train_time:2056ms step_avg:36.07ms step:58/1600 train_time:2093ms step_avg:36.09ms step:59/1600 train_time:2124ms step_avg:36.00ms step:60/1600 train_time:2161ms step_avg:36.02ms step:61/1600 train_time:2192ms step_avg:35.94ms step:62/1600 train_time:2229ms step_avg:35.96ms step:63/1600 train_time:2261ms step_avg:35.88ms step:64/1600 train_time:2298ms step_avg:35.90ms step:65/1600 train_time:2328ms step_avg:35.82ms step:66/1600 train_time:2366ms step_avg:35.85ms step:67/1600 train_time:2396ms step_avg:35.77ms step:68/1600 train_time:2433ms step_avg:35.78ms step:69/1600 train_time:2464ms step_avg:35.71ms step:70/1600 train_time:2501ms step_avg:35.73ms step:71/1600 train_time:2532ms step_avg:35.67ms step:72/1600 train_time:2570ms step_avg:35.69ms step:73/1600 train_time:2601ms step_avg:35.63ms step:74/1600 train_time:2638ms step_avg:35.64ms step:75/1600 train_time:2669ms step_avg:35.59ms step:76/1600 train_time:2706ms step_avg:35.61ms step:77/1600 train_time:2737ms step_avg:35.55ms step:78/1600 train_time:2774ms step_avg:35.57ms step:79/1600 train_time:2805ms step_avg:35.51ms step:80/1600 train_time:2843ms step_avg:35.53ms step:81/1600 train_time:2874ms step_avg:35.48ms step:82/1600 train_time:2911ms step_avg:35.50ms step:83/1600 train_time:2942ms step_avg:35.44ms step:84/1600 train_time:2979ms step_avg:35.46ms step:85/1600 train_time:3010ms step_avg:35.41ms step:86/1600 train_time:3048ms step_avg:35.44ms step:87/1600 train_time:3079ms step_avg:35.39ms step:88/1600 train_time:3115ms step_avg:35.40ms step:89/1600 train_time:3146ms step_avg:35.35ms step:90/1600 train_time:3184ms step_avg:35.38ms step:91/1600 train_time:3215ms step_avg:35.33ms step:92/1600 train_time:3252ms step_avg:35.35ms step:93/1600 train_time:3283ms step_avg:35.30ms step:94/1600 train_time:3319ms step_avg:35.31ms step:95/1600 train_time:3351ms step_avg:35.27ms step:96/1600 train_time:3389ms step_avg:35.30ms step:97/1600 train_time:3420ms step_avg:35.25ms step:98/1600 train_time:3456ms step_avg:35.27ms step:99/1600 train_time:3487ms step_avg:35.23ms step:100/1600 train_time:3525ms step_avg:35.25ms step:101/1600 train_time:3556ms step_avg:35.20ms step:102/1600 train_time:3592ms step_avg:35.22ms step:103/1600 train_time:3624ms step_avg:35.19ms step:104/1600 train_time:3661ms step_avg:35.20ms step:105/1600 train_time:3692ms step_avg:35.16ms step:106/1600 train_time:3729ms step_avg:35.18ms step:107/1600 train_time:3760ms step_avg:35.14ms step:108/1600 train_time:3797ms step_avg:35.16ms step:109/1600 train_time:3828ms step_avg:35.12ms step:110/1600 train_time:3866ms step_avg:35.14ms step:111/1600 train_time:3896ms step_avg:35.10ms step:112/1600 train_time:3934ms step_avg:35.12ms step:113/1600 train_time:3965ms step_avg:35.09ms step:114/1600 train_time:4002ms step_avg:35.11ms step:115/1600 train_time:4033ms step_avg:35.07ms step:116/1600 train_time:4070ms step_avg:35.09ms step:117/1600 train_time:4101ms step_avg:35.05ms step:118/1600 train_time:4138ms step_avg:35.07ms step:119/1600 train_time:4169ms step_avg:35.03ms step:120/1600 train_time:4206ms step_avg:35.05ms step:121/1600 train_time:4237ms step_avg:35.02ms step:122/1600 train_time:4274ms step_avg:35.03ms step:123/1600 train_time:4305ms step_avg:35.00ms step:124/1600 train_time:4342ms step_avg:35.01ms step:125/1600 train_time:4373ms step_avg:34.98ms step:126/1600 train_time:4410ms step_avg:35.00ms step:127/1600 train_time:4441ms step_avg:34.97ms step:128/1600 train_time:4478ms step_avg:34.98ms step:129/1600 train_time:4509ms step_avg:34.96ms step:130/1600 train_time:4547ms step_avg:34.98ms step:131/1600 train_time:4578ms step_avg:34.95ms step:132/1600 train_time:4615ms step_avg:34.96ms step:133/1600 train_time:4646ms step_avg:34.93ms step:134/1600 train_time:4683ms step_avg:34.95ms step:135/1600 train_time:4713ms step_avg:34.91ms step:136/1600 train_time:4750ms step_avg:34.93ms step:137/1600 train_time:4781ms step_avg:34.90ms step:138/1600 train_time:4818ms step_avg:34.91ms step:139/1600 train_time:4849ms step_avg:34.88ms step:140/1600 train_time:4886ms step_avg:34.90ms step:141/1600 train_time:4917ms step_avg:34.87ms step:142/1600 train_time:4954ms step_avg:34.89ms step:143/1600 train_time:4985ms step_avg:34.86ms step:144/1600 train_time:5022ms step_avg:34.87ms step:145/1600 train_time:5053ms step_avg:34.85ms step:146/1600 train_time:5090ms step_avg:34.86ms step:147/1600 train_time:5121ms step_avg:34.84ms step:148/1600 train_time:5158ms step_avg:34.85ms step:149/1600 train_time:5189ms step_avg:34.83ms step:150/1600 train_time:5226ms step_avg:34.84ms step:151/1600 train_time:5257ms step_avg:34.82ms step:152/1600 train_time:5294ms step_avg:34.83ms step:153/1600 train_time:5325ms step_avg:34.80ms step:154/1600 train_time:5362ms step_avg:34.82ms step:155/1600 train_time:5393ms step_avg:34.79ms step:156/1600 train_time:5430ms step_avg:34.81ms step:157/1600 train_time:5461ms step_avg:34.78ms step:158/1600 train_time:5498ms step_avg:34.80ms step:159/1600 train_time:5529ms step_avg:34.77ms step:160/1600 train_time:5566ms step_avg:34.79ms step:161/1600 train_time:5597ms step_avg:34.76ms step:162/1600 train_time:5634ms step_avg:34.78ms step:163/1600 train_time:5665ms step_avg:34.75ms step:164/1600 train_time:5702ms step_avg:34.77ms step:165/1600 train_time:5733ms step_avg:34.75ms step:166/1600 train_time:5770ms step_avg:34.76ms step:167/1600 train_time:5801ms step_avg:34.74ms step:168/1600 train_time:5838ms step_avg:34.75ms step:169/1600 train_time:5869ms step_avg:34.73ms step:170/1600 train_time:5906ms step_avg:34.74ms step:171/1600 train_time:5937ms step_avg:34.72ms step:172/1600 train_time:5974ms step_avg:34.73ms step:173/1600 train_time:6004ms step_avg:34.71ms step:174/1600 train_time:6041ms step_avg:34.72ms step:175/1600 train_time:6072ms step_avg:34.70ms step:176/1600 train_time:6110ms step_avg:34.71ms step:177/1600 train_time:6140ms step_avg:34.69ms step:178/1600 train_time:6177ms step_avg:34.70ms step:179/1600 train_time:6208ms step_avg:34.68ms step:180/1600 train_time:6245ms step_avg:34.70ms step:181/1600 train_time:6276ms step_avg:34.68ms step:182/1600 train_time:6313ms step_avg:34.69ms step:183/1600 train_time:6344ms step_avg:34.67ms step:184/1600 train_time:6381ms step_avg:34.68ms step:185/1600 train_time:6412ms step_avg:34.66ms step:186/1600 train_time:6449ms step_avg:34.67ms step:187/1600 train_time:6479ms step_avg:34.65ms step:188/1600 train_time:6516ms step_avg:34.66ms step:189/1600 train_time:6547ms step_avg:34.64ms step:190/1600 train_time:6584ms step_avg:34.66ms step:191/1600 train_time:6615ms step_avg:34.64ms step:192/1600 train_time:6653ms step_avg:34.65ms step:193/1600 train_time:6684ms step_avg:34.63ms step:194/1600 train_time:6721ms step_avg:34.65ms step:195/1600 train_time:6752ms step_avg:34.63ms step:196/1600 train_time:6789ms step_avg:34.64ms step:197/1600 train_time:6820ms step_avg:34.62ms step:198/1600 train_time:6856ms step_avg:34.63ms step:199/1600 train_time:6888ms step_avg:34.61ms step:200/1600 train_time:6925ms step_avg:34.62ms step:201/1600 train_time:6956ms step_avg:34.61ms step:202/1600 train_time:6993ms step_avg:34.62ms step:203/1600 train_time:7024ms step_avg:34.60ms step:204/1600 train_time:7061ms step_avg:34.61ms step:205/1600 train_time:7092ms step_avg:34.60ms step:206/1600 train_time:7129ms step_avg:34.61ms step:207/1600 train_time:7160ms step_avg:34.59ms step:208/1600 train_time:7196ms step_avg:34.60ms step:209/1600 train_time:7227ms step_avg:34.58ms step:210/1600 train_time:7265ms step_avg:34.59ms step:211/1600 train_time:7296ms step_avg:34.58ms step:212/1600 train_time:7333ms step_avg:34.59ms step:213/1600 train_time:7364ms step_avg:34.57ms step:214/1600 train_time:7401ms step_avg:34.58ms step:215/1600 train_time:7432ms step_avg:34.57ms step:216/1600 train_time:7470ms step_avg:34.58ms step:217/1600 train_time:7500ms step_avg:34.56ms step:218/1600 train_time:7537ms step_avg:34.57ms step:219/1600 train_time:7568ms step_avg:34.56ms step:220/1600 train_time:7605ms step_avg:34.57ms step:221/1600 train_time:7636ms step_avg:34.55ms step:222/1600 train_time:7673ms step_avg:34.56ms step:223/1600 train_time:7703ms step_avg:34.54ms step:224/1600 train_time:7740ms step_avg:34.55ms step:225/1600 train_time:7771ms step_avg:34.54ms step:226/1600 train_time:7808ms step_avg:34.55ms step:227/1600 train_time:7839ms step_avg:34.53ms step:228/1600 train_time:7876ms step_avg:34.54ms step:229/1600 train_time:7906ms step_avg:34.53ms step:230/1600 train_time:7944ms step_avg:34.54ms step:231/1600 train_time:7974ms step_avg:34.52ms step:232/1600 train_time:8011ms step_avg:34.53ms step:233/1600 train_time:8042ms step_avg:34.52ms step:234/1600 train_time:8079ms step_avg:34.53ms step:235/1600 train_time:8110ms step_avg:34.51ms step:236/1600 train_time:8148ms step_avg:34.52ms step:237/1600 train_time:8179ms step_avg:34.51ms step:238/1600 train_time:8215ms step_avg:34.52ms step:239/1600 train_time:8247ms step_avg:34.50ms step:240/1600 train_time:8284ms step_avg:34.52ms step:241/1600 train_time:8315ms step_avg:34.50ms step:242/1600 train_time:8352ms step_avg:34.51ms step:243/1600 train_time:8383ms step_avg:34.50ms step:244/1600 train_time:8420ms step_avg:34.51ms step:245/1600 train_time:8451ms step_avg:34.49ms step:246/1600 train_time:8488ms step_avg:34.51ms step:247/1600 train_time:8519ms step_avg:34.49ms step:248/1600 train_time:8556ms step_avg:34.50ms step:249/1600 train_time:8587ms step_avg:34.49ms step:250/1600 train_time:8625ms step_avg:34.50ms step:250/1600 val_loss:4.5740 train_time:8673ms step_avg:34.69ms step:251/1600 train_time:8691ms step_avg:34.62ms step:252/1600 train_time:8709ms step_avg:34.56ms step:253/1600 train_time:8727ms step_avg:34.49ms step:254/1600 train_time:8765ms step_avg:34.51ms step:255/1600 train_time:8798ms step_avg:34.50ms step:256/1600 train_time:8836ms step_avg:34.51ms step:257/1600 train_time:8867ms step_avg:34.50ms step:258/1600 train_time:8904ms step_avg:34.51ms step:259/1600 train_time:8935ms step_avg:34.50ms step:260/1600 train_time:8972ms step_avg:34.51ms step:261/1600 train_time:9003ms step_avg:34.50ms step:262/1600 train_time:9040ms step_avg:34.50ms step:263/1600 train_time:9071ms step_avg:34.49ms step:264/1600 train_time:9107ms step_avg:34.50ms step:265/1600 train_time:9138ms step_avg:34.48ms step:266/1600 train_time:9175ms step_avg:34.49ms step:267/1600 train_time:9206ms step_avg:34.48ms step:268/1600 train_time:9243ms step_avg:34.49ms step:269/1600 train_time:9273ms step_avg:34.47ms step:270/1600 train_time:9310ms step_avg:34.48ms step:271/1600 train_time:9341ms step_avg:34.47ms step:272/1600 train_time:9378ms step_avg:34.48ms step:273/1600 train_time:9409ms step_avg:34.46ms step:274/1600 train_time:9445ms step_avg:34.47ms step:275/1600 train_time:9476ms step_avg:34.46ms step:276/1600 train_time:9513ms step_avg:34.47ms step:277/1600 train_time:9544ms step_avg:34.46ms step:278/1600 train_time:9581ms step_avg:34.46ms step:279/1600 train_time:9612ms step_avg:34.45ms step:280/1600 train_time:9649ms step_avg:34.46ms step:281/1600 train_time:9680ms step_avg:34.45ms step:282/1600 train_time:9717ms step_avg:34.46ms step:283/1600 train_time:9748ms step_avg:34.45ms step:284/1600 train_time:9786ms step_avg:34.46ms step:285/1600 train_time:9817ms step_avg:34.44ms step:286/1600 train_time:9854ms step_avg:34.45ms step:287/1600 train_time:9884ms step_avg:34.44ms step:288/1600 train_time:9922ms step_avg:34.45ms step:289/1600 train_time:9953ms step_avg:34.44ms step:290/1600 train_time:9990ms step_avg:34.45ms step:291/1600 train_time:10021ms step_avg:34.43ms step:292/1600 train_time:10057ms step_avg:34.44ms step:293/1600 train_time:10088ms step_avg:34.43ms step:294/1600 train_time:10125ms step_avg:34.44ms step:295/1600 train_time:10156ms step_avg:34.43ms step:296/1600 train_time:10193ms step_avg:34.44ms step:297/1600 train_time:10224ms step_avg:34.42ms step:298/1600 train_time:10261ms step_avg:34.43ms step:299/1600 train_time:10291ms step_avg:34.42ms step:300/1600 train_time:10328ms step_avg:34.43ms step:301/1600 train_time:10359ms step_avg:34.41ms step:302/1600 train_time:10396ms step_avg:34.42ms step:303/1600 train_time:10427ms step_avg:34.41ms step:304/1600 train_time:10464ms step_avg:34.42ms step:305/1600 train_time:10495ms step_avg:34.41ms step:306/1600 train_time:10532ms step_avg:34.42ms step:307/1600 train_time:10563ms step_avg:34.41ms step:308/1600 train_time:10599ms step_avg:34.41ms step:309/1600 train_time:10630ms step_avg:34.40ms step:310/1600 train_time:10667ms step_avg:34.41ms step:311/1600 train_time:10698ms step_avg:34.40ms step:312/1600 train_time:10735ms step_avg:34.41ms step:313/1600 train_time:10765ms step_avg:34.39ms step:314/1600 train_time:10802ms step_avg:34.40ms step:315/1600 train_time:10833ms step_avg:34.39ms step:316/1600 train_time:10870ms step_avg:34.40ms step:317/1600 train_time:10901ms step_avg:34.39ms step:318/1600 train_time:10938ms step_avg:34.40ms step:319/1600 train_time:10970ms step_avg:34.39ms step:320/1600 train_time:11006ms step_avg:34.39ms step:321/1600 train_time:11037ms step_avg:34.38ms step:322/1600 train_time:11074ms step_avg:34.39ms step:323/1600 train_time:11105ms step_avg:34.38ms step:324/1600 train_time:11142ms step_avg:34.39ms step:325/1600 train_time:11173ms step_avg:34.38ms step:326/1600 train_time:11210ms step_avg:34.39ms step:327/1600 train_time:11241ms step_avg:34.38ms step:328/1600 train_time:11278ms step_avg:34.38ms step:329/1600 train_time:11309ms step_avg:34.37ms step:330/1600 train_time:11346ms step_avg:34.38ms step:331/1600 train_time:11377ms step_avg:34.37ms step:332/1600 train_time:11414ms step_avg:34.38ms step:333/1600 train_time:11445ms step_avg:34.37ms step:334/1600 train_time:11482ms step_avg:34.38ms step:335/1600 train_time:11512ms step_avg:34.37ms step:336/1600 train_time:11550ms step_avg:34.38ms step:337/1600 train_time:11580ms step_avg:34.36ms step:338/1600 train_time:11618ms step_avg:34.37ms step:339/1600 train_time:11648ms step_avg:34.36ms step:340/1600 train_time:11685ms step_avg:34.37ms step:341/1600 train_time:11716ms step_avg:34.36ms step:342/1600 train_time:11753ms step_avg:34.37ms step:343/1600 train_time:11784ms step_avg:34.36ms step:344/1600 train_time:11821ms step_avg:34.36ms step:345/1600 train_time:11852ms step_avg:34.35ms step:346/1600 train_time:11890ms step_avg:34.36ms step:347/1600 train_time:11921ms step_avg:34.35ms step:348/1600 train_time:11958ms step_avg:34.36ms step:349/1600 train_time:11989ms step_avg:34.35ms step:350/1600 train_time:12026ms step_avg:34.36ms step:351/1600 train_time:12057ms step_avg:34.35ms step:352/1600 train_time:12094ms step_avg:34.36ms step:353/1600 train_time:12125ms step_avg:34.35ms step:354/1600 train_time:12161ms step_avg:34.35ms step:355/1600 train_time:12193ms step_avg:34.35ms step:356/1600 train_time:12230ms step_avg:34.35ms step:357/1600 train_time:12261ms step_avg:34.34ms step:358/1600 train_time:12298ms step_avg:34.35ms step:359/1600 train_time:12328ms step_avg:34.34ms step:360/1600 train_time:12365ms step_avg:34.35ms step:361/1600 train_time:12397ms step_avg:34.34ms step:362/1600 train_time:12434ms step_avg:34.35ms step:363/1600 train_time:12465ms step_avg:34.34ms step:364/1600 train_time:12502ms step_avg:34.35ms step:365/1600 train_time:12532ms step_avg:34.34ms step:366/1600 train_time:12569ms step_avg:34.34ms step:367/1600 train_time:12600ms step_avg:34.33ms step:368/1600 train_time:12637ms step_avg:34.34ms step:369/1600 train_time:12668ms step_avg:34.33ms step:370/1600 train_time:12705ms step_avg:34.34ms step:371/1600 train_time:12735ms step_avg:34.33ms step:372/1600 train_time:12772ms step_avg:34.33ms step:373/1600 train_time:12803ms step_avg:34.32ms step:374/1600 train_time:12839ms step_avg:34.33ms step:375/1600 train_time:12870ms step_avg:34.32ms step:376/1600 train_time:12907ms step_avg:34.33ms step:377/1600 train_time:12938ms step_avg:34.32ms step:378/1600 train_time:12975ms step_avg:34.33ms step:379/1600 train_time:13005ms step_avg:34.32ms step:380/1600 train_time:13042ms step_avg:34.32ms step:381/1600 train_time:13073ms step_avg:34.31ms step:382/1600 train_time:13110ms step_avg:34.32ms step:383/1600 train_time:13142ms step_avg:34.31ms step:384/1600 train_time:13179ms step_avg:34.32ms step:385/1600 train_time:13210ms step_avg:34.31ms step:386/1600 train_time:13247ms step_avg:34.32ms step:387/1600 train_time:13277ms step_avg:34.31ms step:388/1600 train_time:13314ms step_avg:34.31ms step:389/1600 train_time:13345ms step_avg:34.31ms step:390/1600 train_time:13382ms step_avg:34.31ms step:391/1600 train_time:13412ms step_avg:34.30ms step:392/1600 train_time:13450ms step_avg:34.31ms step:393/1600 train_time:13480ms step_avg:34.30ms step:394/1600 train_time:13517ms step_avg:34.31ms step:395/1600 train_time:13548ms step_avg:34.30ms step:396/1600 train_time:13585ms step_avg:34.30ms step:397/1600 train_time:13615ms step_avg:34.30ms step:398/1600 train_time:13652ms step_avg:34.30ms step:399/1600 train_time:13683ms step_avg:34.29ms step:400/1600 train_time:13720ms step_avg:34.30ms step:401/1600 train_time:13751ms step_avg:34.29ms step:402/1600 train_time:13788ms step_avg:34.30ms step:403/1600 train_time:13819ms step_avg:34.29ms step:404/1600 train_time:13856ms step_avg:34.30ms step:405/1600 train_time:13886ms step_avg:34.29ms step:406/1600 train_time:13923ms step_avg:34.29ms step:407/1600 train_time:13954ms step_avg:34.29ms step:408/1600 train_time:13991ms step_avg:34.29ms step:409/1600 train_time:14022ms step_avg:34.28ms step:410/1600 train_time:14059ms step_avg:34.29ms step:411/1600 train_time:14090ms step_avg:34.28ms step:412/1600 train_time:14127ms step_avg:34.29ms step:413/1600 train_time:14158ms step_avg:34.28ms step:414/1600 train_time:14195ms step_avg:34.29ms step:415/1600 train_time:14226ms step_avg:34.28ms step:416/1600 train_time:14263ms step_avg:34.29ms step:417/1600 train_time:14294ms step_avg:34.28ms step:418/1600 train_time:14331ms step_avg:34.28ms step:419/1600 train_time:14362ms step_avg:34.28ms step:420/1600 train_time:14399ms step_avg:34.28ms step:421/1600 train_time:14429ms step_avg:34.27ms step:422/1600 train_time:14467ms step_avg:34.28ms step:423/1600 train_time:14498ms step_avg:34.27ms step:424/1600 train_time:14535ms step_avg:34.28ms step:425/1600 train_time:14565ms step_avg:34.27ms step:426/1600 train_time:14602ms step_avg:34.28ms step:427/1600 train_time:14633ms step_avg:34.27ms step:428/1600 train_time:14670ms step_avg:34.28ms step:429/1600 train_time:14701ms step_avg:34.27ms step:430/1600 train_time:14738ms step_avg:34.27ms step:431/1600 train_time:14768ms step_avg:34.27ms step:432/1600 train_time:14805ms step_avg:34.27ms step:433/1600 train_time:14836ms step_avg:34.26ms step:434/1600 train_time:14873ms step_avg:34.27ms step:435/1600 train_time:14903ms step_avg:34.26ms step:436/1600 train_time:14940ms step_avg:34.27ms step:437/1600 train_time:14971ms step_avg:34.26ms step:438/1600 train_time:15008ms step_avg:34.27ms step:439/1600 train_time:15039ms step_avg:34.26ms step:440/1600 train_time:15076ms step_avg:34.26ms step:441/1600 train_time:15107ms step_avg:34.26ms step:442/1600 train_time:15143ms step_avg:34.26ms step:443/1600 train_time:15174ms step_avg:34.25ms step:444/1600 train_time:15212ms step_avg:34.26ms step:445/1600 train_time:15243ms step_avg:34.25ms step:446/1600 train_time:15280ms step_avg:34.26ms step:447/1600 train_time:15310ms step_avg:34.25ms step:448/1600 train_time:15347ms step_avg:34.26ms step:449/1600 train_time:15378ms step_avg:34.25ms step:450/1600 train_time:15416ms step_avg:34.26ms step:451/1600 train_time:15447ms step_avg:34.25ms step:452/1600 train_time:15483ms step_avg:34.25ms step:453/1600 train_time:15514ms step_avg:34.25ms step:454/1600 train_time:15551ms step_avg:34.25ms step:455/1600 train_time:15582ms step_avg:34.25ms step:456/1600 train_time:15620ms step_avg:34.25ms step:457/1600 train_time:15650ms step_avg:34.25ms step:458/1600 train_time:15687ms step_avg:34.25ms step:459/1600 train_time:15718ms step_avg:34.24ms step:460/1600 train_time:15755ms step_avg:34.25ms step:461/1600 train_time:15785ms step_avg:34.24ms step:462/1600 train_time:15822ms step_avg:34.25ms step:463/1600 train_time:15854ms step_avg:34.24ms step:464/1600 train_time:15891ms step_avg:34.25ms step:465/1600 train_time:15922ms step_avg:34.24ms step:466/1600 train_time:15959ms step_avg:34.25ms step:467/1600 train_time:15989ms step_avg:34.24ms step:468/1600 train_time:16026ms step_avg:34.24ms step:469/1600 train_time:16057ms step_avg:34.24ms step:470/1600 train_time:16094ms step_avg:34.24ms step:471/1600 train_time:16125ms step_avg:34.24ms step:472/1600 train_time:16162ms step_avg:34.24ms step:473/1600 train_time:16193ms step_avg:34.23ms step:474/1600 train_time:16229ms step_avg:34.24ms step:475/1600 train_time:16260ms step_avg:34.23ms step:476/1600 train_time:16297ms step_avg:34.24ms step:477/1600 train_time:16328ms step_avg:34.23ms step:478/1600 train_time:16365ms step_avg:34.24ms step:479/1600 train_time:16395ms step_avg:34.23ms step:480/1600 train_time:16432ms step_avg:34.23ms step:481/1600 train_time:16463ms step_avg:34.23ms step:482/1600 train_time:16500ms step_avg:34.23ms step:483/1600 train_time:16531ms step_avg:34.23ms step:484/1600 train_time:16568ms step_avg:34.23ms step:485/1600 train_time:16599ms step_avg:34.23ms step:486/1600 train_time:16637ms step_avg:34.23ms step:487/1600 train_time:16668ms step_avg:34.23ms step:488/1600 train_time:16705ms step_avg:34.23ms step:489/1600 train_time:16736ms step_avg:34.22ms step:490/1600 train_time:16773ms step_avg:34.23ms step:491/1600 train_time:16804ms step_avg:34.22ms step:492/1600 train_time:16840ms step_avg:34.23ms step:493/1600 train_time:16871ms step_avg:34.22ms step:494/1600 train_time:16907ms step_avg:34.23ms step:495/1600 train_time:16938ms step_avg:34.22ms step:496/1600 train_time:16976ms step_avg:34.23ms step:497/1600 train_time:17006ms step_avg:34.22ms step:498/1600 train_time:17043ms step_avg:34.22ms step:499/1600 train_time:17074ms step_avg:34.22ms step:500/1600 train_time:17112ms step_avg:34.22ms step:500/1600 val_loss:4.2580 train_time:17160ms step_avg:34.32ms step:501/1600 train_time:17179ms step_avg:34.29ms step:502/1600 train_time:17200ms step_avg:34.26ms step:503/1600 train_time:17219ms step_avg:34.23ms step:504/1600 train_time:17251ms step_avg:34.23ms step:505/1600 train_time:17283ms step_avg:34.22ms step:506/1600 train_time:17322ms step_avg:34.23ms step:507/1600 train_time:17354ms step_avg:34.23ms step:508/1600 train_time:17390ms step_avg:34.23ms step:509/1600 train_time:17421ms step_avg:34.23ms step:510/1600 train_time:17459ms step_avg:34.23ms step:511/1600 train_time:17490ms step_avg:34.23ms step:512/1600 train_time:17527ms step_avg:34.23ms step:513/1600 train_time:17558ms step_avg:34.23ms step:514/1600 train_time:17595ms step_avg:34.23ms step:515/1600 train_time:17626ms step_avg:34.22ms step:516/1600 train_time:17663ms step_avg:34.23ms step:517/1600 train_time:17693ms step_avg:34.22ms step:518/1600 train_time:17730ms step_avg:34.23ms step:519/1600 train_time:17760ms step_avg:34.22ms step:520/1600 train_time:17797ms step_avg:34.23ms step:521/1600 train_time:17867ms step_avg:34.29ms step:522/1600 train_time:17926ms step_avg:34.34ms step:523/1600 train_time:17986ms step_avg:34.39ms step:524/1600 train_time:18045ms step_avg:34.44ms step:525/1600 train_time:18105ms step_avg:34.49ms step:526/1600 train_time:18163ms step_avg:34.53ms step:527/1600 train_time:18228ms step_avg:34.59ms step:528/1600 train_time:18288ms step_avg:34.64ms step:529/1600 train_time:18350ms step_avg:34.69ms step:530/1600 train_time:18409ms step_avg:34.73ms step:531/1600 train_time:18471ms step_avg:34.79ms step:532/1600 train_time:18530ms step_avg:34.83ms step:533/1600 train_time:18594ms step_avg:34.88ms step:534/1600 train_time:18653ms step_avg:34.93ms step:535/1600 train_time:18714ms step_avg:34.98ms step:536/1600 train_time:18772ms step_avg:35.02ms step:537/1600 train_time:18834ms step_avg:35.07ms step:538/1600 train_time:18893ms step_avg:35.12ms step:539/1600 train_time:18956ms step_avg:35.17ms step:540/1600 train_time:19014ms step_avg:35.21ms step:541/1600 train_time:19076ms step_avg:35.26ms step:542/1600 train_time:19134ms step_avg:35.30ms step:543/1600 train_time:19197ms step_avg:35.35ms step:544/1600 train_time:19257ms step_avg:35.40ms step:545/1600 train_time:19320ms step_avg:35.45ms step:546/1600 train_time:19379ms step_avg:35.49ms step:547/1600 train_time:19442ms step_avg:35.54ms step:548/1600 train_time:19501ms step_avg:35.59ms step:549/1600 train_time:19564ms step_avg:35.64ms step:550/1600 train_time:19623ms step_avg:35.68ms step:551/1600 train_time:19686ms step_avg:35.73ms step:552/1600 train_time:19745ms step_avg:35.77ms step:553/1600 train_time:19809ms step_avg:35.82ms step:554/1600 train_time:19867ms step_avg:35.86ms step:555/1600 train_time:19930ms step_avg:35.91ms step:556/1600 train_time:19989ms step_avg:35.95ms step:557/1600 train_time:20051ms step_avg:36.00ms step:558/1600 train_time:20110ms step_avg:36.04ms step:559/1600 train_time:20172ms step_avg:36.09ms step:560/1600 train_time:20231ms step_avg:36.13ms step:561/1600 train_time:20293ms step_avg:36.17ms step:562/1600 train_time:20353ms step_avg:36.21ms step:563/1600 train_time:20415ms step_avg:36.26ms step:564/1600 train_time:20473ms step_avg:36.30ms step:565/1600 train_time:20536ms step_avg:36.35ms step:566/1600 train_time:20595ms step_avg:36.39ms step:567/1600 train_time:20657ms step_avg:36.43ms step:568/1600 train_time:20717ms step_avg:36.47ms step:569/1600 train_time:20781ms step_avg:36.52ms step:570/1600 train_time:20839ms step_avg:36.56ms step:571/1600 train_time:20901ms step_avg:36.60ms step:572/1600 train_time:20960ms step_avg:36.64ms step:573/1600 train_time:21027ms step_avg:36.70ms step:574/1600 train_time:21085ms step_avg:36.73ms step:575/1600 train_time:21146ms step_avg:36.78ms step:576/1600 train_time:21205ms step_avg:36.81ms step:577/1600 train_time:21268ms step_avg:36.86ms step:578/1600 train_time:21327ms step_avg:36.90ms step:579/1600 train_time:21389ms step_avg:36.94ms step:580/1600 train_time:21448ms step_avg:36.98ms step:581/1600 train_time:21511ms step_avg:37.02ms step:582/1600 train_time:21569ms step_avg:37.06ms step:583/1600 train_time:21632ms step_avg:37.10ms step:584/1600 train_time:21691ms step_avg:37.14ms step:585/1600 train_time:21752ms step_avg:37.18ms step:586/1600 train_time:21812ms step_avg:37.22ms step:587/1600 train_time:21873ms step_avg:37.26ms step:588/1600 train_time:21932ms step_avg:37.30ms step:589/1600 train_time:21993ms step_avg:37.34ms step:590/1600 train_time:22053ms step_avg:37.38ms step:591/1600 train_time:22115ms step_avg:37.42ms step:592/1600 train_time:22174ms step_avg:37.46ms step:593/1600 train_time:22237ms step_avg:37.50ms step:594/1600 train_time:22296ms step_avg:37.54ms step:595/1600 train_time:22359ms step_avg:37.58ms step:596/1600 train_time:22418ms step_avg:37.61ms step:597/1600 train_time:22481ms step_avg:37.66ms step:598/1600 train_time:22540ms step_avg:37.69ms step:599/1600 train_time:22602ms step_avg:37.73ms step:600/1600 train_time:22661ms step_avg:37.77ms step:601/1600 train_time:22724ms step_avg:37.81ms step:602/1600 train_time:22784ms step_avg:37.85ms step:603/1600 train_time:22847ms step_avg:37.89ms step:604/1600 train_time:22906ms step_avg:37.92ms step:605/1600 train_time:22969ms step_avg:37.96ms step:606/1600 train_time:23028ms step_avg:38.00ms step:607/1600 train_time:23090ms step_avg:38.04ms step:608/1600 train_time:23149ms step_avg:38.07ms step:609/1600 train_time:23212ms step_avg:38.11ms step:610/1600 train_time:23271ms step_avg:38.15ms step:611/1600 train_time:23333ms step_avg:38.19ms step:612/1600 train_time:23391ms step_avg:38.22ms step:613/1600 train_time:23453ms step_avg:38.26ms step:614/1600 train_time:23512ms step_avg:38.29ms step:615/1600 train_time:23574ms step_avg:38.33ms step:616/1600 train_time:23633ms step_avg:38.37ms step:617/1600 train_time:23695ms step_avg:38.40ms step:618/1600 train_time:23754ms step_avg:38.44ms step:619/1600 train_time:23817ms step_avg:38.48ms step:620/1600 train_time:23877ms step_avg:38.51ms step:621/1600 train_time:23939ms step_avg:38.55ms step:622/1600 train_time:23998ms step_avg:38.58ms step:623/1600 train_time:24061ms step_avg:38.62ms step:624/1600 train_time:24120ms step_avg:38.65ms step:625/1600 train_time:24183ms step_avg:38.69ms step:626/1600 train_time:24242ms step_avg:38.73ms step:627/1600 train_time:24305ms step_avg:38.76ms step:628/1600 train_time:24365ms step_avg:38.80ms step:629/1600 train_time:24428ms step_avg:38.84ms step:630/1600 train_time:24487ms step_avg:38.87ms step:631/1600 train_time:24549ms step_avg:38.91ms step:632/1600 train_time:24608ms step_avg:38.94ms step:633/1600 train_time:24671ms step_avg:38.97ms step:634/1600 train_time:24729ms step_avg:39.00ms step:635/1600 train_time:24791ms step_avg:39.04ms step:636/1600 train_time:24850ms step_avg:39.07ms step:637/1600 train_time:24912ms step_avg:39.11ms step:638/1600 train_time:24970ms step_avg:39.14ms step:639/1600 train_time:25032ms step_avg:39.17ms step:640/1600 train_time:25091ms step_avg:39.21ms step:641/1600 train_time:25154ms step_avg:39.24ms step:642/1600 train_time:25215ms step_avg:39.28ms step:643/1600 train_time:25278ms step_avg:39.31ms step:644/1600 train_time:25335ms step_avg:39.34ms step:645/1600 train_time:25400ms step_avg:39.38ms step:646/1600 train_time:25458ms step_avg:39.41ms step:647/1600 train_time:25521ms step_avg:39.44ms step:648/1600 train_time:25579ms step_avg:39.47ms step:649/1600 train_time:25643ms step_avg:39.51ms step:650/1600 train_time:25702ms step_avg:39.54ms step:651/1600 train_time:25764ms step_avg:39.58ms step:652/1600 train_time:25824ms step_avg:39.61ms step:653/1600 train_time:25886ms step_avg:39.64ms step:654/1600 train_time:25945ms step_avg:39.67ms step:655/1600 train_time:26008ms step_avg:39.71ms step:656/1600 train_time:26066ms step_avg:39.73ms step:657/1600 train_time:26129ms step_avg:39.77ms step:658/1600 train_time:26188ms step_avg:39.80ms step:659/1600 train_time:26253ms step_avg:39.84ms step:660/1600 train_time:26312ms step_avg:39.87ms step:661/1600 train_time:26372ms step_avg:39.90ms step:662/1600 train_time:26430ms step_avg:39.92ms step:663/1600 train_time:26492ms step_avg:39.96ms step:664/1600 train_time:26551ms step_avg:39.99ms step:665/1600 train_time:26612ms step_avg:40.02ms step:666/1600 train_time:26671ms step_avg:40.05ms step:667/1600 train_time:26732ms step_avg:40.08ms step:668/1600 train_time:26791ms step_avg:40.11ms step:669/1600 train_time:26854ms step_avg:40.14ms step:670/1600 train_time:26913ms step_avg:40.17ms step:671/1600 train_time:26976ms step_avg:40.20ms step:672/1600 train_time:27035ms step_avg:40.23ms step:673/1600 train_time:27098ms step_avg:40.26ms step:674/1600 train_time:27156ms step_avg:40.29ms step:675/1600 train_time:27219ms step_avg:40.33ms step:676/1600 train_time:27280ms step_avg:40.35ms step:677/1600 train_time:27341ms step_avg:40.39ms step:678/1600 train_time:27401ms step_avg:40.41ms step:679/1600 train_time:27462ms step_avg:40.44ms step:680/1600 train_time:27521ms step_avg:40.47ms step:681/1600 train_time:27583ms step_avg:40.50ms step:682/1600 train_time:27642ms step_avg:40.53ms step:683/1600 train_time:27705ms step_avg:40.56ms step:684/1600 train_time:27764ms step_avg:40.59ms step:685/1600 train_time:27826ms step_avg:40.62ms step:686/1600 train_time:27886ms step_avg:40.65ms step:687/1600 train_time:27949ms step_avg:40.68ms step:688/1600 train_time:28007ms step_avg:40.71ms step:689/1600 train_time:28069ms step_avg:40.74ms step:690/1600 train_time:28129ms step_avg:40.77ms step:691/1600 train_time:28192ms step_avg:40.80ms step:692/1600 train_time:28250ms step_avg:40.82ms step:693/1600 train_time:28312ms step_avg:40.85ms step:694/1600 train_time:28371ms step_avg:40.88ms step:695/1600 train_time:28434ms step_avg:40.91ms step:696/1600 train_time:28493ms step_avg:40.94ms step:697/1600 train_time:28555ms step_avg:40.97ms step:698/1600 train_time:28613ms step_avg:40.99ms step:699/1600 train_time:28676ms step_avg:41.02ms step:700/1600 train_time:28735ms step_avg:41.05ms step:701/1600 train_time:28798ms step_avg:41.08ms step:702/1600 train_time:28857ms step_avg:41.11ms step:703/1600 train_time:28919ms step_avg:41.14ms step:704/1600 train_time:28978ms step_avg:41.16ms step:705/1600 train_time:29041ms step_avg:41.19ms step:706/1600 train_time:29100ms step_avg:41.22ms step:707/1600 train_time:29163ms step_avg:41.25ms step:708/1600 train_time:29222ms step_avg:41.27ms step:709/1600 train_time:29285ms step_avg:41.30ms step:710/1600 train_time:29344ms step_avg:41.33ms step:711/1600 train_time:29407ms step_avg:41.36ms step:712/1600 train_time:29466ms step_avg:41.38ms step:713/1600 train_time:29529ms step_avg:41.41ms step:714/1600 train_time:29588ms step_avg:41.44ms step:715/1600 train_time:29650ms step_avg:41.47ms step:716/1600 train_time:29709ms step_avg:41.49ms step:717/1600 train_time:29771ms step_avg:41.52ms step:718/1600 train_time:29830ms step_avg:41.55ms step:719/1600 train_time:29892ms step_avg:41.57ms step:720/1600 train_time:29952ms step_avg:41.60ms step:721/1600 train_time:30013ms step_avg:41.63ms step:722/1600 train_time:30072ms step_avg:41.65ms step:723/1600 train_time:30135ms step_avg:41.68ms step:724/1600 train_time:30194ms step_avg:41.70ms step:725/1600 train_time:30257ms step_avg:41.73ms step:726/1600 train_time:30317ms step_avg:41.76ms step:727/1600 train_time:30379ms step_avg:41.79ms step:728/1600 train_time:30438ms step_avg:41.81ms step:729/1600 train_time:30500ms step_avg:41.84ms step:730/1600 train_time:30559ms step_avg:41.86ms step:731/1600 train_time:30622ms step_avg:41.89ms step:732/1600 train_time:30681ms step_avg:41.91ms step:733/1600 train_time:30745ms step_avg:41.94ms step:734/1600 train_time:30804ms step_avg:41.97ms step:735/1600 train_time:30867ms step_avg:42.00ms step:736/1600 train_time:30926ms step_avg:42.02ms step:737/1600 train_time:30990ms step_avg:42.05ms step:738/1600 train_time:31048ms step_avg:42.07ms step:739/1600 train_time:31110ms step_avg:42.10ms step:740/1600 train_time:31168ms step_avg:42.12ms step:741/1600 train_time:31231ms step_avg:42.15ms step:742/1600 train_time:31290ms step_avg:42.17ms step:743/1600 train_time:31357ms step_avg:42.20ms step:744/1600 train_time:31415ms step_avg:42.22ms step:745/1600 train_time:31474ms step_avg:42.25ms step:746/1600 train_time:31533ms step_avg:42.27ms step:747/1600 train_time:31597ms step_avg:42.30ms step:748/1600 train_time:31655ms step_avg:42.32ms step:749/1600 train_time:31717ms step_avg:42.35ms step:750/1600 train_time:31776ms step_avg:42.37ms step:750/1600 val_loss:3.8892 train_time:31824ms step_avg:42.43ms step:751/1600 train_time:31843ms step_avg:42.40ms step:752/1600 train_time:31902ms step_avg:42.42ms step:753/1600 train_time:31965ms step_avg:42.45ms step:754/1600 train_time:32028ms step_avg:42.48ms step:755/1600 train_time:32091ms step_avg:42.50ms step:756/1600 train_time:32151ms step_avg:42.53ms step:757/1600 train_time:32212ms step_avg:42.55ms step:758/1600 train_time:32271ms step_avg:42.57ms step:759/1600 train_time:32333ms step_avg:42.60ms step:760/1600 train_time:32392ms step_avg:42.62ms step:761/1600 train_time:32454ms step_avg:42.65ms step:762/1600 train_time:32514ms step_avg:42.67ms step:763/1600 train_time:32575ms step_avg:42.69ms step:764/1600 train_time:32634ms step_avg:42.72ms step:765/1600 train_time:32697ms step_avg:42.74ms step:766/1600 train_time:32756ms step_avg:42.76ms step:767/1600 train_time:32819ms step_avg:42.79ms step:768/1600 train_time:32879ms step_avg:42.81ms step:769/1600 train_time:32942ms step_avg:42.84ms step:770/1600 train_time:33001ms step_avg:42.86ms step:771/1600 train_time:33064ms step_avg:42.88ms step:772/1600 train_time:33123ms step_avg:42.91ms step:773/1600 train_time:33186ms step_avg:42.93ms step:774/1600 train_time:33245ms step_avg:42.95ms step:775/1600 train_time:33308ms step_avg:42.98ms step:776/1600 train_time:33367ms step_avg:43.00ms step:777/1600 train_time:33429ms step_avg:43.02ms step:778/1600 train_time:33488ms step_avg:43.04ms step:779/1600 train_time:33552ms step_avg:43.07ms step:780/1600 train_time:33609ms step_avg:43.09ms step:781/1600 train_time:33671ms step_avg:43.11ms step:782/1600 train_time:33730ms step_avg:43.13ms step:783/1600 train_time:33794ms step_avg:43.16ms step:784/1600 train_time:33854ms step_avg:43.18ms step:785/1600 train_time:33919ms step_avg:43.21ms step:786/1600 train_time:33978ms step_avg:43.23ms step:787/1600 train_time:34041ms step_avg:43.25ms step:788/1600 train_time:34099ms step_avg:43.27ms step:789/1600 train_time:34162ms step_avg:43.30ms step:790/1600 train_time:34221ms step_avg:43.32ms step:791/1600 train_time:34283ms step_avg:43.34ms step:792/1600 train_time:34342ms step_avg:43.36ms step:793/1600 train_time:34404ms step_avg:43.38ms step:794/1600 train_time:34463ms step_avg:43.40ms step:795/1600 train_time:34525ms step_avg:43.43ms step:796/1600 train_time:34584ms step_avg:43.45ms step:797/1600 train_time:34647ms step_avg:43.47ms step:798/1600 train_time:34705ms step_avg:43.49ms step:799/1600 train_time:34769ms step_avg:43.52ms step:800/1600 train_time:34829ms step_avg:43.54ms step:801/1600 train_time:34893ms step_avg:43.56ms step:802/1600 train_time:34952ms step_avg:43.58ms step:803/1600 train_time:35016ms step_avg:43.61ms step:804/1600 train_time:35075ms step_avg:43.63ms step:805/1600 train_time:35138ms step_avg:43.65ms step:806/1600 train_time:35198ms step_avg:43.67ms step:807/1600 train_time:35260ms step_avg:43.69ms step:808/1600 train_time:35319ms step_avg:43.71ms step:809/1600 train_time:35381ms step_avg:43.73ms step:810/1600 train_time:35440ms step_avg:43.75ms step:811/1600 train_time:35502ms step_avg:43.78ms step:812/1600 train_time:35561ms step_avg:43.79ms step:813/1600 train_time:35626ms step_avg:43.82ms step:814/1600 train_time:35686ms step_avg:43.84ms step:815/1600 train_time:35746ms step_avg:43.86ms step:816/1600 train_time:35805ms step_avg:43.88ms step:817/1600 train_time:35868ms step_avg:43.90ms step:818/1600 train_time:35928ms step_avg:43.92ms step:819/1600 train_time:35991ms step_avg:43.95ms step:820/1600 train_time:36051ms step_avg:43.96ms step:821/1600 train_time:36113ms step_avg:43.99ms step:822/1600 train_time:36173ms step_avg:44.01ms step:823/1600 train_time:36236ms step_avg:44.03ms step:824/1600 train_time:36295ms step_avg:44.05ms step:825/1600 train_time:36358ms step_avg:44.07ms step:826/1600 train_time:36417ms step_avg:44.09ms step:827/1600 train_time:36479ms step_avg:44.11ms step:828/1600 train_time:36538ms step_avg:44.13ms step:829/1600 train_time:36601ms step_avg:44.15ms step:830/1600 train_time:36659ms step_avg:44.17ms step:831/1600 train_time:36722ms step_avg:44.19ms step:832/1600 train_time:36782ms step_avg:44.21ms step:833/1600 train_time:36846ms step_avg:44.23ms step:834/1600 train_time:36905ms step_avg:44.25ms step:835/1600 train_time:36967ms step_avg:44.27ms step:836/1600 train_time:37026ms step_avg:44.29ms step:837/1600 train_time:37089ms step_avg:44.31ms step:838/1600 train_time:37148ms step_avg:44.33ms step:839/1600 train_time:37211ms step_avg:44.35ms step:840/1600 train_time:37270ms step_avg:44.37ms step:841/1600 train_time:37333ms step_avg:44.39ms step:842/1600 train_time:37393ms step_avg:44.41ms step:843/1600 train_time:37455ms step_avg:44.43ms step:844/1600 train_time:37514ms step_avg:44.45ms step:845/1600 train_time:37576ms step_avg:44.47ms step:846/1600 train_time:37636ms step_avg:44.49ms step:847/1600 train_time:37698ms step_avg:44.51ms step:848/1600 train_time:37757ms step_avg:44.53ms step:849/1600 train_time:37820ms step_avg:44.55ms step:850/1600 train_time:37879ms step_avg:44.56ms step:851/1600 train_time:37941ms step_avg:44.58ms step:852/1600 train_time:38001ms step_avg:44.60ms step:853/1600 train_time:38063ms step_avg:44.62ms step:854/1600 train_time:38125ms step_avg:44.64ms step:855/1600 train_time:38185ms step_avg:44.66ms step:856/1600 train_time:38244ms step_avg:44.68ms step:857/1600 train_time:38307ms step_avg:44.70ms step:858/1600 train_time:38367ms step_avg:44.72ms step:859/1600 train_time:38430ms step_avg:44.74ms step:860/1600 train_time:38489ms step_avg:44.75ms step:861/1600 train_time:38551ms step_avg:44.78ms step:862/1600 train_time:38611ms step_avg:44.79ms step:863/1600 train_time:38673ms step_avg:44.81ms step:864/1600 train_time:38733ms step_avg:44.83ms step:865/1600 train_time:38796ms step_avg:44.85ms step:866/1600 train_time:38855ms step_avg:44.87ms step:867/1600 train_time:38919ms step_avg:44.89ms step:868/1600 train_time:38978ms step_avg:44.91ms step:869/1600 train_time:39040ms step_avg:44.93ms step:870/1600 train_time:39099ms step_avg:44.94ms step:871/1600 train_time:39162ms step_avg:44.96ms step:872/1600 train_time:39221ms step_avg:44.98ms step:873/1600 train_time:39283ms step_avg:45.00ms step:874/1600 train_time:39342ms step_avg:45.01ms step:875/1600 train_time:39404ms step_avg:45.03ms step:876/1600 train_time:39462ms step_avg:45.05ms step:877/1600 train_time:39525ms step_avg:45.07ms step:878/1600 train_time:39585ms step_avg:45.08ms step:879/1600 train_time:39647ms step_avg:45.10ms step:880/1600 train_time:39707ms step_avg:45.12ms step:881/1600 train_time:39769ms step_avg:45.14ms step:882/1600 train_time:39830ms step_avg:45.16ms step:883/1600 train_time:39893ms step_avg:45.18ms step:884/1600 train_time:39952ms step_avg:45.19ms step:885/1600 train_time:40016ms step_avg:45.22ms step:886/1600 train_time:40075ms step_avg:45.23ms step:887/1600 train_time:40139ms step_avg:45.25ms step:888/1600 train_time:40198ms step_avg:45.27ms step:889/1600 train_time:40260ms step_avg:45.29ms step:890/1600 train_time:40324ms step_avg:45.31ms step:891/1600 train_time:40383ms step_avg:45.32ms step:892/1600 train_time:40442ms step_avg:45.34ms step:893/1600 train_time:40503ms step_avg:45.36ms step:894/1600 train_time:40562ms step_avg:45.37ms step:895/1600 train_time:40624ms step_avg:45.39ms step:896/1600 train_time:40684ms step_avg:45.41ms step:897/1600 train_time:40748ms step_avg:45.43ms step:898/1600 train_time:40807ms step_avg:45.44ms step:899/1600 train_time:40868ms step_avg:45.46ms step:900/1600 train_time:40927ms step_avg:45.47ms step:901/1600 train_time:40990ms step_avg:45.49ms step:902/1600 train_time:41050ms step_avg:45.51ms step:903/1600 train_time:41113ms step_avg:45.53ms step:904/1600 train_time:41173ms step_avg:45.54ms step:905/1600 train_time:41236ms step_avg:45.56ms step:906/1600 train_time:41295ms step_avg:45.58ms step:907/1600 train_time:41358ms step_avg:45.60ms step:908/1600 train_time:41420ms step_avg:45.62ms step:909/1600 train_time:41482ms step_avg:45.63ms step:910/1600 train_time:41539ms step_avg:45.65ms step:911/1600 train_time:41601ms step_avg:45.66ms step:912/1600 train_time:41660ms step_avg:45.68ms step:913/1600 train_time:41723ms step_avg:45.70ms step:914/1600 train_time:41782ms step_avg:45.71ms step:915/1600 train_time:41845ms step_avg:45.73ms step:916/1600 train_time:41904ms step_avg:45.75ms step:917/1600 train_time:41967ms step_avg:45.77ms step:918/1600 train_time:42026ms step_avg:45.78ms step:919/1600 train_time:42089ms step_avg:45.80ms step:920/1600 train_time:42148ms step_avg:45.81ms step:921/1600 train_time:42212ms step_avg:45.83ms step:922/1600 train_time:42272ms step_avg:45.85ms step:923/1600 train_time:42335ms step_avg:45.87ms step:924/1600 train_time:42394ms step_avg:45.88ms step:925/1600 train_time:42457ms step_avg:45.90ms step:926/1600 train_time:42516ms step_avg:45.91ms step:927/1600 train_time:42579ms step_avg:45.93ms step:928/1600 train_time:42638ms step_avg:45.95ms step:929/1600 train_time:42700ms step_avg:45.96ms step:930/1600 train_time:42760ms step_avg:45.98ms step:931/1600 train_time:42822ms step_avg:46.00ms step:932/1600 train_time:42881ms step_avg:46.01ms step:933/1600 train_time:42943ms step_avg:46.03ms step:934/1600 train_time:43002ms step_avg:46.04ms step:935/1600 train_time:43065ms step_avg:46.06ms step:936/1600 train_time:43124ms step_avg:46.07ms step:937/1600 train_time:43187ms step_avg:46.09ms step:938/1600 train_time:43247ms step_avg:46.11ms step:939/1600 train_time:43310ms step_avg:46.12ms step:940/1600 train_time:43370ms step_avg:46.14ms step:941/1600 train_time:43433ms step_avg:46.16ms step:942/1600 train_time:43492ms step_avg:46.17ms step:943/1600 train_time:43554ms step_avg:46.19ms step:944/1600 train_time:43614ms step_avg:46.20ms step:945/1600 train_time:43677ms step_avg:46.22ms step:946/1600 train_time:43737ms step_avg:46.23ms step:947/1600 train_time:43799ms step_avg:46.25ms step:948/1600 train_time:43859ms step_avg:46.26ms step:949/1600 train_time:43921ms step_avg:46.28ms step:950/1600 train_time:43980ms step_avg:46.29ms step:951/1600 train_time:44042ms step_avg:46.31ms step:952/1600 train_time:44101ms step_avg:46.32ms step:953/1600 train_time:44163ms step_avg:46.34ms step:954/1600 train_time:44222ms step_avg:46.35ms step:955/1600 train_time:44285ms step_avg:46.37ms step:956/1600 train_time:44343ms step_avg:46.38ms step:957/1600 train_time:44406ms step_avg:46.40ms step:958/1600 train_time:44465ms step_avg:46.41ms step:959/1600 train_time:44529ms step_avg:46.43ms step:960/1600 train_time:44588ms step_avg:46.45ms step:961/1600 train_time:44651ms step_avg:46.46ms step:962/1600 train_time:44710ms step_avg:46.48ms step:963/1600 train_time:44772ms step_avg:46.49ms step:964/1600 train_time:44832ms step_avg:46.51ms step:965/1600 train_time:44895ms step_avg:46.52ms step:966/1600 train_time:44955ms step_avg:46.54ms step:967/1600 train_time:45019ms step_avg:46.56ms step:968/1600 train_time:45078ms step_avg:46.57ms step:969/1600 train_time:45140ms step_avg:46.58ms step:970/1600 train_time:45199ms step_avg:46.60ms step:971/1600 train_time:45261ms step_avg:46.61ms step:972/1600 train_time:45320ms step_avg:46.63ms step:973/1600 train_time:45383ms step_avg:46.64ms step:974/1600 train_time:45442ms step_avg:46.65ms step:975/1600 train_time:45504ms step_avg:46.67ms step:976/1600 train_time:45563ms step_avg:46.68ms step:977/1600 train_time:45627ms step_avg:46.70ms step:978/1600 train_time:45686ms step_avg:46.71ms step:979/1600 train_time:45749ms step_avg:46.73ms step:980/1600 train_time:45810ms step_avg:46.75ms step:981/1600 train_time:45875ms step_avg:46.76ms step:982/1600 train_time:45935ms step_avg:46.78ms step:983/1600 train_time:45995ms step_avg:46.79ms step:984/1600 train_time:46054ms step_avg:46.80ms step:985/1600 train_time:46117ms step_avg:46.82ms step:986/1600 train_time:46176ms step_avg:46.83ms step:987/1600 train_time:46239ms step_avg:46.85ms step:988/1600 train_time:46301ms step_avg:46.86ms step:989/1600 train_time:46361ms step_avg:46.88ms step:990/1600 train_time:46420ms step_avg:46.89ms step:991/1600 train_time:46481ms step_avg:46.90ms step:992/1600 train_time:46540ms step_avg:46.92ms step:993/1600 train_time:46603ms step_avg:46.93ms step:994/1600 train_time:46661ms step_avg:46.94ms step:995/1600 train_time:46723ms step_avg:46.96ms step:996/1600 train_time:46783ms step_avg:46.97ms step:997/1600 train_time:46845ms step_avg:46.99ms step:998/1600 train_time:46906ms step_avg:47.00ms step:999/1600 train_time:46968ms step_avg:47.02ms step:1000/1600 train_time:47028ms step_avg:47.03ms step:1000/1600 val_loss:3.5991 train_time:47075ms step_avg:47.08ms step:1001/1600 train_time:47094ms step_avg:47.05ms step:1002/1600 train_time:47152ms step_avg:47.06ms step:1003/1600 train_time:47217ms step_avg:47.08ms step:1004/1600 train_time:47281ms step_avg:47.09ms step:1005/1600 train_time:47342ms step_avg:47.11ms step:1006/1600 train_time:47401ms step_avg:47.12ms step:1007/1600 train_time:47462ms step_avg:47.13ms step:1008/1600 train_time:47521ms step_avg:47.14ms step:1009/1600 train_time:47584ms step_avg:47.16ms step:1010/1600 train_time:47642ms step_avg:47.17ms step:1011/1600 train_time:47704ms step_avg:47.18ms step:1012/1600 train_time:47762ms step_avg:47.20ms step:1013/1600 train_time:47823ms step_avg:47.21ms step:1014/1600 train_time:47882ms step_avg:47.22ms step:1015/1600 train_time:47944ms step_avg:47.24ms step:1016/1600 train_time:48003ms step_avg:47.25ms step:1017/1600 train_time:48066ms step_avg:47.26ms step:1018/1600 train_time:48127ms step_avg:47.28ms step:1019/1600 train_time:48191ms step_avg:47.29ms step:1020/1600 train_time:48251ms step_avg:47.30ms step:1021/1600 train_time:48314ms step_avg:47.32ms step:1022/1600 train_time:48372ms step_avg:47.33ms step:1023/1600 train_time:48435ms step_avg:47.35ms step:1024/1600 train_time:48494ms step_avg:47.36ms step:1025/1600 train_time:48557ms step_avg:47.37ms step:1026/1600 train_time:48615ms step_avg:47.38ms step:1027/1600 train_time:48677ms step_avg:47.40ms step:1028/1600 train_time:48737ms step_avg:47.41ms step:1029/1600 train_time:48799ms step_avg:47.42ms step:1030/1600 train_time:48858ms step_avg:47.44ms step:1031/1600 train_time:48921ms step_avg:47.45ms step:1032/1600 train_time:48980ms step_avg:47.46ms step:1033/1600 train_time:49044ms step_avg:47.48ms step:1034/1600 train_time:49103ms step_avg:47.49ms step:1035/1600 train_time:49167ms step_avg:47.50ms step:1036/1600 train_time:49226ms step_avg:47.52ms step:1037/1600 train_time:49288ms step_avg:47.53ms step:1038/1600 train_time:49347ms step_avg:47.54ms step:1039/1600 train_time:49409ms step_avg:47.55ms step:1040/1600 train_time:49468ms step_avg:47.57ms step:1041/1600 train_time:49539ms step_avg:47.59ms step:1042/1600 train_time:49622ms step_avg:47.62ms step:1043/1600 train_time:49712ms step_avg:47.66ms step:1044/1600 train_time:49796ms step_avg:47.70ms step:1045/1600 train_time:49885ms step_avg:47.74ms step:1046/1600 train_time:49970ms step_avg:47.77ms step:1047/1600 train_time:50061ms step_avg:47.81ms step:1048/1600 train_time:50146ms step_avg:47.85ms step:1049/1600 train_time:50233ms step_avg:47.89ms step:1050/1600 train_time:50318ms step_avg:47.92ms step:1051/1600 train_time:50408ms step_avg:47.96ms step:1052/1600 train_time:50494ms step_avg:48.00ms step:1053/1600 train_time:50582ms step_avg:48.04ms step:1054/1600 train_time:50667ms step_avg:48.07ms step:1055/1600 train_time:50755ms step_avg:48.11ms step:1056/1600 train_time:50840ms step_avg:48.14ms step:1057/1600 train_time:50929ms step_avg:48.18ms step:1058/1600 train_time:51014ms step_avg:48.22ms step:1059/1600 train_time:51102ms step_avg:48.25ms step:1060/1600 train_time:51188ms step_avg:48.29ms step:1061/1600 train_time:51277ms step_avg:48.33ms step:1062/1600 train_time:51362ms step_avg:48.36ms step:1063/1600 train_time:51450ms step_avg:48.40ms step:1064/1600 train_time:51534ms step_avg:48.43ms step:1065/1600 train_time:51623ms step_avg:48.47ms step:1066/1600 train_time:51709ms step_avg:48.51ms step:1067/1600 train_time:51798ms step_avg:48.55ms step:1068/1600 train_time:51884ms step_avg:48.58ms step:1069/1600 train_time:51972ms step_avg:48.62ms step:1070/1600 train_time:52057ms step_avg:48.65ms step:1071/1600 train_time:52146ms step_avg:48.69ms step:1072/1600 train_time:52231ms step_avg:48.72ms step:1073/1600 train_time:52320ms step_avg:48.76ms step:1074/1600 train_time:52405ms step_avg:48.79ms step:1075/1600 train_time:52494ms step_avg:48.83ms step:1076/1600 train_time:52578ms step_avg:48.86ms step:1077/1600 train_time:52668ms step_avg:48.90ms step:1078/1600 train_time:52752ms step_avg:48.94ms step:1079/1600 train_time:52840ms step_avg:48.97ms step:1080/1600 train_time:52926ms step_avg:49.01ms step:1081/1600 train_time:53014ms step_avg:49.04ms step:1082/1600 train_time:53099ms step_avg:49.08ms step:1083/1600 train_time:53188ms step_avg:49.11ms step:1084/1600 train_time:53274ms step_avg:49.15ms step:1085/1600 train_time:53362ms step_avg:49.18ms step:1086/1600 train_time:53448ms step_avg:49.22ms step:1087/1600 train_time:53536ms step_avg:49.25ms step:1088/1600 train_time:53621ms step_avg:49.28ms step:1089/1600 train_time:53710ms step_avg:49.32ms step:1090/1600 train_time:53795ms step_avg:49.35ms step:1091/1600 train_time:53883ms step_avg:49.39ms step:1092/1600 train_time:53969ms step_avg:49.42ms step:1093/1600 train_time:54058ms step_avg:49.46ms step:1094/1600 train_time:54144ms step_avg:49.49ms step:1095/1600 train_time:54232ms step_avg:49.53ms step:1096/1600 train_time:54317ms step_avg:49.56ms step:1097/1600 train_time:54405ms step_avg:49.59ms step:1098/1600 train_time:54490ms step_avg:49.63ms step:1099/1600 train_time:54579ms step_avg:49.66ms step:1100/1600 train_time:54665ms step_avg:49.70ms step:1101/1600 train_time:54753ms step_avg:49.73ms step:1102/1600 train_time:54838ms step_avg:49.76ms step:1103/1600 train_time:54926ms step_avg:49.80ms step:1104/1600 train_time:55011ms step_avg:49.83ms step:1105/1600 train_time:55099ms step_avg:49.86ms step:1106/1600 train_time:55187ms step_avg:49.90ms step:1107/1600 train_time:55274ms step_avg:49.93ms step:1108/1600 train_time:55358ms step_avg:49.96ms step:1109/1600 train_time:55447ms step_avg:50.00ms step:1110/1600 train_time:55532ms step_avg:50.03ms step:1111/1600 train_time:55620ms step_avg:50.06ms step:1112/1600 train_time:55706ms step_avg:50.10ms step:1113/1600 train_time:55794ms step_avg:50.13ms step:1114/1600 train_time:55878ms step_avg:50.16ms step:1115/1600 train_time:55967ms step_avg:50.19ms step:1116/1600 train_time:56052ms step_avg:50.23ms step:1117/1600 train_time:56140ms step_avg:50.26ms step:1118/1600 train_time:56227ms step_avg:50.29ms step:1119/1600 train_time:56315ms step_avg:50.33ms step:1120/1600 train_time:56400ms step_avg:50.36ms step:1121/1600 train_time:56489ms step_avg:50.39ms step:1122/1600 train_time:56573ms step_avg:50.42ms step:1123/1600 train_time:56662ms step_avg:50.46ms step:1124/1600 train_time:56747ms step_avg:50.49ms step:1125/1600 train_time:56836ms step_avg:50.52ms step:1126/1600 train_time:56921ms step_avg:50.55ms step:1127/1600 train_time:57009ms step_avg:50.59ms step:1128/1600 train_time:57095ms step_avg:50.62ms step:1129/1600 train_time:57184ms step_avg:50.65ms step:1130/1600 train_time:57272ms step_avg:50.68ms step:1131/1600 train_time:57357ms step_avg:50.71ms step:1132/1600 train_time:57452ms step_avg:50.75ms step:1133/1600 train_time:57537ms step_avg:50.78ms step:1134/1600 train_time:57622ms step_avg:50.81ms step:1135/1600 train_time:57710ms step_avg:50.85ms step:1136/1600 train_time:57794ms step_avg:50.88ms step:1137/1600 train_time:57883ms step_avg:50.91ms step:1138/1600 train_time:57967ms step_avg:50.94ms step:1139/1600 train_time:58051ms step_avg:50.97ms step:1140/1600 train_time:58136ms step_avg:51.00ms step:1141/1600 train_time:58225ms step_avg:51.03ms step:1142/1600 train_time:58312ms step_avg:51.06ms step:1143/1600 train_time:58399ms step_avg:51.09ms step:1144/1600 train_time:58484ms step_avg:51.12ms step:1145/1600 train_time:58572ms step_avg:51.15ms step:1146/1600 train_time:58657ms step_avg:51.18ms step:1147/1600 train_time:58745ms step_avg:51.22ms step:1148/1600 train_time:58831ms step_avg:51.25ms step:1149/1600 train_time:58919ms step_avg:51.28ms step:1150/1600 train_time:59004ms step_avg:51.31ms step:1151/1600 train_time:59092ms step_avg:51.34ms step:1152/1600 train_time:59177ms step_avg:51.37ms step:1153/1600 train_time:59266ms step_avg:51.40ms step:1154/1600 train_time:59355ms step_avg:51.43ms step:1155/1600 train_time:59441ms step_avg:51.46ms step:1156/1600 train_time:59527ms step_avg:51.49ms step:1157/1600 train_time:59615ms step_avg:51.53ms step:1158/1600 train_time:59700ms step_avg:51.55ms step:1159/1600 train_time:59788ms step_avg:51.59ms step:1160/1600 train_time:59874ms step_avg:51.62ms step:1161/1600 train_time:59961ms step_avg:51.65ms step:1162/1600 train_time:60047ms step_avg:51.68ms step:1163/1600 train_time:60135ms step_avg:51.71ms step:1164/1600 train_time:60220ms step_avg:51.74ms step:1165/1600 train_time:60312ms step_avg:51.77ms step:1166/1600 train_time:60395ms step_avg:51.80ms step:1167/1600 train_time:60482ms step_avg:51.83ms step:1168/1600 train_time:60569ms step_avg:51.86ms step:1169/1600 train_time:60657ms step_avg:51.89ms step:1170/1600 train_time:60742ms step_avg:51.92ms step:1171/1600 train_time:60831ms step_avg:51.95ms step:1172/1600 train_time:60915ms step_avg:51.98ms step:1173/1600 train_time:61004ms step_avg:52.01ms step:1174/1600 train_time:61089ms step_avg:52.03ms step:1175/1600 train_time:61178ms step_avg:52.07ms step:1176/1600 train_time:61264ms step_avg:52.09ms step:1177/1600 train_time:61352ms step_avg:52.13ms step:1178/1600 train_time:61437ms step_avg:52.15ms step:1179/1600 train_time:61525ms step_avg:52.18ms step:1180/1600 train_time:61610ms step_avg:52.21ms step:1181/1600 train_time:61699ms step_avg:52.24ms step:1182/1600 train_time:61785ms step_avg:52.27ms step:1183/1600 train_time:61874ms step_avg:52.30ms step:1184/1600 train_time:61958ms step_avg:52.33ms step:1185/1600 train_time:62047ms step_avg:52.36ms step:1186/1600 train_time:62131ms step_avg:52.39ms step:1187/1600 train_time:62220ms step_avg:52.42ms step:1188/1600 train_time:62305ms step_avg:52.45ms step:1189/1600 train_time:62394ms step_avg:52.48ms step:1190/1600 train_time:62478ms step_avg:52.50ms step:1191/1600 train_time:62567ms step_avg:52.53ms step:1192/1600 train_time:62651ms step_avg:52.56ms step:1193/1600 train_time:62740ms step_avg:52.59ms step:1194/1600 train_time:62825ms step_avg:52.62ms step:1195/1600 train_time:62913ms step_avg:52.65ms step:1196/1600 train_time:62998ms step_avg:52.67ms step:1197/1600 train_time:63087ms step_avg:52.70ms step:1198/1600 train_time:63174ms step_avg:52.73ms step:1199/1600 train_time:63262ms step_avg:52.76ms step:1200/1600 train_time:63347ms step_avg:52.79ms step:1201/1600 train_time:63435ms step_avg:52.82ms step:1202/1600 train_time:63520ms step_avg:52.85ms step:1203/1600 train_time:63609ms step_avg:52.87ms step:1204/1600 train_time:63693ms step_avg:52.90ms step:1205/1600 train_time:63781ms step_avg:52.93ms step:1206/1600 train_time:63866ms step_avg:52.96ms step:1207/1600 train_time:63956ms step_avg:52.99ms step:1208/1600 train_time:64042ms step_avg:53.01ms step:1209/1600 train_time:64131ms step_avg:53.04ms step:1210/1600 train_time:64216ms step_avg:53.07ms step:1211/1600 train_time:64304ms step_avg:53.10ms step:1212/1600 train_time:64389ms step_avg:53.13ms step:1213/1600 train_time:64478ms step_avg:53.16ms step:1214/1600 train_time:64564ms step_avg:53.18ms step:1215/1600 train_time:64652ms step_avg:53.21ms step:1216/1600 train_time:64737ms step_avg:53.24ms step:1217/1600 train_time:64825ms step_avg:53.27ms step:1218/1600 train_time:64911ms step_avg:53.29ms step:1219/1600 train_time:65000ms step_avg:53.32ms step:1220/1600 train_time:65086ms step_avg:53.35ms step:1221/1600 train_time:65175ms step_avg:53.38ms step:1222/1600 train_time:65260ms step_avg:53.40ms step:1223/1600 train_time:65349ms step_avg:53.43ms step:1224/1600 train_time:65437ms step_avg:53.46ms step:1225/1600 train_time:65526ms step_avg:53.49ms step:1226/1600 train_time:65609ms step_avg:53.51ms step:1227/1600 train_time:65697ms step_avg:53.54ms step:1228/1600 train_time:65782ms step_avg:53.57ms step:1229/1600 train_time:65870ms step_avg:53.60ms step:1230/1600 train_time:65954ms step_avg:53.62ms step:1231/1600 train_time:66043ms step_avg:53.65ms step:1232/1600 train_time:66127ms step_avg:53.67ms step:1233/1600 train_time:66216ms step_avg:53.70ms step:1234/1600 train_time:66300ms step_avg:53.73ms step:1235/1600 train_time:66389ms step_avg:53.76ms step:1236/1600 train_time:66475ms step_avg:53.78ms step:1237/1600 train_time:66562ms step_avg:53.81ms step:1238/1600 train_time:66648ms step_avg:53.83ms step:1239/1600 train_time:66736ms step_avg:53.86ms step:1240/1600 train_time:66821ms step_avg:53.89ms step:1241/1600 train_time:66909ms step_avg:53.92ms step:1242/1600 train_time:66994ms step_avg:53.94ms step:1243/1600 train_time:67082ms step_avg:53.97ms step:1244/1600 train_time:67166ms step_avg:53.99ms step:1245/1600 train_time:67255ms step_avg:54.02ms step:1246/1600 train_time:67340ms step_avg:54.04ms step:1247/1600 train_time:67429ms step_avg:54.07ms step:1248/1600 train_time:67514ms step_avg:54.10ms step:1249/1600 train_time:67603ms step_avg:54.13ms step:1250/1600 train_time:67687ms step_avg:54.15ms step:1250/1600 val_loss:3.4164 train_time:67761ms step_avg:54.21ms step:1251/1600 train_time:67781ms step_avg:54.18ms step:1252/1600 train_time:67868ms step_avg:54.21ms step:1253/1600 train_time:67960ms step_avg:54.24ms step:1254/1600 train_time:68046ms step_avg:54.26ms step:1255/1600 train_time:68133ms step_avg:54.29ms step:1256/1600 train_time:68218ms step_avg:54.31ms step:1257/1600 train_time:68304ms step_avg:54.34ms step:1258/1600 train_time:68388ms step_avg:54.36ms step:1259/1600 train_time:68475ms step_avg:54.39ms step:1260/1600 train_time:68560ms step_avg:54.41ms step:1261/1600 train_time:68648ms step_avg:54.44ms step:1262/1600 train_time:68733ms step_avg:54.46ms step:1263/1600 train_time:68825ms step_avg:54.49ms step:1264/1600 train_time:68913ms step_avg:54.52ms step:1265/1600 train_time:69002ms step_avg:54.55ms step:1266/1600 train_time:69088ms step_avg:54.57ms step:1267/1600 train_time:69176ms step_avg:54.60ms step:1268/1600 train_time:69260ms step_avg:54.62ms step:1269/1600 train_time:69348ms step_avg:54.65ms step:1270/1600 train_time:69432ms step_avg:54.67ms step:1271/1600 train_time:69520ms step_avg:54.70ms step:1272/1600 train_time:69606ms step_avg:54.72ms step:1273/1600 train_time:69697ms step_avg:54.75ms step:1274/1600 train_time:69782ms step_avg:54.77ms step:1275/1600 train_time:69871ms step_avg:54.80ms step:1276/1600 train_time:69956ms step_avg:54.82ms step:1277/1600 train_time:70046ms step_avg:54.85ms step:1278/1600 train_time:70131ms step_avg:54.88ms step:1279/1600 train_time:70219ms step_avg:54.90ms step:1280/1600 train_time:70304ms step_avg:54.92ms step:1281/1600 train_time:70392ms step_avg:54.95ms step:1282/1600 train_time:70477ms step_avg:54.97ms step:1283/1600 train_time:70565ms step_avg:55.00ms step:1284/1600 train_time:70650ms step_avg:55.02ms step:1285/1600 train_time:70738ms step_avg:55.05ms step:1286/1600 train_time:70824ms step_avg:55.07ms step:1287/1600 train_time:70913ms step_avg:55.10ms step:1288/1600 train_time:70999ms step_avg:55.12ms step:1289/1600 train_time:71088ms step_avg:55.15ms step:1290/1600 train_time:71173ms step_avg:55.17ms step:1291/1600 train_time:71261ms step_avg:55.20ms step:1292/1600 train_time:71346ms step_avg:55.22ms step:1293/1600 train_time:71434ms step_avg:55.25ms step:1294/1600 train_time:71518ms step_avg:55.27ms step:1295/1600 train_time:71607ms step_avg:55.29ms step:1296/1600 train_time:71691ms step_avg:55.32ms step:1297/1600 train_time:71780ms step_avg:55.34ms step:1298/1600 train_time:71866ms step_avg:55.37ms step:1299/1600 train_time:71956ms step_avg:55.39ms step:1300/1600 train_time:72041ms step_avg:55.42ms step:1301/1600 train_time:72130ms step_avg:55.44ms step:1302/1600 train_time:72215ms step_avg:55.46ms step:1303/1600 train_time:72303ms step_avg:55.49ms step:1304/1600 train_time:72388ms step_avg:55.51ms step:1305/1600 train_time:72475ms step_avg:55.54ms step:1306/1600 train_time:72560ms step_avg:55.56ms step:1307/1600 train_time:72648ms step_avg:55.58ms step:1308/1600 train_time:72733ms step_avg:55.61ms step:1309/1600 train_time:72821ms step_avg:55.63ms step:1310/1600 train_time:72907ms step_avg:55.65ms step:1311/1600 train_time:72995ms step_avg:55.68ms step:1312/1600 train_time:73082ms step_avg:55.70ms step:1313/1600 train_time:73170ms step_avg:55.73ms step:1314/1600 train_time:73255ms step_avg:55.75ms step:1315/1600 train_time:73343ms step_avg:55.77ms step:1316/1600 train_time:73427ms step_avg:55.80ms step:1317/1600 train_time:73515ms step_avg:55.82ms step:1318/1600 train_time:73600ms step_avg:55.84ms step:1319/1600 train_time:73687ms step_avg:55.87ms step:1320/1600 train_time:73772ms step_avg:55.89ms step:1321/1600 train_time:73861ms step_avg:55.91ms step:1322/1600 train_time:73946ms step_avg:55.94ms step:1323/1600 train_time:74035ms step_avg:55.96ms step:1324/1600 train_time:74121ms step_avg:55.98ms step:1325/1600 train_time:74209ms step_avg:56.01ms step:1326/1600 train_time:74294ms step_avg:56.03ms step:1327/1600 train_time:74383ms step_avg:56.05ms step:1328/1600 train_time:74467ms step_avg:56.07ms step:1329/1600 train_time:74555ms step_avg:56.10ms step:1330/1600 train_time:74640ms step_avg:56.12ms step:1331/1600 train_time:74729ms step_avg:56.14ms step:1332/1600 train_time:74814ms step_avg:56.17ms step:1333/1600 train_time:74902ms step_avg:56.19ms step:1334/1600 train_time:74987ms step_avg:56.21ms step:1335/1600 train_time:75076ms step_avg:56.24ms step:1336/1600 train_time:75162ms step_avg:56.26ms step:1337/1600 train_time:75250ms step_avg:56.28ms step:1338/1600 train_time:75335ms step_avg:56.30ms step:1339/1600 train_time:75424ms step_avg:56.33ms step:1340/1600 train_time:75508ms step_avg:56.35ms step:1341/1600 train_time:75596ms step_avg:56.37ms step:1342/1600 train_time:75682ms step_avg:56.39ms step:1343/1600 train_time:75771ms step_avg:56.42ms step:1344/1600 train_time:75855ms step_avg:56.44ms step:1345/1600 train_time:75944ms step_avg:56.46ms step:1346/1600 train_time:76029ms step_avg:56.49ms step:1347/1600 train_time:76117ms step_avg:56.51ms step:1348/1600 train_time:76203ms step_avg:56.53ms step:1349/1600 train_time:76292ms step_avg:56.55ms step:1350/1600 train_time:76377ms step_avg:56.58ms step:1351/1600 train_time:76465ms step_avg:56.60ms step:1352/1600 train_time:76549ms step_avg:56.62ms step:1353/1600 train_time:76637ms step_avg:56.64ms step:1354/1600 train_time:76724ms step_avg:56.66ms step:1355/1600 train_time:76812ms step_avg:56.69ms step:1356/1600 train_time:76897ms step_avg:56.71ms step:1357/1600 train_time:76987ms step_avg:56.73ms step:1358/1600 train_time:77071ms step_avg:56.75ms step:1359/1600 train_time:77160ms step_avg:56.78ms step:1360/1600 train_time:77245ms step_avg:56.80ms step:1361/1600 train_time:77334ms step_avg:56.82ms step:1362/1600 train_time:77419ms step_avg:56.84ms step:1363/1600 train_time:77507ms step_avg:56.87ms step:1364/1600 train_time:77591ms step_avg:56.89ms step:1365/1600 train_time:77680ms step_avg:56.91ms step:1366/1600 train_time:77770ms step_avg:56.93ms step:1367/1600 train_time:77856ms step_avg:56.95ms step:1368/1600 train_time:77942ms step_avg:56.98ms step:1369/1600 train_time:78031ms step_avg:57.00ms step:1370/1600 train_time:78116ms step_avg:57.02ms step:1371/1600 train_time:78205ms step_avg:57.04ms step:1372/1600 train_time:78290ms step_avg:57.06ms step:1373/1600 train_time:78378ms step_avg:57.09ms step:1374/1600 train_time:78463ms step_avg:57.11ms step:1375/1600 train_time:78551ms step_avg:57.13ms step:1376/1600 train_time:78636ms step_avg:57.15ms step:1377/1600 train_time:78725ms step_avg:57.17ms step:1378/1600 train_time:78810ms step_avg:57.19ms step:1379/1600 train_time:78898ms step_avg:57.21ms step:1380/1600 train_time:78984ms step_avg:57.23ms step:1381/1600 train_time:79072ms step_avg:57.26ms step:1382/1600 train_time:79157ms step_avg:57.28ms step:1383/1600 train_time:79247ms step_avg:57.30ms step:1384/1600 train_time:79333ms step_avg:57.32ms step:1385/1600 train_time:79421ms step_avg:57.34ms step:1386/1600 train_time:79505ms step_avg:57.36ms step:1387/1600 train_time:79593ms step_avg:57.38ms step:1388/1600 train_time:79680ms step_avg:57.41ms step:1389/1600 train_time:79768ms step_avg:57.43ms step:1390/1600 train_time:79853ms step_avg:57.45ms step:1391/1600 train_time:79942ms step_avg:57.47ms step:1392/1600 train_time:80029ms step_avg:57.49ms step:1393/1600 train_time:80117ms step_avg:57.51ms step:1394/1600 train_time:80202ms step_avg:57.53ms step:1395/1600 train_time:80290ms step_avg:57.56ms step:1396/1600 train_time:80375ms step_avg:57.58ms step:1397/1600 train_time:80463ms step_avg:57.60ms step:1398/1600 train_time:80548ms step_avg:57.62ms step:1399/1600 train_time:80636ms step_avg:57.64ms step:1400/1600 train_time:80724ms step_avg:57.66ms step:1401/1600 train_time:80811ms step_avg:57.68ms step:1402/1600 train_time:80896ms step_avg:57.70ms step:1403/1600 train_time:80985ms step_avg:57.72ms step:1404/1600 train_time:81071ms step_avg:57.74ms step:1405/1600 train_time:81159ms step_avg:57.76ms step:1406/1600 train_time:81244ms step_avg:57.78ms step:1407/1600 train_time:81332ms step_avg:57.81ms step:1408/1600 train_time:81418ms step_avg:57.83ms step:1409/1600 train_time:81507ms step_avg:57.85ms step:1410/1600 train_time:81593ms step_avg:57.87ms step:1411/1600 train_time:81681ms step_avg:57.89ms step:1412/1600 train_time:81766ms step_avg:57.91ms step:1413/1600 train_time:81854ms step_avg:57.93ms step:1414/1600 train_time:81940ms step_avg:57.95ms step:1415/1600 train_time:82029ms step_avg:57.97ms step:1416/1600 train_time:82113ms step_avg:57.99ms step:1417/1600 train_time:82201ms step_avg:58.01ms step:1418/1600 train_time:82287ms step_avg:58.03ms step:1419/1600 train_time:82374ms step_avg:58.05ms step:1420/1600 train_time:82460ms step_avg:58.07ms step:1421/1600 train_time:82549ms step_avg:58.09ms step:1422/1600 train_time:82633ms step_avg:58.11ms step:1423/1600 train_time:82722ms step_avg:58.13ms step:1424/1600 train_time:82808ms step_avg:58.15ms step:1425/1600 train_time:82896ms step_avg:58.17ms step:1426/1600 train_time:82981ms step_avg:58.19ms step:1427/1600 train_time:83070ms step_avg:58.21ms step:1428/1600 train_time:83155ms step_avg:58.23ms step:1429/1600 train_time:83244ms step_avg:58.25ms step:1430/1600 train_time:83329ms step_avg:58.27ms step:1431/1600 train_time:83417ms step_avg:58.29ms step:1432/1600 train_time:83502ms step_avg:58.31ms step:1433/1600 train_time:83590ms step_avg:58.33ms step:1434/1600 train_time:83675ms step_avg:58.35ms step:1435/1600 train_time:83764ms step_avg:58.37ms step:1436/1600 train_time:83849ms step_avg:58.39ms step:1437/1600 train_time:83938ms step_avg:58.41ms step:1438/1600 train_time:84023ms step_avg:58.43ms step:1439/1600 train_time:84114ms step_avg:58.45ms step:1440/1600 train_time:84199ms step_avg:58.47ms step:1441/1600 train_time:84287ms step_avg:58.49ms step:1442/1600 train_time:84372ms step_avg:58.51ms step:1443/1600 train_time:84460ms step_avg:58.53ms step:1444/1600 train_time:84545ms step_avg:58.55ms step:1445/1600 train_time:84635ms step_avg:58.57ms step:1446/1600 train_time:84721ms step_avg:58.59ms step:1447/1600 train_time:84808ms step_avg:58.61ms step:1448/1600 train_time:84892ms step_avg:58.63ms step:1449/1600 train_time:84980ms step_avg:58.65ms step:1450/1600 train_time:85068ms step_avg:58.67ms step:1451/1600 train_time:85154ms step_avg:58.69ms step:1452/1600 train_time:85239ms step_avg:58.70ms step:1453/1600 train_time:85328ms step_avg:58.73ms step:1454/1600 train_time:85413ms step_avg:58.74ms step:1455/1600 train_time:85501ms step_avg:58.76ms step:1456/1600 train_time:85586ms step_avg:58.78ms step:1457/1600 train_time:85674ms step_avg:58.80ms step:1458/1600 train_time:85760ms step_avg:58.82ms step:1459/1600 train_time:85848ms step_avg:58.84ms step:1460/1600 train_time:85933ms step_avg:58.86ms step:1461/1600 train_time:86022ms step_avg:58.88ms step:1462/1600 train_time:86108ms step_avg:58.90ms step:1463/1600 train_time:86196ms step_avg:58.92ms step:1464/1600 train_time:86282ms step_avg:58.94ms step:1465/1600 train_time:86371ms step_avg:58.96ms step:1466/1600 train_time:86456ms step_avg:58.97ms step:1467/1600 train_time:86544ms step_avg:58.99ms step:1468/1600 train_time:86629ms step_avg:59.01ms step:1469/1600 train_time:86717ms step_avg:59.03ms step:1470/1600 train_time:86802ms step_avg:59.05ms step:1471/1600 train_time:86891ms step_avg:59.07ms step:1472/1600 train_time:86976ms step_avg:59.09ms step:1473/1600 train_time:87065ms step_avg:59.11ms step:1474/1600 train_time:87150ms step_avg:59.12ms step:1475/1600 train_time:87238ms step_avg:59.14ms step:1476/1600 train_time:87323ms step_avg:59.16ms step:1477/1600 train_time:87412ms step_avg:59.18ms step:1478/1600 train_time:87497ms step_avg:59.20ms step:1479/1600 train_time:87585ms step_avg:59.22ms step:1480/1600 train_time:87671ms step_avg:59.24ms step:1481/1600 train_time:87760ms step_avg:59.26ms step:1482/1600 train_time:87845ms step_avg:59.27ms step:1483/1600 train_time:87935ms step_avg:59.30ms step:1484/1600 train_time:88020ms step_avg:59.31ms step:1485/1600 train_time:88108ms step_avg:59.33ms step:1486/1600 train_time:88193ms step_avg:59.35ms step:1487/1600 train_time:88282ms step_avg:59.37ms step:1488/1600 train_time:88367ms step_avg:59.39ms step:1489/1600 train_time:88455ms step_avg:59.41ms step:1490/1600 train_time:88540ms step_avg:59.42ms step:1491/1600 train_time:88628ms step_avg:59.44ms step:1492/1600 train_time:88714ms step_avg:59.46ms step:1493/1600 train_time:88804ms step_avg:59.48ms step:1494/1600 train_time:88888ms step_avg:59.50ms step:1495/1600 train_time:88976ms step_avg:59.52ms step:1496/1600 train_time:89062ms step_avg:59.53ms step:1497/1600 train_time:89150ms step_avg:59.55ms step:1498/1600 train_time:89235ms step_avg:59.57ms step:1499/1600 train_time:89324ms step_avg:59.59ms step:1500/1600 train_time:89409ms step_avg:59.61ms step:1500/1600 val_loss:3.3067 train_time:89482ms step_avg:59.65ms step:1501/1600 train_time:89503ms step_avg:59.63ms step:1502/1600 train_time:89585ms step_avg:59.64ms step:1503/1600 train_time:89678ms step_avg:59.67ms step:1504/1600 train_time:89763ms step_avg:59.68ms step:1505/1600 train_time:89850ms step_avg:59.70ms step:1506/1600 train_time:89934ms step_avg:59.72ms step:1507/1600 train_time:90022ms step_avg:59.74ms step:1508/1600 train_time:90108ms step_avg:59.75ms step:1509/1600 train_time:90195ms step_avg:59.77ms step:1510/1600 train_time:90280ms step_avg:59.79ms step:1511/1600 train_time:90368ms step_avg:59.81ms step:1512/1600 train_time:90455ms step_avg:59.82ms step:1513/1600 train_time:90547ms step_avg:59.85ms step:1514/1600 train_time:90633ms step_avg:59.86ms step:1515/1600 train_time:90723ms step_avg:59.88ms step:1516/1600 train_time:90809ms step_avg:59.90ms step:1517/1600 train_time:90897ms step_avg:59.92ms step:1518/1600 train_time:90981ms step_avg:59.93ms step:1519/1600 train_time:91068ms step_avg:59.95ms step:1520/1600 train_time:91152ms step_avg:59.97ms step:1521/1600 train_time:91240ms step_avg:59.99ms step:1522/1600 train_time:91324ms step_avg:60.00ms step:1523/1600 train_time:91412ms step_avg:60.02ms step:1524/1600 train_time:91499ms step_avg:60.04ms step:1525/1600 train_time:91588ms step_avg:60.06ms step:1526/1600 train_time:91674ms step_avg:60.07ms step:1527/1600 train_time:91763ms step_avg:60.09ms step:1528/1600 train_time:91848ms step_avg:60.11ms step:1529/1600 train_time:91936ms step_avg:60.13ms step:1530/1600 train_time:92020ms step_avg:60.14ms step:1531/1600 train_time:92107ms step_avg:60.16ms step:1532/1600 train_time:92192ms step_avg:60.18ms step:1533/1600 train_time:92280ms step_avg:60.20ms step:1534/1600 train_time:92364ms step_avg:60.21ms step:1535/1600 train_time:92453ms step_avg:60.23ms step:1536/1600 train_time:92539ms step_avg:60.25ms step:1537/1600 train_time:92628ms step_avg:60.27ms step:1538/1600 train_time:92714ms step_avg:60.28ms step:1539/1600 train_time:92803ms step_avg:60.30ms step:1540/1600 train_time:92889ms step_avg:60.32ms step:1541/1600 train_time:92976ms step_avg:60.33ms step:1542/1600 train_time:93060ms step_avg:60.35ms step:1543/1600 train_time:93148ms step_avg:60.37ms step:1544/1600 train_time:93232ms step_avg:60.38ms step:1545/1600 train_time:93320ms step_avg:60.40ms step:1546/1600 train_time:93405ms step_avg:60.42ms step:1547/1600 train_time:93493ms step_avg:60.44ms step:1548/1600 train_time:93579ms step_avg:60.45ms step:1549/1600 train_time:93669ms step_avg:60.47ms step:1550/1600 train_time:93754ms step_avg:60.49ms step:1551/1600 train_time:93844ms step_avg:60.51ms step:1552/1600 train_time:93929ms step_avg:60.52ms step:1553/1600 train_time:94017ms step_avg:60.54ms step:1554/1600 train_time:94102ms step_avg:60.55ms step:1555/1600 train_time:94191ms step_avg:60.57ms step:1556/1600 train_time:94275ms step_avg:60.59ms step:1557/1600 train_time:94366ms step_avg:60.61ms step:1558/1600 train_time:94453ms step_avg:60.62ms step:1559/1600 train_time:94539ms step_avg:60.64ms step:1560/1600 train_time:94625ms step_avg:60.66ms step:1561/1600 train_time:94721ms step_avg:60.68ms step:1562/1600 train_time:94805ms step_avg:60.69ms step:1563/1600 train_time:94893ms step_avg:60.71ms step:1564/1600 train_time:94978ms step_avg:60.73ms step:1565/1600 train_time:95066ms step_avg:60.75ms step:1566/1600 train_time:95152ms step_avg:60.76ms step:1567/1600 train_time:95240ms step_avg:60.78ms step:1568/1600 train_time:95325ms step_avg:60.79ms step:1569/1600 train_time:95415ms step_avg:60.81ms step:1570/1600 train_time:95501ms step_avg:60.83ms step:1571/1600 train_time:95590ms step_avg:60.85ms step:1572/1600 train_time:95675ms step_avg:60.86ms step:1573/1600 train_time:95765ms step_avg:60.88ms step:1574/1600 train_time:95851ms step_avg:60.90ms step:1575/1600 train_time:95939ms step_avg:60.91ms step:1576/1600 train_time:96026ms step_avg:60.93ms step:1577/1600 train_time:96117ms step_avg:60.95ms step:1578/1600 train_time:96201ms step_avg:60.96ms step:1579/1600 train_time:96289ms step_avg:60.98ms step:1580/1600 train_time:96374ms step_avg:61.00ms step:1581/1600 train_time:96462ms step_avg:61.01ms step:1582/1600 train_time:96548ms step_avg:61.03ms step:1583/1600 train_time:96638ms step_avg:61.05ms step:1584/1600 train_time:96723ms step_avg:61.06ms step:1585/1600 train_time:96812ms step_avg:61.08ms step:1586/1600 train_time:96898ms step_avg:61.10ms step:1587/1600 train_time:96986ms step_avg:61.11ms step:1588/1600 train_time:97072ms step_avg:61.13ms step:1589/1600 train_time:97160ms step_avg:61.15ms step:1590/1600 train_time:97246ms step_avg:61.16ms step:1591/1600 train_time:97335ms step_avg:61.18ms step:1592/1600 train_time:97420ms step_avg:61.19ms step:1593/1600 train_time:97509ms step_avg:61.21ms step:1594/1600 train_time:97595ms step_avg:61.23ms step:1595/1600 train_time:97684ms step_avg:61.24ms step:1596/1600 train_time:97770ms step_avg:61.26ms step:1597/1600 train_time:97860ms step_avg:61.28ms step:1598/1600 train_time:97945ms step_avg:61.29ms step:1599/1600 train_time:98034ms step_avg:61.31ms step:1600/1600 train_time:98120ms step_avg:61.32ms step:1600/1600 val_loss:3.2772 train_time:98193ms step_avg:61.37ms peak memory allocated: 30264 MiB reserved: 46240 MiB