import os import sys # Read the current file and the kernels file code ASAP, for logging with open(sys.argv[0], 'r') as f: code = f.read() with open(os.path.join(os.path.dirname(sys.argv[0]), 'triton_kernels.py'), 'r') as f: code += f"\n\n{'-'*40}\n# triton_kernels.py\n{'-'*40}\n\n" code += f.read() import copy import glob import math import threading import time import uuid from dataclasses import dataclass from itertools import accumulate from pathlib import Path import gc os.environ["PYTORCH_ALLOC_CONF"] = "expandable_segments:True" import torch import triton torch.empty( 1, device=f"cuda:{os.environ['LOCAL_RANK']}", requires_grad=True ).backward() # prevents a bug on some systems import torch._dynamo as dynamo import torch.distributed as dist import torch.nn.functional as F # torch._inductor.config.coordinate_descent_tuning = True # we have banned this flag for new records because it causes compilation to take 30min from kernels import get_kernel from torch import Tensor, nn from triton_kernels import XXT, ba_plus_cAA, FusedLinearReLUSquareFunction, FusedSoftcappedCrossEntropy dynamo.config.recompile_limit = 64 # ----------------------------------------------------------------------------- # Custom operators: FP8 matmul by @YouJiacheng # Transposed layout by @ChrisJMcCormick allows for faster gradient accumulation. @torch.library.custom_op("nanogpt::mm_t", mutates_args=()) def mm_t_op(x: Tensor, w: Tensor, x_s: float, w_s: float, grad_s: float) -> tuple[Tensor, Tensor, Tensor]: """Computes y = x @ w with F8 weights stored as (in_features, out_features).""" @torch.compile def impl(x: Tensor, w: Tensor): assert x.is_contiguous() and w.is_contiguous() assert x.shape[1] == w.shape[0] # x: (batch, in), w: (in, out) x_f8 = x.div(x_s).to(torch.float8_e4m3fn) w_f8 = w.div(w_s).to(torch.float8_e4m3fn) # _scaled_mm requires column-major B. w_f8 is row-major (in, out). # .T.contiguous().T creates a column-major view without changing logical shape. w_f8_col_major = w_f8.T.contiguous().T out = torch._scaled_mm( x_f8, w_f8_col_major, out_dtype=torch.bfloat16, scale_a=x.new_tensor(x_s, dtype=torch.float32), scale_b=x.new_tensor(w_s, dtype=torch.float32), use_fast_accum=True, ) return out, x_f8, w_f8 return impl(x, w) @mm_t_op.register_fake def _(x: Tensor, w: Tensor, *_): assert x.ndim == w.ndim == 2 assert x.shape[1] == w.shape[0] assert x.device == w.device assert x.is_contiguous() and w.is_contiguous() return x @ w, x.to(torch.float8_e4m3fn), w.to(torch.float8_e4m3fn) @torch.library.custom_op("nanogpt::mm_t_backward", mutates_args=()) def mm_t_backward_op(g: Tensor, x_f8: Tensor, w_f8: Tensor, x_s: float, w_s: float, grad_s: float) -> tuple[Tensor, Tensor]: @torch.compile def impl(grad: Tensor, x_f8: Tensor, w_f8: Tensor): assert grad.is_contiguous() x_scale = grad.new_tensor(x_s, dtype=torch.float32) w_scale = grad.new_tensor(w_s, dtype=torch.float32) grad_scale = grad.new_tensor(grad_s, dtype=torch.float32) grad_f8 = grad.div(grad_s).to(torch.float8_e5m2) # grad_x = grad @ w.T grad_x = torch._scaled_mm( grad_f8, w_f8.T, out_dtype=torch.bfloat16, scale_a=grad_scale, scale_b=w_scale, use_fast_accum=False, ) # grad_w = x.T @ grad # Result is (in, out), naturally matching weight storage. No final .T needed. grad_w = torch._scaled_mm( x_f8.T.contiguous(), grad_f8.T.contiguous().T, out_dtype=torch.float32, scale_a=x_scale, scale_b=grad_scale, use_fast_accum=False, ) return grad_x, grad_w grad_x, grad_w = impl(g, x_f8, w_f8) return grad_x, grad_w @mm_t_backward_op.register_fake def _(g: Tensor, x_f8: Tensor, w_f8: Tensor, *_): return x_f8.to(torch.bfloat16), w_f8.to(torch.float32) def backward_t(ctx, grad_out: Tensor, *_): x_f8, w_f8 = ctx.saved_tensors x_s, w_s, grad_s = ctx.scales grad_x, grad_w = torch.ops.nanogpt.mm_t_backward( grad_out, x_f8, w_f8, x_s, w_s, grad_s ) return grad_x, grad_w, None, None, None def setup_context_t(ctx: torch.autograd.function.FunctionCtx, inputs, output): *_, x_s, w_s, grad_s = inputs _, x_f8, w_f8 = output ctx.save_for_backward(x_f8, w_f8) ctx.scales = x_s, w_s, grad_s ctx.set_materialize_grads(False) mm_t_op.register_autograd(backward_t, setup_context=setup_context_t) # ----------------------------------------------------------------------------- # Polar Express # Computed for num_iters=5, safety_factor=2e-2, cushion=2 polar_express_coeffs = [ (8.156554524902461, -22.48329292557795, 15.878769915207462), (4.042929935166739, -2.808917465908714, 0.5000178451051316), (3.8916678022926607, -2.772484153217685, 0.5060648178503393), (3.285753657755655, -2.3681294933425376, 0.46449024233003106), (2.3465413258596377, -1.7097828382687081, 0.42323551169305323) ] @torch.compile(dynamic=False, fullgraph=True) # Must use dynamic=False or else it's much slower def polar_express(G: torch.Tensor, split_baddbmm: bool = False): """ Polar Express Sign Method: https://arxiv.org/pdf/2505.16932 by Noah Amsel, David Persson, Christopher Musco, Robert M. Gower. """ X = G.bfloat16() if G.size(-2) > G.size(-1): X = X.mT # Ensure spectral norm is at most 1 X = X / (X.norm(dim=(-2, -1), keepdim=True) * (1 + 2e-2) + 1e-6) # Allocate buffers X = X.contiguous() A = torch.empty((*X.shape[:-1], X.size(-2)), device=X.device, dtype=X.dtype) B = torch.empty_like(A) C = torch.empty_like(X) # Select batched vs unbatched if split_baddbmm: BX_matmul = torch.bmm if X.ndim > 2 else torch.mm else: aX_plus_BX = torch.baddbmm if X.ndim > 2 else torch.addmm # Perform the iterations for a, b, c in polar_express_coeffs: XXT(X, out=A) # A = X @ X.mT ba_plus_cAA(A, alpha=c, beta=b, out=B) # B = b * A + c * A @ A # Referencing X twice causes pytorch to make a defensive copy, # resulting in a cudaMemcpyAsync in baddbmm. # For large matrices (i.e., the mlp weights), it's faster to split # the operation into two kernels to avoid this. if split_baddbmm: BX_matmul(B, X, out=C) # C = B @ X C.add_(X, alpha=a) # C = C + a*X (in-place, X only read) else: aX_plus_BX(X, B, X, beta=a, out=C) # C = a * X + B @ X X, C = C, X # Swap references to avoid unnecessary copies if G.size(-2) > G.size(-1): X = X.mT return X # ----------------------------------------------------------------------------- # Combined NorMuon + Adam Optimizer @dataclass class ParamConfig: """Per-parameter configuration for NorMuonAndAdam optimizer.""" label: str optim: str # "adam" or "normuon" comms: str # "none", "replicated", or "sharded" adam_betas: tuple[float, float] | None lr_mul: float wd_mul: float lr: float initial_lr: float weight_decay: float # Adam-specific eps: float | None = None # NorMuon-specific reshape: tuple | None = None chunk_size: int | None = None momentum: float | None = None beta2: float | None = None per_matrix_lr_mul: list[float] | None = None class NorMuonAndAdam: """ Combined optimizer that handles both NorMuon (for projection matrices) and Adam (for embeddings/scalars/gate weights). Muon - MomentUm Orthogonalized by Newton-schulz https://kellerjordan.github.io/posts/muon/ Muon internally runs standard SGD-momentum, and then performs an orthogonalization post- processing step, in which each 2D parameter's update is replaced with the nearest orthogonal matrix. To efficiently orthogonalize each update, Muon uses a Newton-Schulz iteration (replaced here with Polar Express), which has the advantage that it can be stably run in bfloat16 on the GPU. Muon is applied only to the projection matrices in the attention and MLP layers, and is not recommended for embeddings, scalars, or individual weight vectors (e.g., bias terms or gate weights). Differences from standard Muon: - Newton-Shulz is replaced with Polar Express for the orthogonalization step - NorMuon adds a low-rank variance estimator similar to Adafactor. https://arxiv.org/pdf/2510.05491 - Cautious weight decay, a gated version of decoupled weight decay - Mantissa tracking for precision Adam (for embeddings/scalars/gates): - Standard Adam with bias correction - Cautious weight decay Configuration: Unlike torch.optim.Optimizer, this class uses per-parameter configs from a `param_table` dict and does not include parameter "groups". All parameters require a .label attribute, and a corresponding entry in the param_table to specify their hyperparameters (lr_mul, wd_mul, adam_betas, etc.). Communication and ordering: Gradient communication is explicitly scheduled rather than hook-driven. Reductions are launched in `scatter_order`, while update math and final gathers are executed in `work_order`. These orders are independent and must each contain every parameter label exactly once. Two communication modes are supported per parameter: - 'replicated': Gradients are all-reduced and each rank computes the full update. - 'sharded': Gradients are reduce-scattered, each rank updates its shard, and results are all-gathered. Adam parameters may be freely sharded. NorMuon operates on full matrices; sharding is supported by grouping matrices into parameter banks. NorMuon parameters must have a `.reshape` attribute that reshapes the bank so that the leading dimension is divisible by world_size. # Contributors include @YouJiacheng, @KonstantinWilleke, @alexrgilbert, @adricarda, # @tuttyfrutyee, @vdlad, @ryanyang0, @vagrawal, @varunneal, @chrisjmccormick """ def __init__(self, named_params, param_table: dict, scatter_order: list, work_order: list, adam_defaults: dict, normuon_defaults: dict): self.world_size = dist.get_world_size() if dist.is_initialized() else 1 # Store defaults for each optimizer type self.adam_defaults = adam_defaults self.normuon_defaults = normuon_defaults self.param_table = param_table self.scatter_order = scatter_order self.work_order = work_order # Collect params by label and build config self.param_cfgs: dict[nn.Parameter, ParamConfig] = {} self.param_states: dict[nn.Parameter, dict] = {} self._param_by_label: dict[str, nn.Parameter] = {} for name, param in named_params: label = getattr(param, "label", None) assert label is not None and label in param_table # all params must have valid label assert label not in self._param_by_label # exactly one param per label self._param_by_label[label] = param self._build_param_cfg(param, label) # Assert scatter_order and work_order match present labels exactly present = set(self._param_by_label.keys()) assert set(scatter_order) == present and set(work_order) == present # Handle world_size=1: overwrite comms to "none" if self.world_size == 1: for p_cfg in self.param_cfgs.values(): p_cfg.comms = "none" # Initialize state for all params self._init_state() # 0-D CPU tensors to avoid recompilation self._step_size_t = torch.tensor(0.0, dtype=torch.float32, device="cpu") self._eff_wd_t = torch.tensor(0.0, dtype=torch.float32, device="cpu") self._eff_lr_t = torch.tensor(0.0, dtype=torch.float32, device="cpu") # Track async operations self._reduce_futures: dict[nn.Parameter, tuple] = {} # Embed/lm_head tying state self.split_embed = False self._lm_head_param = self._param_by_label.get("lm_head") self._embed_param = self._param_by_label.get("embed") def _build_param_cfg(self, param: nn.Parameter, label: str): """Build config for a single parameter from param_table.""" table_entry = self.param_table[label] optim = table_entry["optim"] comms = table_entry["comms"] adam_betas = table_entry.get("adam_betas") lr_mul = table_entry.get("lr_mul", 1.0) wd_mul = table_entry.get("wd_mul", 1.0) if optim == "adam": chunk_size = param.shape[0] // self.world_size if comms == "sharded" else None p_cfg = ParamConfig( label=label, optim=optim, comms=comms, adam_betas=tuple(adam_betas) if adam_betas else None, lr_mul=lr_mul, wd_mul=wd_mul, lr=self.adam_defaults["lr"], initial_lr=self.adam_defaults["lr"], weight_decay=self.adam_defaults["weight_decay"], eps=self.adam_defaults["eps"], chunk_size=chunk_size, ) elif optim == "normuon": reshape = getattr(param, "reshape", None) if reshape is None: raise ValueError(f"NorMuon param {label} must have .reshape attribute") if reshape[0] % self.world_size != 0: raise ValueError(f"reshape[0]={reshape[0]} must be divisible by world_size") chunk_size = reshape[0] // self.world_size chunk_shape = (chunk_size, *reshape[1:]) # Shape-based LR multiplier for NorMuon shape_mult = max(1.0, chunk_shape[-2] / chunk_shape[-1]) ** 0.5 if len(chunk_shape) >= 2 else 1.0 lr_mul = shape_mult * lr_mul # Per-matrix LR multipliers for MLP c_proj (2x LR on odd indices) per_matrix_lr_mul = None if label == "mlp": rank = dist.get_rank() if dist.is_initialized() else 0 start_idx = rank * chunk_size per_matrix_lr_mul = [] for i in range(chunk_size): global_idx = start_idx + i is_c_proj = (global_idx % 2 == 1) per_matrix_lr_mul.append(2.0 if is_c_proj else 1.0) p_cfg = ParamConfig( label=label, optim=optim, comms=comms, adam_betas=tuple(adam_betas) if adam_betas else None, lr_mul=lr_mul, wd_mul=wd_mul, lr=self.normuon_defaults["lr"], initial_lr=self.normuon_defaults["lr"], weight_decay=self.normuon_defaults["weight_decay"], reshape=reshape, chunk_size=chunk_size, momentum=self.normuon_defaults["momentum"], beta2=self.normuon_defaults["beta2"], per_matrix_lr_mul=per_matrix_lr_mul, ) else: raise ValueError(f"Unknown optim type: {optim}") self.param_cfgs[param] = p_cfg def _init_state(self): """Initialize optimizer state for all parameters.""" for param, p_cfg in self.param_cfgs.items(): if p_cfg.optim == "adam": # Sharded params use chunk state, replicated use full state if p_cfg.comms == "sharded": chunk = param[:p_cfg.chunk_size] else: chunk = param exp_avg = torch.zeros_like(chunk, dtype=torch.float32, device=param.device) self.param_states[param] = dict(step=0, exp_avg=exp_avg, exp_avg_sq=torch.zeros_like(exp_avg)) elif p_cfg.optim == "normuon": chunk_shape = (p_cfg.chunk_size, *p_cfg.reshape[1:]) # Momentum buffer (FP32 for precision) momentum_buffer = torch.zeros( chunk_shape, dtype=torch.float32, device=param.device ) # Second momentum buffer - reduced along one dimension if chunk_shape[-2] >= chunk_shape[-1]: second_mom_shape = (*chunk_shape[:-1], 1) else: second_mom_shape = (*chunk_shape[:-2], 1, chunk_shape[-1]) second_momentum_buffer = torch.zeros( second_mom_shape, dtype=torch.float32, device=param.device ) # Mantissa buffer for precision tracking mantissa = torch.zeros( chunk_shape, dtype=torch.uint16, device=param.device ) self.param_states[param] = dict( momentum_buffer=momentum_buffer, second_momentum_buffer=second_momentum_buffer, mantissa=mantissa, ) # ----------------------------------- # Reduce/Gather operations def _launch_reduce(self, param: nn.Parameter, grad: Tensor): """Launch async reduce for a parameter based on its comms policy.""" p_cfg = self.param_cfgs[param] if p_cfg.comms == "none": if p_cfg.optim == "normuon": # NorMuon needs reshaped gradient even without communication grad = grad.view(p_cfg.reshape) self._reduce_futures[param] = (None, grad) elif p_cfg.comms == "replicated": future = dist.all_reduce(grad, op=dist.ReduceOp.AVG, async_op=True).get_future() self._reduce_futures[param] = (future, grad) elif p_cfg.comms == "sharded": if p_cfg.optim == "normuon": # NorMuon: reshape before reduce_scatter grad_reshaped = grad.view(p_cfg.reshape) grad_chunk = torch.empty( (p_cfg.chunk_size, *grad_reshaped.shape[1:]), dtype=grad.dtype, device=grad.device ) future = dist.reduce_scatter_tensor( grad_chunk, grad_reshaped.contiguous(), op=dist.ReduceOp.AVG, async_op=True ).get_future() self._reduce_futures[param] = (future, grad_chunk) else: # Adam: simple reduce_scatter grad_chunk = torch.empty_like(grad[:p_cfg.chunk_size]) future = dist.reduce_scatter_tensor( grad_chunk, grad, op=dist.ReduceOp.AVG, async_op=True ).get_future() self._reduce_futures[param] = (future, grad_chunk) def _launch_gather(self, param: nn.Parameter, p_slice: Tensor) -> "torch.futures.Future": """Launch async all_gather for a sharded parameter.""" p_cfg = self.param_cfgs[param] if p_cfg.optim == "normuon": full_param = param.data.view(p_cfg.reshape) assert full_param.is_contiguous() return dist.all_gather_into_tensor( full_param, p_slice.contiguous(), async_op=True ).get_future() else: return dist.all_gather_into_tensor( param, p_slice.contiguous(), async_op=True ).get_future() # ----------------------------------- # State management def reset(self): """Reset NorMuon momentum buffers and split_embed state (called on training reset).""" self.split_embed = False for param, p_cfg in self.param_cfgs.items(): if p_cfg.optim == "normuon": p_state = self.param_states[param] p_state["momentum_buffer"].zero_() p_state["mantissa"].zero_() p_state["second_momentum_buffer"].zero_() def copy_lm_state_to_embed(self): """ Copy the optimizer state from the lm_head to the embed at the untie point. This requires an all-gather + reshard because of different sharding: - lm_head (768, 50304) is sharded to (96, 50304) per rank (along model_dim) - embed (50304, 768) is sharded to (6288, 768) per rank (along vocab_size) We all-gather the lm_head momentum, transpose it, then each rank takes their embed shard to get the correct momentum state. """ lm_head = self._lm_head_param embed = self._embed_param lm_state = self.param_states[lm_head] embed_state = self.param_states[embed] lm_cfg = self.param_cfgs[lm_head] embed_cfg = self.param_cfgs[embed] embed_state['step'] = lm_state['step'] # Preserve step count for bias correction # Copy optimizer state with all-gather + transpose + reshard if self.world_size > 1: rank = dist.get_rank() lm_chunk_size = lm_cfg.chunk_size # 96 embed_chunk_size = embed_cfg.chunk_size # 6288 # All-gather lm_head momentum to get full (768, 50304) tensor for key in ["exp_avg", "exp_avg_sq"]: lm_chunk = lm_state[key] # (96, 50304) full_lm = torch.empty(lm_head.shape[0], lm_head.shape[1], dtype=lm_chunk.dtype, device=lm_chunk.device) dist.all_gather_into_tensor(full_lm, lm_chunk.contiguous()) embed_state[key].copy_(full_lm.T[rank * embed_chunk_size:(rank + 1) * embed_chunk_size]) else: # Single GPU: simple transpose for key in ["exp_avg", "exp_avg_sq"]: embed_state[key].copy_(lm_state[key].T) # Mark as split self.split_embed = True def state_dict(self): """Return the optimizer state as a dict.""" return { "param_states": {id(p): s for p, s in self.param_states.items()}, "param_cfgs": {id(p): s for p, s in self.param_cfgs.items()}, } def load_state_dict(self, state_dict): """Load optimizer state from a dict.""" # Build id->param mapping id_to_param = {id(p): p for p in self.param_cfgs.keys()} # Load state, preserving dtypes for param_id, saved_p_state in state_dict["param_states"].items(): if param_id in id_to_param: param = id_to_param[param_id] p_state = self.param_states[param] for k, v in saved_p_state.items(): if isinstance(v, torch.Tensor) and k in p_state: target_dtype = p_state[k].dtype p_state[k] = v.to(dtype=target_dtype, device=p_state[k].device) else: p_state[k] = v # ----------------------------------- # Unified optimizer step with explicit ordering @torch.no_grad() def step(self, do_adam: bool = True): """ Combined optimizer step with explicit ordering. Args: do_adam: If True, update Adam params. NorMuon params always updated. Flow: 1. Scatter phase: Launch reduces in scatter_order 2. Work phase: Process updates in work_order - Wait for reduce, compute update, launch gather 3. Finalize phase: Wait for gathers While the embeddings are tied: - Comms and update math are only done on lm_head. - We add embed.grad.T into lm_head.grad before comms. - After lm_head gather, we copy lm_head.data.T --> embed.data """ rank = dist.get_rank() if dist.is_initialized() else 0 lm_param, embed_param = self._lm_head_param, self._embed_param # ===== Phase 1: Launch reduces in scatter_order ===== for label in self.scatter_order: param = self._param_by_label[label] p_cfg = self.param_cfgs[param] if p_cfg.optim == "adam" and not do_adam: continue if param.grad is None: continue # lm_head when tied: aggregate embed.grad.T (transposed shapes) if label == "lm_head" and do_adam and not self.split_embed: if embed_param is not None and embed_param.grad is not None: param.grad.add_(embed_param.grad.T) # Skip embed when tied (copied from lm_head after gather) if label == "embed" and not self.split_embed: continue self._launch_reduce(param, param.grad) # ===== Phase 2: Process updates in work_order ===== gather_futures = [] lm_head_gather_future = None for label in self.work_order: param = self._param_by_label[label] if param not in self._reduce_futures: continue p_cfg = self.param_cfgs[param] if p_cfg.optim == "adam" and not do_adam: continue # Wait for reduce future, grad_chunk = self._reduce_futures[param] if future is not None: future.wait() # Apply update based on optim type if p_cfg.optim == "adam": p_slice = self._adam_update(param, grad_chunk, p_cfg, rank) else: p_slice = self._normuon_update(param, grad_chunk, p_cfg, rank) # Launch gather for sharded params if p_cfg.comms == "sharded" and self.world_size > 1: gather_fut = self._launch_gather(param, p_slice) if label == "lm_head": lm_head_gather_future = gather_fut else: gather_futures.append(gather_fut) # ===== Phase 3: Wait for gathers, sync embed if tied ===== # Wait for lm_head gather first so we can copy to embed while other gathers complete if lm_head_gather_future is not None: lm_head_gather_future.wait() # When tied: copy lm_head.T to embed if do_adam and not self.split_embed and embed_param is not None and lm_param is not None: embed_param.data.copy_(lm_param.data.T) # Wait for remaining gathers for fut in gather_futures: fut.wait() self._reduce_futures.clear() # Clear grads for updated params for param, p_cfg in self.param_cfgs.items(): if p_cfg.optim == "adam" and not do_adam: continue # Don't clear Adam grads on even steps param.grad = None # ----------------------------------- # Adam update def _adam_update(self, param: nn.Parameter, grad_chunk: Tensor, p_cfg: ParamConfig, rank: int) -> Tensor: """Apply Adam update to a parameter. Returns the updated p_slice.""" beta1, beta2 = p_cfg.adam_betas lr = p_cfg.lr * p_cfg.lr_mul # Get parameter slice if p_cfg.comms == "sharded": p_slice = param[rank * p_cfg.chunk_size:(rank + 1) * p_cfg.chunk_size] else: p_slice = param p_state = self.param_states[param] p_state["step"] += 1 t = p_state["step"] bias1, bias2 = 1 - beta1 ** t, 1 - beta2 ** t self._step_size_t.fill_(lr * (bias2 ** 0.5 / bias1)) self._eff_wd_t.fill_(lr * lr * p_cfg.weight_decay * p_cfg.wd_mul) NorMuonAndAdam._adam_update_step( p_slice, grad_chunk, p_state["exp_avg"], p_state["exp_avg_sq"], beta1, beta2, p_cfg.eps, self._step_size_t, self._eff_wd_t ) return p_slice @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _adam_update_step(p_slice, g_slice, exp_avg, exp_avg_sq, beta1, beta2, eps, step_size_t, eff_wd_t): """Compiled Adam update step.""" exp_avg.mul_(beta1).add_(g_slice, alpha=1 - beta1) exp_avg_sq.mul_(beta2).addcmul_(g_slice, g_slice, value=1 - beta2) update = exp_avg.div(exp_avg_sq.sqrt().add_(eps)).mul_(step_size_t) # Cautious weight decay mask = (update * p_slice) > 0 update.addcmul_(p_slice, mask, value=eff_wd_t) p_slice.add_(other=update, alpha=-1.0) # ----------------------------------- # NorMuon update def _normuon_update(self, param: nn.Parameter, grad_chunk: Tensor, p_cfg: ParamConfig, rank: int) -> Tensor: """Apply NorMuon update to a parameter. Returns the updated p_slice.""" chunk_shape = grad_chunk.shape p_state = self.param_states[param] grad_chunk = grad_chunk.float() # FP32 for momentum # Momentum update momentum_buffer = p_state["momentum_buffer"] momentum_buffer.lerp_(grad_chunk, 1 - p_cfg.momentum) updated_grads = grad_chunk.lerp_(momentum_buffer, p_cfg.momentum) self._eff_lr_t.fill_(p_cfg.lr_mul * p_cfg.lr) self._eff_wd_t.fill_(p_cfg.wd_mul * p_cfg.weight_decay * p_cfg.lr) # Polar Express orthogonalization is_large_matrix = chunk_shape[-2] > 1024 v_chunk = polar_express(updated_grads, split_baddbmm=is_large_matrix) # Variance reduction red_dim = -1 if chunk_shape[-2] >= chunk_shape[-1] else -2 v_chunk = NorMuonAndAdam._apply_normuon_variance_reduction( v_chunk, p_state["second_momentum_buffer"], p_cfg.beta2, red_dim ) # Update parameter, in place, with cautious weight decay param_view = param.data.view(p_cfg.reshape) p_slice = param_view[rank * p_cfg.chunk_size:(rank + 1) * p_cfg.chunk_size] # MLP has per-matrix LR multipliers (c_proj gets 2x LR) if p_cfg.per_matrix_lr_mul is not None: for mat_idx in range(p_cfg.chunk_size): self._eff_lr_t.fill_(p_cfg.lr_mul * p_cfg.per_matrix_lr_mul[mat_idx] * p_cfg.lr) self._eff_wd_t.fill_(p_cfg.wd_mul * p_cfg.weight_decay * p_cfg.lr) NorMuonAndAdam._cautious_wd_and_update_inplace( p_slice[mat_idx].view(torch.uint16), p_state["mantissa"][mat_idx], v_chunk[mat_idx], self._eff_wd_t, self._eff_lr_t ) else: NorMuonAndAdam._cautious_wd_and_update_inplace( p_slice.view(torch.uint16), p_state["mantissa"], v_chunk, self._eff_wd_t, self._eff_lr_t ) return p_slice @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _cautious_wd_and_update_inplace(p, mantissa, grad, wd_tensor, lr_tensor): """ Cautious weight decay + parameter update. wd_tensor and lr_tensor are 0-D CPU tensors. Mantissa is tracked to enable higher precision updates on bfloat16 parameters. bfloat16 format: 1 sign bit + 8 exponent bits + 7 mantissa bits = 16 bits total float32 format: 1 sign bit + 8 exponent bits + 23 mantissa bits = 32 bits total """ assert p.dtype == mantissa.dtype == torch.uint16 grad = grad.float() wd_factor = wd_tensor.to(torch.float32) lr_factor = lr_tensor.to(torch.float32) p_precise_raw = (p.to(torch.uint32) << 16) | mantissa.to(torch.uint32) p_precise = p_precise_raw.view(torch.float32) mask = (grad * p_precise) >= 0 p_precise.copy_(p_precise - (p_precise * mask * wd_factor * lr_factor) - (grad * lr_factor)) p.copy_((p_precise_raw >> 16).to(torch.uint16)) mantissa.copy_(p_precise_raw.to(torch.uint16)) @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _apply_normuon_variance_reduction(v_chunk, second_momentum_buffer, beta2, red_dim): """NorMuon variance reduction. Algebraically fuses the normalization steps to minimize memory ops.""" v_mean = v_chunk.float().square().mean(dim=red_dim, keepdim=True) red_dim_size = v_chunk.size(red_dim) v_norm_sq = v_mean.sum(dim=(-2, -1), keepdim=True).mul_(red_dim_size) v_norm = v_norm_sq.sqrt_() second_momentum_buffer.lerp_(v_mean.to(dtype=second_momentum_buffer.dtype), 1 - beta2) step_size = second_momentum_buffer.clamp_min(1e-10).rsqrt_() scaled_sq_sum = (v_mean * red_dim_size) * step_size.float().square() v_norm_new = scaled_sq_sum.sum(dim=(-2, -1), keepdim=True).sqrt_() final_scale = step_size * (v_norm / v_norm_new.clamp_min_(1e-10)) return v_chunk.mul_(final_scale.type_as(v_chunk)) # ----------------------------------------------------------------------------- # PyTorch nn.Module definitions for the model def norm(x: Tensor): return F.rms_norm(x, (x.size(-1),)) class CastedLinearT(nn.Module): """ Linear layer with transposed weight storage (in_features, out_features) which addresses the slow kernel that was used for gradient accumulation. @chrisjmccormick """ def __init__(self, in_features: int, out_features: int, use_fp8=False, x_s=1.0, w_s=1.0, grad_s=1.0): super().__init__() self.in_features = in_features self.out_features = out_features self.use_fp8 = use_fp8 self.x_s = x_s self.w_s = w_s self.grad_s = grad_s self.weight = nn.Parameter(torch.empty(in_features, out_features, dtype=torch.bfloat16)) self.reset_parameters() def reset_parameters(self) -> None: with torch.no_grad(): nn.init.zeros_(self.weight) # @Grad62304977 and others def forward(self, x: Tensor): if self.use_fp8 and self.training: _x = x.flatten(0, -2) out = torch.ops.nanogpt.mm_t(_x, self.weight, x_s=self.x_s, w_s=self.w_s, grad_s=self.grad_s)[0] return out.reshape(*x.shape[:-1], -1) else: return x @ self.weight.type_as(x) # ----------------------------------------------------------------------------- # PyTorch nn.Module definitions for the model class Yarn(nn.Module): def __init__(self, head_dim, max_seq_len): super().__init__() self.head_dim = head_dim self.max_seq_len = max_seq_len self.reset() def rotary(self, x_BTHD): assert self.factor1.size(0) >= x_BTHD.size(-3) factor1, factor2 = ( self.factor1[None, : x_BTHD.size(-3), None, :], self.factor2[None, : x_BTHD.size(-3), None, :], ) x_flip = x_BTHD.view(*x_BTHD.shape[:-1], x_BTHD.shape[-1] // 2, 2).flip(-1).view(x_BTHD.shape) return factor1 * x_BTHD + factor2 * x_flip def reset(self): angular_freq = (1 / 1024) ** torch.linspace(0, 1, steps=self.head_dim//4, dtype=torch.float32, device=device) angular_freq = angular_freq.repeat_interleave(2) # half-truncate RoPE by @YouJiacheng (w/ base freq tuning) angular_freq = torch.cat([angular_freq, angular_freq.new_zeros(self.head_dim//2)]) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=device) theta = torch.outer(t, angular_freq) self.factor1 = nn.Buffer( theta.cos().to(torch.bfloat16), persistent=False ) self.factor2 = nn.Buffer( theta.sin().to(torch.bfloat16), persistent=False ) self.factor2[..., 1::2] *= -1 self.angular_freq = angular_freq # start with 0.1, inspired by 0.12 from @leloykun and learnable scalars used by @brendanh0gan https://x.com/hi_tysam/status/1879693583898591283 self.attn_scale = 0.1 def apply(self, old_window: int, new_window: int, alpha: int=1, beta: int=32): rotations = args.block_size * old_window * self.angular_freq / (2 * torch.pi) scaling_factor = old_window / new_window interpolation_weight = torch.clamp((rotations - alpha) / (beta - alpha), 0, 1) self.angular_freq *= scaling_factor + interpolation_weight * (1 - scaling_factor) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=self.angular_freq.device) theta = torch.outer(t, self.angular_freq) self.factor1.copy_(theta.cos()) self.factor2.copy_(theta.sin()) self.factor2[..., 1::2] *= -1 self.attn_scale *= 0.2 * math.log(new_window / old_window) + 1 class YarnPairedHead(nn.Module): def __init__(self, head_dim, max_seq_len): super().__init__() self.head_dim = head_dim self.max_seq_len = max_seq_len self.reset() def rotary(self, x_BTHD): assert self.factor1.size(0) >= x_BTHD.size(-3) factor1, factor2 = ( self.factor1[None, : x_BTHD.size(-3), None, :], self.factor2[None, : x_BTHD.size(-3), None, :], ) x_flip = x_BTHD.view(*x_BTHD.shape[:-1], x_BTHD.shape[-1] // 2, 2).flip(-1).view(x_BTHD.shape) return factor1 * x_BTHD + factor2 * x_flip def reset(self): angular_freq = (1 / 1024) ** torch.linspace(0, 1, steps=self.head_dim//4, dtype=torch.float32, device=device) angular_freq = angular_freq.repeat_interleave(2) angular_freq = torch.cat([angular_freq, angular_freq.new_zeros(self.head_dim//2)]) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=device) t_even = 2 * t t_odd = 2 * t + 1 theta1 = torch.outer(t_even, angular_freq) theta2 = torch.outer(t_odd, angular_freq) self.factor1 = nn.Buffer( torch.cat((theta1.cos(),theta2.cos()), dim=-1).to(torch.bfloat16), persistent=False ) self.factor2 = nn.Buffer( torch.cat((theta1.sin(),theta2.sin()), dim=-1).to(torch.bfloat16), persistent=False ) self.factor2[..., 1::2] *= -1 self.angular_freq = angular_freq # start with 0.1, inspired by 0.12 from @leloykun and learnable scalars used by @brendanh0gan https://x.com/hi_tysam/status/1879693583898591283 self.attn_scale = 0.1 def apply(self, old_window: int, new_window: int, alpha: int=1, beta: int=32): rotations = args.block_size * old_window * self.angular_freq / (2 * torch.pi) scaling_factor = old_window / new_window interpolation_weight = torch.clamp((rotations - alpha) / (beta - alpha), 0, 1) self.angular_freq *= scaling_factor + interpolation_weight * (1 - scaling_factor) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=self.angular_freq.device) t_even = 2 * t t_odd = 2 * t + 1 theta1 = torch.outer(t_even, self.angular_freq) theta2 = torch.outer(t_odd, self.angular_freq) self.factor1.copy_(torch.cat((theta1.cos(),theta2.cos()), dim=-1)) self.factor2.copy_( torch.cat((theta1.sin(),theta2.sin()), dim=-1)) self.factor2[..., 1::2] *= -1 self.attn_scale *= 0.2 * math.log(new_window / old_window) + 1 @dataclass class AttnArgs: ve: torch.Tensor sa_lambdas: torch.Tensor seqlens: torch.Tensor bm_size: int yarn: Yarn key_offset: bool attn_gate_w: torch.Tensor ve_gate_w: torch.Tensor flash_attn_interface = get_kernel('varunneal/flash-attention-3').flash_attn_interface class CausalSelfAttention(nn.Module): def __init__(self, dim: int, head_dim: int, num_heads: int): super().__init__() self.num_heads = num_heads self.head_dim = head_dim self.dim = dim self.hdim = num_heads * head_dim assert self.hdim == self.dim, "num_heads * head_dim must equal model_dim" # Weights are stored in parameter banks and passed via forward() def forward(self, x: Tensor, attn_args: AttnArgs, qkvo_w: Tensor): B, T = x.size(0), x.size(1) # batch size, sequence length assert B == 1, "varlen sequences requires B == 1" assert T % 16 == 0 # unpack attention args yarn = attn_args.yarn ve, sa_lambdas, key_offset = attn_args.ve, attn_args.sa_lambdas, attn_args.key_offset seqlens, bm_size = attn_args.seqlens, attn_args.bm_size # sparse gated attention to enable context based no-op by @classiclarryd # only include gates on layers with value embeds used on forward pass attn_gate_w, ve_gate_w = attn_args.attn_gate_w, attn_args.ve_gate_w q, k, v = F.linear(x, sa_lambdas[0] * qkvo_w[:self.dim * 3].type_as(x)).view(B, T, 3 * self.num_heads, self.head_dim).chunk(3, dim=-2) q, k = norm(q), norm(k) # QK norm @Grad62304977 q, k = yarn.rotary(q), yarn.rotary(k) if key_offset: # shift keys forward for the stationary head dims. Enables 1-layer induction. k[:, 1:, :, self.head_dim // 2:] = k[:, :-1, :, self.head_dim // 2:] if ve is not None: ve_gate_out = 2 * torch.sigmoid(F.linear(x[..., :12], ve_gate_w)).view(B, T, self.num_heads, 1) v = v + ve_gate_out * ve.view_as(v) # @ KoszarskyB & @Grad62304977 max_len = args.train_max_seq_len if self.training else (args.val_batch_size // (grad_accum_steps * world_size)) # use flash_attn over flex_attn @varunneal. flash_attn_varlen suggested by @YouJiacheng y = flash_attn_interface.flash_attn_varlen_func(q[0], k[0], v[0], cu_seqlens_q=seqlens, cu_seqlens_k=seqlens, max_seqlen_q=max_len, max_seqlen_k=max_len, causal=True, softmax_scale=yarn.attn_scale, window_size=(bm_size, 0)) y = y.view(B, T, self.num_heads, self.head_dim) y = y * torch.sigmoid(F.linear(x[..., :12], attn_gate_w)).view(B, T, self.num_heads, 1) y = y.contiguous().view(B, T, self.num_heads * self.head_dim) # re-assemble all head outputs side by side y = F.linear(y, sa_lambdas[1] * qkvo_w[self.dim * 3:].type_as(y)) # sa_lambdas[1] pre-multiplied to O @shenberg return y class PairedHeadCausalSelfAttention(nn.Module): """ Pairs up attention heads such that queries from head 1 can attend to keys in head 2, and vice-versa. Implemented by interleaving the k, q, and v for pairs of heads to form twice as long sequences EG [k1_h1, k2_h1, k3_h1], [k1_h2, k2_h2, k3_h2] -> [k1_h1, k1_h2, k2_h1, k2_h2, k3_h1, k3_h2], repeat for q and v """ def __init__(self, dim: int, head_dim: int, num_heads: int): super().__init__() self.num_heads = num_heads self.head_dim = head_dim self.dim = dim self.hdim = num_heads * head_dim assert self.hdim == self.dim, "num_heads * head_dim must equal model_dim" # Weights are stored in parameter banks and passed via forward() def forward(self, x: Tensor, attn_args: AttnArgs, qkvo_w: Tensor): B, T = x.size(0), x.size(1) # batch size, sequence length assert B == 1, "varlen sequences requires B == 1" assert T % 16 == 0 # unpack attention args yarn = attn_args.yarn ve, sa_lambdas = attn_args.ve, attn_args.sa_lambdas seqlens, bm_size = attn_args.seqlens, attn_args.bm_size attn_gate_w, ve_gate_w = attn_args.attn_gate_w, attn_args.ve_gate_w q, k, v = F.linear(x, sa_lambdas[0] * qkvo_w[:self.dim * 3].type_as(x)).view(B, T, 3 * self.num_heads, self.head_dim).chunk(3, dim=-2) q, k = norm(q), norm(k) # delay q,k reshape until rotary makes data contiguous, to enable view (non-copy) q = q.view(B, T, self.num_heads // 2, self.head_dim * 2) k = k.view(B, T, self.num_heads // 2, self.head_dim * 2) v = v.reshape(B, T*2, self.num_heads//2, self.head_dim) q, k = yarn.rotary(q), yarn.rotary(k) q = q.view(B, T*2, self.num_heads//2, self.head_dim) k = k.view(B, T*2, self.num_heads//2, self.head_dim) if ve is not None: ve_gate_out = 2 * torch.sigmoid(F.linear(x[..., :12], ve_gate_w)).view(B, T*2, self.num_heads//2, 1) v = v + ve_gate_out * ve.view_as(v) max_len = args.train_max_seq_len if self.training else (args.val_batch_size // (grad_accum_steps * world_size)) # paired head correction seqlens = 2 * seqlens max_len = 2 * max_len y = flash_attn_interface.flash_attn_varlen_func(q[0], k[0], v[0], cu_seqlens_q=seqlens, cu_seqlens_k=seqlens, max_seqlen_q=max_len, max_seqlen_k=max_len, causal=True, softmax_scale=yarn.attn_scale, window_size=(bm_size, 0)) y = y.view(B, T, self.num_heads, self.head_dim) y = y * torch.sigmoid(F.linear(x[..., :12], attn_gate_w)).view(B, T, self.num_heads, 1) y = y.contiguous().view(B, T, self.num_heads * self.head_dim) y = F.linear(y, sa_lambdas[1] * qkvo_w[self.dim * 3:].type_as(y)) return y class MLP(nn.Module): def __init__(self): super().__init__() # Weights are stored in parameter banks and passed via forward() def forward(self, x: Tensor, c_fc: Tensor, c_proj: Tensor): # relu(x)^2: # https://arxiv.org/abs/2109.08668v2; ~1-2% better than GELU; suggested by @SKYLINEZ007 and @Grad62304977 # Fused triton kernel for relu(x @ W1.T)^2 @ W2.T return FusedLinearReLUSquareFunction.apply(x, c_fc, c_proj) class Block(nn.Module): def __init__(self, dim: int, head_dim: int, num_heads: int, has_attn: bool, has_mlp: bool, use_paired_head: bool): super().__init__() # skip attention of blocks.6 (the 7th layer) by @YouJiacheng if has_attn: if use_paired_head: self.attn = PairedHeadCausalSelfAttention(dim, head_dim, num_heads) else: self.attn = CausalSelfAttention(dim, head_dim, num_heads) else: self.attn = None # skip MLP blocks for first MLP layer by @EmelyanenkoK self.mlp = MLP() if has_mlp else None def forward(self, x: Tensor, attn_args: AttnArgs, qkvo_w: Tensor = None, c_fc: Tensor = None, c_proj: Tensor = None): if self.attn is not None: x = x + self.attn(norm(x), attn_args, qkvo_w) if self.mlp is not None: x = x + self.mlp(norm(x), c_fc, c_proj) return x # ----------------------------------------------------------------------------- # The main model def next_multiple_of_n(v: float | int, *, n: int): return next(x for x in range(n, int(v) + 1 + n, n) if x >= v) @dataclass class ForwardScheduleConfig: mtp_weights: torch.Tensor ws_short: int ws_long: int class GPT(nn.Module): def __init__(self, vocab_size: int, num_layers: int, num_heads: int, head_dim: int, model_dim: int, max_seq_len: int): super().__init__() self.num_layers = num_layers vocab_size = next_multiple_of_n(vocab_size, n=128) self.smear_gate = nn.Linear(12, 1, bias=False) nn.init.zeros_(self.smear_gate.weight) self.smear_gate.weight.label = 'smear_gate' self.skip_gate = nn.Linear(12, 1, bias=False) nn.init.zeros_(self.skip_gate.weight) self.skip_gate.weight.label = 'skip_gate' # token value embeddings by @KoszarskyB - inspired by @Grad62304977's value residual implementation following https://arxiv.org/abs/2410.17897 # value embedding code simplification inspired by @ragulpr https://github.com/KellerJordan/modded-nanogpt/pull/78 self.value_embeds = nn.ModuleList([nn.Embedding(vocab_size, model_dim) for _ in range(5)]) for embed in self.value_embeds: nn.init.zeros_(embed.weight) for i, ve in enumerate(self.value_embeds): ve.weight.label = f've{i}' # ve0, ve1, ve2, ve3, ve4 # parameter banks for attention and value embedding gate weights self.attn_gate_bank = nn.Parameter(torch.zeros(10, num_heads, 12)) # 10 layers self.attn_gate_bank.label = 'attn_gate_bank' self.ve_gate_bank = nn.Parameter(torch.zeros(5, num_heads, 12)) # 5 unique gates self.ve_gate_bank.label = 've_gate_bank' # ----------------------------------- # Parameter banks for sharded optimization, by @chrisjmccormick # Identify which layers have attention/MLP # Attention is skipped in layer 6 by @YouJiacheng self.attn_layer_indices = [i for i in range(num_layers) if i != 6] # All layers have MLP (At 11 layers--dropped first layer @EmelyanenkoK) self.mlp_layer_indices = list(range(num_layers)) hdim = num_heads * head_dim mlp_hdim = 4 * model_dim # Create index mappings: layer_idx -> bank_idx self.layer_to_attn_idx = {layer_idx: bank_idx for bank_idx, layer_idx in enumerate(self.attn_layer_indices)} self.layer_to_mlp_idx = {layer_idx: bank_idx for bank_idx, layer_idx in enumerate(self.mlp_layer_indices)} # Attention bank: stores QKVO weights for all attention layers # merged QKVO weights: suggested by many, implemented by @fernbear.bsky.social, and further improved by @YouJiacheng # https://x.com/hi_tysam/status/1879699187107033311 # Simplified layout by @chrisjmccormick # Shape: (num_attn_layers, 4*model_dim, hdim) = (10, 3072, 768) # Reshape for sharding: (40, 768, 768) for even distribution across 8 GPUs self.attn_bank = nn.Parameter(torch.empty(len(self.attn_layer_indices), 4 * model_dim, hdim)) self.attn_bank.label = 'attn' self.attn_bank.reshape = (len(self.attn_layer_indices) * 4, hdim, hdim) # (40, 768, 768) # MLP bank: stores c_fc and c_proj for all MLP layers # Shape: (num_mlp_layers + padding, 2, mlp_hdim, model_dim) = (12, 2, 3072, 768) # We add 1 padding layer (index 11) to get 12*2=24 matrices for even distribution across 8 GPUs # Reshape for sharding: (24, 3072, 768) num_mlp_with_padding = len(self.mlp_layer_indices) + 1 # 11 + 1 = 12 self.mlp_bank = nn.Parameter(torch.empty(num_mlp_with_padding, 2, mlp_hdim, model_dim)) self.mlp_bank.label = 'mlp' self.mlp_bank.reshape = (num_mlp_with_padding * 2, mlp_hdim, model_dim) # (24, 3072, 768) # improved init scale by @YouJiacheng # Attention uses dim^-0.5, MLP uses 0.5 * dim^-0.5 attn_std = model_dim ** -0.5 attn_bound = (3 ** 0.5) * attn_std mlp_std = 0.5 * (model_dim ** -0.5) mlp_bound = (3 ** 0.5) * mlp_std with torch.no_grad(): # Init attention bank (QKV uniform, O zero) self.attn_bank[:, :model_dim * 3, :].uniform_(-attn_bound, attn_bound) self.attn_bank[:, model_dim * 3:, :].zero_() # Init MLP bank (c_fc uniform, c_proj zero) self.mlp_bank[:, 0, :, :].uniform_(-mlp_bound, mlp_bound) # c_fc self.mlp_bank[:, 1, :, :].zero_() # c_proj - zero init suggested by @Grad62304977 # Create blocks with has_attn/has_mlp flags self.paired_head_layers = [0, 2, 5, 9] self.blocks = nn.ModuleList([ Block(model_dim, head_dim, num_heads, has_attn=(i in self.layer_to_attn_idx), has_mlp=(i in self.layer_to_mlp_idx), use_paired_head=(i in self.paired_head_layers)) for i in range(num_layers) ]) self.yarn = Yarn(head_dim, max_seq_len) self.yarn_paired_head = YarnPairedHead(head_dim, max_seq_len) # there are only 50257 unique GPT-2 tokens; we extend to nearest multiple of 128 for efficiency. # suggested to me by @Grad62304977. this originates from Karpathy's experiments. use_fp8 = not os.environ.get("DISABLE_FP8", False) # Transposed weight storage for faster gradient accumulation self.lm_head = CastedLinearT(model_dim, vocab_size, use_fp8=use_fp8, x_s=100/448, w_s=1.6/448, grad_s=0.75/448) nn.init.normal_(self.lm_head.weight, mean=0, std=0.005) self.lm_head.weight.label = 'lm_head' self.embed = nn.Embedding(vocab_size, model_dim) self.embed.weight.label = 'embed' with torch.no_grad(): self.embed.weight.copy_(self.lm_head.weight.T) self.bigram_embed = nn.Embedding(args.bigram_vocab_size, model_dim) self.bigram_embed.weight.label = 'bigram_embed' nn.init.zeros_(self.bigram_embed.weight) # x0_lambdas separated out for different optimizer treatment (no beta smoothing) self.x0_lambdas = nn.Parameter(torch.zeros(num_layers)) self.x0_lambdas.label = 'x0_lambdas' pad = (-num_layers * 3 - 3) % dist.get_world_size() # updated: 3*num_layers instead of 4* self.scalars = nn.Parameter( torch.cat( [ 1.1 * torch.ones(num_layers), # resid lambdas. 1.1 init such that layer i weight is i^(num_layers-i). *[torch.tensor([0.5, 1.0]) for _ in range(num_layers)], # SA lambdas 0.1 * torch.ones(num_layers), # bigram lambdas torch.zeros(1), # smear_lambda 0.5*torch.ones(1), # backout_lambda -1.5 * torch.ones(1), # skip_lambda -> σ(-1.5) ≈ 0.18 torch.ones(pad), ] ) ) self.scalars.label = 'scalars' def forward(self, input_seq: Tensor, target_seq: Tensor, seqlens: Tensor, bigram_input_seq: Tensor, schedule_cfg: ForwardScheduleConfig): assert input_seq.ndim == 1 # unpack schedule_cfg mtp_weights, ws_short, ws_long = schedule_cfg.mtp_weights, schedule_cfg.ws_short, schedule_cfg.ws_long # set configs skip_connections = [] skip_in = [3] # long attention window on layer 3 skip_out = [6] # no attn op on layer 6 x_backout = None backout_layer = 7 # set lambdas resid_lambdas = self.scalars[: 1 * self.num_layers] x0_lambdas = self.x0_lambdas sa_lambdas = self.scalars[1 * self.num_layers: 3 * self.num_layers].view(-1, 2) bigram_lambdas = self.scalars[3 * self.num_layers: 4 * self.num_layers] smear_lambda = self.scalars[4 * self.num_layers] backout_lambda = self.scalars[4 * self.num_layers+1] skip_lambda = self.scalars[4 * self.num_layers+2] # set block masks and key shift short_bm = ws_short * args.block_size long_bm = ws_long * args.block_size bm_sizes = [short_bm, short_bm, short_bm, long_bm, short_bm, short_bm, None, short_bm, short_bm, short_bm, long_bm] assert len(bm_sizes) == self.num_layers key_offset = [b==long_bm for b in bm_sizes] # apply partial key offset to long windows # Embedding lookup - embed is synced from lm_head during tied phase by optimizer x = self.embed(input_seq) x0_bigram = self.bigram_embed(bigram_input_seq)[None] # Value embeddings - always computed (not precomputed) ve = [value_embed(input_seq) for value_embed in self.value_embeds] # 01 ... 01 structure on token value embeddings by @YouJiacheng, improved on @leloykun's U-net structure # shifting first layer updates this to 01 ... 01 @photomz ve = [ve[0], ve[1]] + [None] * (self.num_layers - 5) + [ve[2], ve[3], ve[4]] assert len(ve) == self.num_layers # smear token embed forward 1 position @classiclarryd smear_gate_out = smear_lambda * torch.sigmoid(self.smear_gate(x[1:, :self.smear_gate.weight.size(-1)])) x = torch.cat([x[:1], x[1:] + smear_gate_out * x[:-1]]) x = x0 = norm(x[None]) # unbind gate banks to avoid select_backwards kernel ag = [w.bfloat16() for w in self.attn_gate_bank.unbind(0)] veg = [w.bfloat16() for w in self.ve_gate_bank.unbind(0)] attn_gates = ag[:6] + [None] + ag[6:] ve_gates = [veg[0], veg[1]] + [None] * (self.num_layers - 5) + [veg[2], veg[3], veg[4]] assert len(attn_gates) == self.num_layers assert len(ve_gates) == self.num_layers # unbind weight banks to avoid select_backwards kernel attn_weights = self.attn_bank.unbind(0) # tuple of [4*dim, hdim] tensors mlp_fcs = self.mlp_bank[:, 0, :, :].unbind(0) # tuple of [mlp_hdim, dim] tensors mlp_projs = self.mlp_bank[:, 1, :, :].unbind(0) # tuple of [mlp_hdim, dim] tensors for i in range(self.num_layers): yarn = self.yarn_paired_head if i in self.paired_head_layers else self.yarn attn_args = AttnArgs( ve=ve[i], sa_lambdas=sa_lambdas[i], seqlens=seqlens, bm_size=bm_sizes[i], yarn=yarn, key_offset=key_offset[i], attn_gate_w=attn_gates[i], ve_gate_w=ve_gates[i] ) if i in skip_out: skip_gate_out = torch.sigmoid(skip_lambda) * 2 * torch.sigmoid(self.skip_gate(x0[..., :self.skip_gate.weight.size(-1)])) x = x + skip_gate_out * skip_connections.pop() if i == 0: x = (resid_lambdas[0] + x0_lambdas[0]) * x + bigram_lambdas[0] * x0_bigram else: x = resid_lambdas[i] * x + x0_lambdas[i] * x0 + bigram_lambdas[i] * x0_bigram # Get weights for this layer from banks qkvo_w = attn_weights[self.layer_to_attn_idx[i]] if i in self.layer_to_attn_idx else None c_fc = mlp_fcs[self.layer_to_mlp_idx[i]] if i in self.layer_to_mlp_idx else None c_proj = mlp_projs[self.layer_to_mlp_idx[i]] if i in self.layer_to_mlp_idx else None x = self.blocks[i](x, attn_args, qkvo_w, c_fc, c_proj) if i in skip_in: skip_connections.append(x) if i == backout_layer: x_backout = x # back out contributions from first 7 layers that are only required for downstream context and not direct prediction x -= backout_lambda * x_backout x = norm(x) logits = self.lm_head(x) # @Grad62304977 added tanh softcapping following Gemma 2 paper, @KoszarskyB reduced it from 30 to 15 # @YouJiacheng shifted it by +15 (2*sigmoid(2*x)=tanh(x)+1). @classiclarryd updated to 23*sigmoid((logits+5)/7.5) if self.training: losses = FusedSoftcappedCrossEntropy.apply(logits.view(-1, logits.size(-1)), target_seq, mtp_weights) loss = losses.sum() else: logits = 23 * torch.sigmoid((logits + 5) / 7.5) logits_for_loss = logits.float() loss = F.cross_entropy(logits_for_loss.view(-1, logits_for_loss.size(-1)), target_seq, reduction="mean") return loss # ----------------------------------------------------------------------------- # Distributed data loader def _load_data_shard(file: Path): header = torch.from_file(str(file), False, 256, dtype=torch.int32) # header is 256 int32 assert header[0] == 20240520, "magic number mismatch in the data .bin file" assert header[1] == 1, "unsupported version" num_tokens = int(header[2]) # number of tokens (claimed) with file.open("rb", buffering=0) as f: tokens = torch.empty(num_tokens, dtype=torch.uint16, pin_memory=True) # avoid pin_memory copy by @YouJiacheng f.seek(256 * 4) nbytes = f.readinto(tokens.numpy()) # avoid bytes->array copy by @YouJiacheng assert nbytes == 2 * num_tokens, "number of tokens read does not match header" return tokens BOS_ID = 50256 class BOSFinder: # Helper for getting sequences that start at the beginning of documents by @varunneal based on work by @classiclarryd def __init__(self, tokens: Tensor, world_size: int = 1, quickload: bool = False): # Precompute BOS positions once per shard self.tokens=tokens self.size = tokens.numel() self.quickload = quickload if quickload: # only scan first 4 million tokens, then kickoff async thread to scan rest self.bos_idx = (tokens[:4_000_000] == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self.thread = None self.ready = threading.Event() self.start() else: self.bos_idx = (tokens == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self.i = 0 self.world_size = world_size self.batch_iter = 0 def _load(self): self.bos_idx_async = (self.tokens == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self.ready.set() def start(self): self.ready.clear() self.thread = threading.Thread(target=self._load) self.thread.start() def get(self): if self.thread: self.ready.wait() self.thread.join() self.bos_idx = self.bos_idx_async def next_batch(self, num_tokens_local: int, max_seq_len: int): # if quickload was used, repoint to the full dataset after 5 batches if self.quickload and self.batch_iter==5: self.get() n = len(self.bos_idx) starts = [[] for _ in range(self.world_size)] ends = [[] for _ in range(self.world_size)] idx = self.i for r in range(self.world_size): cur_len = 0 while cur_len <= num_tokens_local: if idx >= n: raise StopIteration(f"Insufficient BOS ahead; hit tail of shard.") cur = self.bos_idx[idx] starts[r].append(cur) end = min(self.bos_idx[idx + 1] if idx + 1 < n else self.size, cur + max_seq_len, cur + num_tokens_local - cur_len + 1) ends[r].append(end) cur_len += end - cur idx += 1 assert cur_len == num_tokens_local + 1 self.i = idx self.batch_iter+=1 return starts, ends class DataPreloader: # Helper for asynchronously loading next shard and indexing bos tokens def __init__(self, file_iter, world_size: int = 1): self.file_iter = file_iter self.world_size = world_size self.thread = None self.data = None self.ready = threading.Event() def _load(self): tokens = _load_data_shard(next(self.file_iter)) self.data = (tokens, BOSFinder(tokens, self.world_size)) self.ready.set() def start(self): self.ready.clear() self.thread = threading.Thread(target=self._load) self.thread.start() def get(self): if self.thread: self.ready.wait() self.thread.join() return self.data def get_bigram_hash(x): """ Computes bigram hash for each position using [prev_token, curr_token]. Multiply by arbitary large ints to get even spread over int32 range. Position 0 is mapped to the reserved index (vocab_size - 1). BOS_tokens within the batch will hash based on last token of prior doc. Masking this ran slower and showed no improvement. """ rand_int_1 = 36313 rand_int_2 = 27191 mod = args.bigram_vocab_size-1 x = x.to(torch.int32).clone() x[0] = mod x[1:] = torch.bitwise_xor(rand_int_1 * x[1:], rand_int_2 * x[:-1]) % mod return x def distributed_data_generator(filename_pattern: str, num_tokens: int, max_seq_len: int, grad_accum_steps: int = 1, align_to_bos: bool = True): # align_to_bos: each sequence begins with Beginning of Sequence token, sequences truncated to max_seq_len rank = dist.get_rank() if dist.is_initialized() else 0 world_size = dist.get_world_size() if dist.is_initialized() else 1 assert num_tokens % (world_size * grad_accum_steps) == 0, "Batch size must be divisible by world size" num_tokens = num_tokens // grad_accum_steps files = [Path(file) for file in sorted(glob.glob(filename_pattern))] if not files: raise FileNotFoundError(f"No files found for pattern: {filename_pattern}") file_iter = iter(files) # Use itertools.cycle(files) for multi-epoch training tokens = _load_data_shard(next(file_iter)) if align_to_bos: finder = BOSFinder(tokens, world_size=world_size, quickload=True) preloader = DataPreloader(file_iter, world_size) preloader.start() else: pos = 0 # for unaligned case while True: num_tokens_local = num_tokens // world_size max_num_docs = next_multiple_of_n(num_tokens_local // 300, n=128) # median doc length is ~400 if align_to_bos: try: seq_starts, seq_ends = finder.next_batch(num_tokens_local, max_seq_len) start_idxs, end_idxs = torch.tensor(seq_starts[rank]), torch.tensor(seq_ends[rank]) except StopIteration: # This shard is exhausted, load the next one in the next loop iteration. tokens, finder = preloader.get() preloader.start() continue buf = torch.cat([tokens[i:j] for i, j in zip(start_idxs, end_idxs)]) _inputs = buf[:-1] _targets = buf[1:] end_idxs[-1] -= 1 # last document was too long to account for _targets offset cum_lengths = (end_idxs - start_idxs).cumsum(0) else: if pos + num_tokens + 1 >= len(tokens): # should not occur for val data tokens, pos = _load_data_shard(next(file_iter)), 0 pos_local = pos + rank * num_tokens_local buf = tokens[pos_local: pos_local + num_tokens_local + 1] _inputs = buf[:-1].view(num_tokens_local, ) _targets = buf[1:].view(num_tokens_local, ) cum_lengths = torch.nonzero(_inputs == BOS_ID)[:, 0] pos += num_tokens _cum_lengths = torch.full((max_num_docs,), num_tokens_local) _cum_lengths[0] = 0 _cum_lengths[1:len(cum_lengths) + 1] = cum_lengths # Cast to int32 on CPU before transfer to avoid dtype conversion during .to() _inputs = _inputs.to(dtype=torch.int32) _targets = _targets.to(dtype=torch.int64) _cum_lengths = _cum_lengths.to(dtype=torch.int32) _bigram_inputs = get_bigram_hash(_inputs) new_params = yield ( _inputs.to(device="cuda", non_blocking=True), _targets.to(device="cuda", non_blocking=True), _cum_lengths.to(device="cuda", non_blocking=True), _bigram_inputs.to(device="cuda", non_blocking=True) ) if new_params is not None: # makes it possible for generator to receive new (num_tokens, max_seq_len, grad_accum_steps) via .send() new_num_tokens, new_max_seq_len, new_grad_accum_steps = new_params assert new_num_tokens % (world_size * new_grad_accum_steps) == 0, "Num tokens must be divisible by world size" num_tokens = new_num_tokens // new_grad_accum_steps max_seq_len = new_max_seq_len # ----------------------------------------------------------------------------- # Training Management def get_bs(step: int): if step >= args.num_scheduled_iterations: return args.train_bs_extension x = step / args.num_scheduled_iterations bs_idx = int(len(args.train_bs_schedule) * x) return args.train_bs_schedule[bs_idx] def get_ws(step: int): # set short window size to half of long window size # Higher ws on "extension" steps if step >= args.num_scheduled_iterations: return args.ws_final // 2, args.ws_final x = step / args.num_scheduled_iterations assert 0 <= x < 1 ws_idx = int(len(args.ws_schedule) * x) return args.ws_schedule[ws_idx] // 2, args.ws_schedule[ws_idx] # learning rate schedule: tied to batch size schedule, with cooldown at the end. def get_lr(step: int): if step > args.num_scheduled_iterations: return 0.1 lr_max = 1.0 x = step / args.num_scheduled_iterations if x > 1/3: lr_max = 1.52 # (16/8)**0.6 if x > 2/3: lr_max = 1.73 # (24/8)**0.5 if x >= 1 - args.cooldown_frac: w = (1 - x) / args.cooldown_frac lr = lr_max * w + (1 - w) * 0.1 return lr return lr_max def get_muon_momentum(step: int, muon_warmup_steps=300, muon_cooldown_steps=50, momentum_min=0.85, momentum_max=0.95): # warmup phase: linearly increase momentum from min to max # cooldown phase: linearly decrease momentum from max to min momentum_cd_start = args.num_iterations - muon_cooldown_steps if step < muon_warmup_steps: frac = step / muon_warmup_steps momentum = momentum_min + frac * (momentum_max - momentum_min) elif step > momentum_cd_start: frac = (step - momentum_cd_start) / muon_cooldown_steps momentum = momentum_max - frac * (momentum_max - momentum_min) else: momentum = momentum_max return momentum class TrainingManager(): """ Manages the NorMuonAndAdam for all parameters with explicit ordering. Notable Features: 1. Scalars are given higher momentum terms to smooth learning @ChrisJMcCormick 2. Adam optimizers are only stepped on odd steps @classiclarryd 3. Explicit scatter_order and work_order for communication scheduling (no backward hooks) 4. Muon has a linear momentum warmup and cooldown schedule 5. Learning rates follow a linear decay schedule 6. Embed is tied to lm_head until split step (2/3 of training), then untied @classiclarryd Manages model architecture, data, and target that changes during training Notable Features: 1. Multi Token Prediction schedule of [1, 0.5, 0.25->0] -> [1, 0.5->0] -> [1] @varunneal 2. Sliding Attention window schedule of [1,3] -> [3,7] -> [5,11] -> [6,13] 3. YaRN updates to RoPE on window changes 4. Split embed and lm_head at 2/3 of training (weights and optimizer state copied) 5. Batch size schedule of 8 -> 16 -> 24 6. Post training extension of long windows from 13 to 20 """ def __init__(self, model): self.mtp_weights_schedule = self._build_mtp_schedule() self.model = model # - Ordering dictates when to launch reduce/reduce_scatter operations # - "sharded" parameters use reduce_scatter/all_gather and "replicated" ones use all_reduce # - lr_mul and wd_mul are per-parameter learning rate and weight decay multipliers self.param_table = { "attn": {"optim": "normuon", "comms": "sharded", "adam_betas": None}, "mlp": {"optim": "normuon", "comms": "sharded", "adam_betas": None}, "scalars": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99], "lr_mul": 5.0, "wd_mul": 0.0}, "ve0": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve1": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve2": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve3": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve4": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "bigram_embed": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "smear_gate": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99], "lr_mul": 0.01, "wd_mul": 0.0}, "skip_gate": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99], "lr_mul": 0.05, "wd_mul": 0.0}, "attn_gate_bank": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99]}, "ve_gate_bank": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99]}, "x0_lambdas": {"optim": "adam", "comms": "replicated", "adam_betas": [0.65, 0.95], "lr_mul": 5.0, "wd_mul": 0.0}, "lm_head": {"optim": "adam", "comms": "sharded", "adam_betas": [0.5, 0.95], "wd_mul": 150.}, "embed": {"optim": "adam", "comms": "sharded", "adam_betas": [0.5, 0.95], "wd_mul": 150.}, } # - Process smaller/faster params first while large reduces complete # - lm_head must complete before embed sync (when tied) self.work_order = [ "scalars", "smear_gate", "skip_gate", "attn_gate_bank", "ve_gate_bank", "x0_lambdas", # Small, fast "ve0", "ve1", "ve2", "ve3", "ve4", "bigram_embed", # Medium "lm_head", "embed", # lm_head must complete before embed sync (when tied) "attn", "mlp", # Large, polar express - process last to maximize overlap ] adam_defaults = dict( lr=0.008, eps=1e-10, weight_decay=0.005, ) normuon_defaults = dict( lr=0.023, momentum=0.95, beta2=0.95, weight_decay=1.2, ) self.optimizer = NorMuonAndAdam( model.named_parameters(), param_table=self.param_table, scatter_order=list(self.param_table.keys()), # Dict order defines scatter priority work_order=self.work_order, adam_defaults=adam_defaults, normuon_defaults=normuon_defaults, ) # Split embed from lm_head at 2/3 of training (on an odd step so Adam updates) self.split_step = math.ceil(args.split_embed_frac * args.num_scheduled_iterations) | 1 self.reset() def _build_mtp_schedule(self): # Precompute MTP weights for all steps to avoid tensor allocation during training # Schedule: [1, 0.5, 0.25->0] -> [1, 0.5->0] -> [1] mtp_weights_schedule = [] for s in range(args.num_iterations + 1): x = s / args.num_scheduled_iterations if x < 1/3: w = [1.0, 0.5, 0.25 * (1 - 3*x)] elif x < 2/3: w = [1.0, 0.5 * (1 - (3*x - 1))] else: w = [1.0] mtp_weights_schedule.append(torch.tensor(w, device=device)) return mtp_weights_schedule def apply_final_ws_ext(self): self.ws_long = args.ws_validate_post_yarn_ext def get_forward_args(self): return ForwardScheduleConfig( mtp_weights = self.mtp_weights, ws_short = self.ws_short, ws_long = self.ws_long ) def _is_adam_step(self, step: int): """Adam params are only updated on odd steps.""" return step % 2 == 1 def get_transition_steps(self): transition_steps = [] ws_short, ws_long = get_ws(0) for step in range(1, args.num_iterations): ws_short, new_ws_long = get_ws(step) if new_ws_long != ws_long: transition_steps.append(step) ws_long = new_ws_long return transition_steps def advance_schedule(self, step: int): self.ws_short, new_ws_long = get_ws(step) if new_ws_long != self.ws_long: self.model.yarn.apply(self.ws_long, new_ws_long) self.model.yarn_paired_head.apply(self.ws_long, new_ws_long) new_batch_size = get_bs(step) if new_batch_size != self.batch_size: self.train_loader_send_args = (new_batch_size, args.train_max_seq_len, grad_accum_steps) self.batch_size = new_batch_size else: self.train_loader_send_args = None self.ws_long = new_ws_long self.mtp_weights = self.mtp_weights_schedule[step] def step_optimizers(self, step: int): step_lr = get_lr(step) muon_momentum = get_muon_momentum(step) do_adam = self._is_adam_step(step) # Update learning rates and momentum for all params for param, p_cfg in self.optimizer.param_cfgs.items(): p_cfg.lr = p_cfg.initial_lr * step_lr if p_cfg.optim == "normuon": p_cfg.momentum = muon_momentum # Step optimizer with do_adam flag self.optimizer.step(do_adam=do_adam) # At split step: copy lm_head optimizer state to embed and mark as split if step == self.split_step: self.optimizer.copy_lm_state_to_embed() def reset(self, state=None): if state is not None: self.optimizer.load_state_dict(state) # Reset NorMuon momentum buffers and split_embed state self.optimizer.reset() self.ws_short, self.ws_long = get_ws(0) self.batch_size = get_bs(0) self.model.yarn.reset() self.model.yarn_paired_head.reset() def get_state(self): return copy.deepcopy(self.optimizer.state_dict()) # ----------------------------------------------------------------------------- # int main @dataclass class Hyperparameters: # data train_files: str = "data/fineweb10B/fineweb_train_*.bin" # input .bin to train on val_files: str = "data/fineweb10B/fineweb_val_*.bin" # input .bin to eval validation loss on val_tokens: int = 10485760 # how many tokens of validation data? it's important to keep this fixed for consistent comparisons # batch sizes train_bs_schedule: tuple = (8 * 2048 * 8, 16 * 2048 * 8, 24 * 2048 * 8) train_bs_extension: int = 24 * 2048 * 8 train_max_seq_len: int = 128 * 16 val_batch_size: int = 4 * 64 * 1024 * 8 # optimization num_scheduled_iterations: int = 1535 # number of steps to complete lr and ws schedule num_extension_iterations: int = 40 # number of steps to continue training at final lr and ws num_iterations: int = num_scheduled_iterations + num_extension_iterations cooldown_frac: float = 0.55 # fraction of num_scheduled_iterations spent cooling down the learning rate split_embed_frac: float = 2/3 # fraction of training when embeddings split from lm_head # evaluation and logging run_id: str = f"{uuid.uuid4()}" val_loss_every: int = 250 # every how many steps to evaluate val loss? 0 for only at the end save_checkpoint: bool = False # attention masking block_size: int = 128 ws_schedule: tuple = (3, 7, 11) ws_final: int = 13 # increase final validation ws, used for YaRN extension and short window size @classiclarryd ws_validate_post_yarn_ext: int = 20 # extend long windows out even further after applying YaRN # bigram hash embedding bigram_vocab_size = 50304 * 5 args = Hyperparameters() data_path = os.environ.get("DATA_PATH", ".") args.train_files = os.path.join(data_path, args.train_files) args.val_files = os.path.join(data_path, args.val_files) # torchrun sets these env variables rank = int(os.environ["RANK"]) world_size = int(os.environ["WORLD_SIZE"]) assert 8 % world_size == 0, "world_size must be a divisor of 8" grad_accum_steps = 8 // world_size assert torch.cuda.is_available() device = torch.device("cuda", int(os.environ["LOCAL_RANK"])) torch.cuda.set_device(device) dist.init_process_group(backend="nccl", device_id=device) dist.barrier() master_process = (rank == 0) # this process will do logging, checkpointing etc. # begin logging logfile = None if master_process: run_id = args.run_id os.makedirs("logs", exist_ok=True) logfile = f"logs/{run_id}.txt" print(logfile) def print0(s, console=False): if master_process: with open(logfile, "a") as f: if console: print(s) print(s, file=f) # begin by printing this file (the Python code) print0(code) print0("="*100) # log information about the hardware/software environment this is running on print0(f"Running Python {sys.version}") print0(f"Running PyTorch {torch.version.__version__} compiled for CUDA {torch.version.cuda}") print0(f"Running Triton version {triton.__version__}") def nvidia_smi(): import subprocess # avoid top level import return subprocess.run(["nvidia-smi"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True).stdout print0(nvidia_smi()) print0("="*100) model: nn.Module = GPT( vocab_size=50257, num_layers=11, num_heads=6, head_dim=128, model_dim=768, max_seq_len=args.val_batch_size // (grad_accum_steps * world_size) ).cuda() for m in model.modules(): if isinstance(m, (nn.Embedding, nn.Linear)): m.weight.data = m.weight.data.bfloat16() model.attn_gate_bank.data = model.attn_gate_bank.data.bfloat16() model.ve_gate_bank.data = model.ve_gate_bank.data.bfloat16() model.attn_bank.data = model.attn_bank.data.bfloat16() model.mlp_bank.data = model.mlp_bank.data.bfloat16() for param in model.parameters(): dist.broadcast(param.detach(), 0) model: nn.Module = torch.compile(model, dynamic=False, fullgraph=True) training_manager = TrainingManager(model) ######################################## # Warmup kernels # ######################################## print0("Compiling model and warming up kernels (~7 minutes on first execution)", console=True) # Warmup the training kernels, then re-initialize the state so we aren't cheating initial_state = dict(model=copy.deepcopy(model.state_dict()), optimizer=training_manager.get_state()) # save the initial state train_loader = distributed_data_generator(args.train_files, args.train_bs_schedule[0], args.train_max_seq_len, grad_accum_steps=grad_accum_steps) val_loader = distributed_data_generator(args.val_files, args.val_batch_size, -1, grad_accum_steps=grad_accum_steps, align_to_bos=False) transition_steps = training_manager.get_transition_steps() # first few steps plus transitions warmup_steps = sorted({0, 1, 2} | set(s + offset for s in transition_steps for offset in [-1, 0, 1] if s + offset >= 0)) print0(f"Sampling steps {warmup_steps} for warmup", console=True) for step in warmup_steps: training_manager.advance_schedule(step) model.eval() with torch.no_grad(): inputs, targets, cum_seqlens, bigram_inputs = next(val_loader) model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) model.train() for idx in range(grad_accum_steps): send_args = training_manager.train_loader_send_args inputs, targets, cum_seqlens, bigram_inputs = train_loader.send(send_args) (model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) / grad_accum_steps).backward() training_manager.step_optimizers(step) print0("Resetting Model", console=True) model.zero_grad(set_to_none=True) model.load_state_dict(initial_state["model"]) training_manager.reset(initial_state["optimizer"]) del val_loader, train_loader, initial_state model.train() ######################################## # Training and validation # ######################################## train_loader = distributed_data_generator(args.train_files, args.train_bs_schedule[0], args.train_max_seq_len, grad_accum_steps=grad_accum_steps) gc.collect() training_time_ms = 0 # start the clock torch.cuda.synchronize() t0 = time.perf_counter() # begin training train_steps = args.num_iterations for step in range(train_steps + 1): last_step = (step == train_steps) training_manager.advance_schedule(step) # --------------- VALIDATION SECTION ----------------- if last_step or (args.val_loss_every > 0 and step % args.val_loss_every == 0): if last_step: training_manager.apply_final_ws_ext() # stop the clock torch.cuda.synchronize() training_time_ms += 1000 * (time.perf_counter() - t0) model.eval() assert args.val_tokens % args.val_batch_size == 0 val_steps = grad_accum_steps * args.val_tokens // args.val_batch_size val_loader = distributed_data_generator(args.val_files, args.val_batch_size, -1, grad_accum_steps=grad_accum_steps, align_to_bos=False) val_loss = 0 with torch.no_grad(): for _ in range(val_steps): inputs, targets, cum_seqlens, bigram_inputs = next(val_loader) val_loss += model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) val_loss /= val_steps del val_loader dist.reduce(val_loss, 0, op=dist.ReduceOp.AVG) print0(f"step:{step}/{train_steps} val_loss:{val_loss:.4f} train_time:{training_time_ms:.0f}ms step_avg:{training_time_ms/max(step, 1):.2f}ms", console=True) model.train() # start the clock again torch.cuda.synchronize() t0 = time.perf_counter() if last_step: if master_process and args.save_checkpoint: log = dict(step=step, code=code, model=model.state_dict(), optimizer=training_manager.get_state()) os.makedirs(f"logs/{run_id}", exist_ok=True) torch.save(log, f"logs/{run_id}/state_step{step:06d}.pt") # the last step only has the validation loop, so break to avoid training break # --------------- TRAINING SECTION ----------------- for idx in range(grad_accum_steps): inputs, targets, cum_seqlens, bigram_inputs = train_loader.send(training_manager.train_loader_send_args) (model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) / grad_accum_steps).backward() training_manager.step_optimizers(step) # logging approx_training_time_ms = training_time_ms + 1000 * (time.perf_counter() - t0) print0(f"step:{step+1}/{train_steps} train_time:{approx_training_time_ms:.0f}ms step_avg:{approx_training_time_ms/(step + 1):.2f}ms", console=True) print0(f"peak memory allocated: {torch.cuda.max_memory_allocated() // 1024 // 1024} MiB " f"reserved: {torch.cuda.max_memory_reserved() // 1024 // 1024} MiB", console=True) dist.destroy_process_group() ---------------------------------------- # triton_kernels.py ---------------------------------------- import torch import triton import triton.language as tl from triton.tools.tensor_descriptor import TensorDescriptor # ----------------------------------------------------------------------------- # Triton kernel for symmetric matrix multiplication by @byronxu99 def _get_autotune_configs(): return [ triton.Config( { "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk, "GROUP_SIZE_M": 8, "LOWER_UPPER": 1, }, num_stages=stages, num_warps=warps, ) for bm in [64, 128] for bn in [64, 128, 256] for bk in [64, 128] for stages, warps in [(3, 4), (3, 8), (4, 4)] if bm // bn <= 2 and bn // bm <= 2 ] @triton.jit def _pid_to_block( pid, M, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, GROUP_SIZE_M: tl.constexpr, ): # Split output matrix into blocks of size (BLOCK_SIZE_M, BLOCK_SIZE_N) num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) num_pid_n = tl.cdiv(M, BLOCK_SIZE_N) # Map PID to a single matrix in batch batch_idx = pid // (num_pid_m * num_pid_n) pid = pid % (num_pid_m * num_pid_n) # Map PID to 2D grid of blocks pid_m = pid // num_pid_n pid_n = pid % num_pid_n pid_m, pid_n = tl.swizzle2d(pid_m, pid_n, num_pid_m, num_pid_n, GROUP_SIZE_M) m_idx = pid_m * BLOCK_SIZE_M n_idx = pid_n * BLOCK_SIZE_N return batch_idx, m_idx, n_idx @triton.autotune( configs=_get_autotune_configs(), key=["M", "K", "a_stride_r", "a_stride_c", "c_stride_r", "c_stride_c"], ) @triton.jit def XXT_kernel( A_ptr, C_ptr, M, K, a_stride_b, a_stride_r, a_stride_c, c_stride_b, c_stride_r, c_stride_c, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, LOWER_UPPER: tl.constexpr, ): pid = tl.program_id(axis=0) batch_idx, m_idx, n_idx = _pid_to_block( pid, M, BLOCK_SIZE_M, BLOCK_SIZE_N, GROUP_SIZE_M ) # Skip blocks that don't need to be computed skip_block_below_diag = (LOWER_UPPER == 0) and (n_idx + BLOCK_SIZE_N <= m_idx) skip_block_above_diag = (LOWER_UPPER != 0) and (m_idx + BLOCK_SIZE_M <= n_idx) if skip_block_below_diag or skip_block_above_diag: return # Index into one matrix of batch A_ptr += batch_idx * a_stride_b C_ptr += batch_idx * c_stride_b # Create pointer arrays for A and A.T offs_m = (m_idx + tl.arange(0, BLOCK_SIZE_M)) % M offs_n = (n_idx + tl.arange(0, BLOCK_SIZE_N)) % M offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = A_ptr + (offs_m[:, None] * a_stride_r + offs_k[None, :] * a_stride_c) at_ptrs = A_ptr + (offs_k[:, None] * a_stride_c + offs_n[None, :] * a_stride_r) accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) # Accumulate over blocks of K for k in tl.range(0, tl.cdiv(K, BLOCK_SIZE_K)): a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0) at = tl.load(at_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0) accumulator = tl.dot(a, at, accumulator) a_ptrs += BLOCK_SIZE_K * a_stride_c at_ptrs += BLOCK_SIZE_K * a_stride_c out_dtype = C_ptr.dtype.element_ty output = accumulator.to(out_dtype) # Store block of C offs_cm = m_idx + tl.arange(0, BLOCK_SIZE_M) offs_cn = n_idx + tl.arange(0, BLOCK_SIZE_N) c_ptrs = C_ptr + (offs_cm[:, None] * c_stride_r + offs_cn[None, :] * c_stride_c) c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < M) tl.store(c_ptrs, output, mask=c_mask) # Store block of C mirrored across the diagonal c_ptrs_t = C_ptr + (offs_cn[:, None] * c_stride_r + offs_cm[None, :] * c_stride_c) c_mask_t = (offs_cn[:, None] < M) & (offs_cm[None, :] < M) tl.store(c_ptrs_t, output.T, mask=c_mask_t) def XXT(A: torch.Tensor, out: torch.Tensor): """ Launch Triton kernel to compute C = A @ A.T """ assert A.ndim == 2 or A.ndim == 3 M, K = A.shape[-2:] assert out.size(-2) == M, "Output matrix has incorrect shape" assert out.size(-1) == M, "Output matrix has incorrect shape" batch_size = A.size(0) if A.ndim == 3 else 1 input_batch_stride = A.stride(0) if A.ndim == 3 else 0 output_batch_stride = out.stride(0) if out.ndim == 3 else 0 grid = lambda meta: ( batch_size * triton.cdiv(M, meta["BLOCK_SIZE_M"]) * triton.cdiv(M, meta["BLOCK_SIZE_N"]), ) XXT_kernel[grid]( A_ptr=A, C_ptr=out, M=M, K=K, a_stride_b=input_batch_stride, a_stride_r=A.stride(-2), a_stride_c=A.stride(-1), c_stride_b=output_batch_stride, c_stride_r=out.stride(-2), c_stride_c=out.stride(-1), ) return out @triton.autotune( configs=_get_autotune_configs(), key=["M", "a_stride_r", "a_stride_c", "c_stride_r", "c_stride_c"], ) @triton.jit def ba_plus_cAA_kernel( A_ptr, C_ptr, M, a_stride_b, a_stride_r, a_stride_c, c_stride_b, c_stride_r, c_stride_c, alpha, beta, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, LOWER_UPPER: tl.constexpr, ): # This is mostly duplicated from XXT_kernel, but also loads and adds a block of A # Performance is slightly slower than XXT_kernel, so we use two separate kernels pid = tl.program_id(axis=0) batch_idx, m_idx, n_idx = _pid_to_block( pid, M, BLOCK_SIZE_M, BLOCK_SIZE_N, GROUP_SIZE_M ) # Skip blocks that don't need to be computed skip_block_below_diag = (LOWER_UPPER == 0) and (n_idx + BLOCK_SIZE_N <= m_idx) skip_block_above_diag = (LOWER_UPPER != 0) and (m_idx + BLOCK_SIZE_M <= n_idx) if skip_block_below_diag or skip_block_above_diag: return # Index into one matrix of batch A_ptr += batch_idx * a_stride_b C_ptr += batch_idx * c_stride_b # Create pointer arrays for A and A.T offs_m = (m_idx + tl.arange(0, BLOCK_SIZE_M)) % M offs_n = (n_idx + tl.arange(0, BLOCK_SIZE_N)) % M offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = A_ptr + (offs_m[:, None] * a_stride_r + offs_k[None, :] * a_stride_c) at_ptrs = A_ptr + (offs_k[:, None] * a_stride_c + offs_n[None, :] * a_stride_r) accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) # Accumulate over blocks of K for k in tl.range(0, tl.cdiv(M, BLOCK_SIZE_K)): a = tl.load(a_ptrs, mask=offs_k[None, :] < M - k * BLOCK_SIZE_K, other=0.0) at = tl.load(at_ptrs, mask=offs_k[:, None] < M - k * BLOCK_SIZE_K, other=0.0) accumulator = tl.dot(a, at, accumulator) a_ptrs += BLOCK_SIZE_K * a_stride_c at_ptrs += BLOCK_SIZE_K * a_stride_c # Load block of A to add (corresponds to the current block of C) offs_am = m_idx + tl.arange(0, BLOCK_SIZE_M) offs_an = n_idx + tl.arange(0, BLOCK_SIZE_N) a_add_ptrs = A_ptr + (offs_am[:, None] * a_stride_r + offs_an[None, :] * a_stride_c) a_add_mask = (offs_am[:, None] < M) & (offs_an[None, :] < M) a_add = tl.load(a_add_ptrs, mask=a_add_mask, other=0.0).to(tl.float32) # Apply alpha and beta accumulator *= alpha accumulator += a_add * beta out_dtype = C_ptr.dtype.element_ty output = accumulator.to(out_dtype) # Store block of C offs_cm = m_idx + tl.arange(0, BLOCK_SIZE_M) offs_cn = n_idx + tl.arange(0, BLOCK_SIZE_N) c_ptrs = C_ptr + (offs_cm[:, None] * c_stride_r + offs_cn[None, :] * c_stride_c) c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < M) tl.store(c_ptrs, output, mask=c_mask) # Store block of C mirrored across the diagonal c_ptrs_t = C_ptr + (offs_cn[:, None] * c_stride_r + offs_cm[None, :] * c_stride_c) c_mask_t = (offs_cn[:, None] < M) & (offs_cm[None, :] < M) tl.store(c_ptrs_t, output.T, mask=c_mask_t) def ba_plus_cAA(A: torch.Tensor, alpha: float, beta: float, out: torch.Tensor): """ Launch Triton kernel to compute C = alpha * A @ A.T + beta * A """ assert A.ndim == 2 or A.ndim == 3 M, K = A.shape[-2:] assert M == K, "Input matrix must be square" assert out.size(-2) == M assert out.size(-1) == M batch_size = A.size(0) if A.ndim == 3 else 1 input_batch_stride = A.stride(0) if A.ndim == 3 else 0 output_batch_stride = out.stride(0) if out.ndim == 3 else 0 grid = lambda meta: ( batch_size * triton.cdiv(M, meta["BLOCK_SIZE_M"]) * triton.cdiv(M, meta["BLOCK_SIZE_N"]), ) ba_plus_cAA_kernel[grid]( A_ptr=A, C_ptr=out, M=M, a_stride_b=input_batch_stride, a_stride_r=A.stride(-2), a_stride_c=A.stride(-1), c_stride_b=output_batch_stride, c_stride_r=out.stride(-2), c_stride_c=out.stride(-1), alpha=alpha, beta=beta, ) return out # ----------------------------------------------------------------------------- # Triton kernel for MLP: relu(x @ W1.T)^2, by @andrewbriand, @jrauvola @triton.jit def linear_relu_square_kernel(a_desc, b_desc, c_desc, aux_desc, M, N, K, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, NUM_SMS: tl.constexpr, FORWARD: tl.constexpr, ): dtype = tl.bfloat16 start_pid = tl.program_id(axis=0) num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) k_tiles = tl.cdiv(K, BLOCK_SIZE_K) num_tiles = num_pid_m * num_pid_n tile_id_c = start_pid - NUM_SMS num_pid_in_group = GROUP_SIZE_M * num_pid_n for tile_id in tl.range(start_pid, num_tiles, NUM_SMS, flatten=True): pid_m = tile_id // num_pid_n pid_n = tile_id % num_pid_n offs_am = pid_m * BLOCK_SIZE_M offs_bn = pid_n * BLOCK_SIZE_N accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) for ki in range(k_tiles): offs_k = ki * BLOCK_SIZE_K a = a_desc.load([offs_am, offs_k]) b = b_desc.load([offs_bn, offs_k]) accumulator = tl.dot(a, b.T, accumulator) tile_id_c += NUM_SMS pid_m = tile_id // num_pid_n pid_n = tile_id % num_pid_n offs_am_c = pid_m * BLOCK_SIZE_M offs_bn_c = pid_n * BLOCK_SIZE_N acc = tl.reshape(accumulator, (BLOCK_SIZE_M, 2, BLOCK_SIZE_N // 2)) acc = tl.permute(acc, (0, 2, 1)) acc0, acc1 = tl.split(acc) c0 = acc0.to(dtype) if not FORWARD: c0_pre = aux_desc.load([offs_am_c, offs_bn_c]) c0 = 2 * c0 * tl.where(c0_pre > 0, c0_pre, 0) c_desc.store([offs_am_c, offs_bn_c], c0) if FORWARD: c0_post = tl.maximum(c0, 0) c0_post = c0_post * c0_post aux_desc.store([offs_am_c, offs_bn_c], c0_post) c1 = acc1.to(dtype) if not FORWARD: c1_pre = aux_desc.load([offs_am_c, offs_bn_c + BLOCK_SIZE_N // 2]) c1 = 2 * c1 * tl.where(c1_pre > 0, c1_pre, 0) c_desc.store([offs_am_c, offs_bn_c + BLOCK_SIZE_N // 2], c1) if FORWARD: c1_post = tl.maximum(c1, 0) c1_post = c1_post * c1_post aux_desc.store([offs_am_c, offs_bn_c + BLOCK_SIZE_N // 2], c1_post) def linear_relu_square(a, b, aux=None): M, K = a.shape N, K = b.shape dtype = a.dtype c = torch.empty((M, N), device=a.device, dtype=dtype) FORWARD = False if aux is None: FORWARD = True aux = torch.empty((M, N), device=a.device, dtype=dtype) NUM_SMS = torch.cuda.get_device_properties("cuda").multi_processor_count BLOCK_SIZE_M = 128 BLOCK_SIZE_N = 256 BLOCK_SIZE_K = 64 num_stages = 4 if FORWARD else 3 num_warps = 8 a_desc = TensorDescriptor.from_tensor(a, [BLOCK_SIZE_M, BLOCK_SIZE_K]) b_desc = TensorDescriptor.from_tensor(b, [BLOCK_SIZE_N, BLOCK_SIZE_K]) c_desc = TensorDescriptor.from_tensor(c, [BLOCK_SIZE_M, BLOCK_SIZE_N // 2]) aux_desc = TensorDescriptor.from_tensor(aux, [BLOCK_SIZE_M, BLOCK_SIZE_N // 2]) def grid(META): return (min( NUM_SMS, triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N), ), ) linear_relu_square_kernel[grid]( a_desc, b_desc, c_desc, aux_desc, M, N, K, BLOCK_SIZE_M=BLOCK_SIZE_M, BLOCK_SIZE_N=BLOCK_SIZE_N, BLOCK_SIZE_K=BLOCK_SIZE_K, GROUP_SIZE_M=1, NUM_SMS=NUM_SMS, FORWARD=FORWARD, num_stages=num_stages, num_warps=num_warps ) if FORWARD: return c, aux else: return c class FusedLinearReLUSquareFunction(torch.autograd.Function): @staticmethod def forward(ctx, x, W1, W2): pre, post = linear_relu_square(x.view((-1, x.shape[-1])), W1) x3 = post @ W2 ctx.save_for_backward(x, W1, W2, pre, post) return x3.view(x.shape) @staticmethod def backward(ctx, grad_output): x, W1, W2, pre, post = ctx.saved_tensors dW2 = post.T @ grad_output dpre = linear_relu_square(grad_output.view((-1, grad_output.shape[-1])), W2, aux=pre) dW1 = dpre.T @ x dx = dpre @ W1 return dx.view(x.shape), dW1, dW2 # ----------------------------------------------------------------------------- # Fused Softcapped Cross Entropy @triton.jit def fused_softcapped_entropy_fwd_kernel( logits_ptr, losses_ptr, lse_ptr, targets_ptr, mtp_weights_ptr, stride_logits_n, stride_logits_v, n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE: tl.constexpr ): row_idx = tl.program_id(0).to(tl.int64) logits_row_ptr = logits_ptr + row_idx * stride_logits_n max_val = -float('inf') sum_exp = 0.0 for off in range(0, n_cols, BLOCK_SIZE): cols = off + tl.arange(0, BLOCK_SIZE) mask = cols < n_cols val = tl.load(logits_row_ptr + cols, mask=mask, other=-float('inf')).to(tl.float32) z = A * tl.sigmoid((val + B) / C) z = tl.where(mask, z, -float('inf')) curr_max = tl.max(z, axis=0) new_max = tl.maximum(max_val, curr_max) sum_exp = sum_exp * tl.exp(max_val - new_max) + tl.sum(tl.exp(z - new_max), axis=0) max_val = new_max lse = max_val + tl.log(sum_exp) tl.store(lse_ptr + row_idx, lse) total_loss = 0.0 for k in range(n_predict): target_idx = row_idx + k if target_idx < n_rows: weight = tl.load(mtp_weights_ptr + k) if weight > 0: target = tl.load(targets_ptr + target_idx).to(tl.int32) if target >= 0 and target < n_cols: val_target = tl.load(logits_row_ptr + target).to(tl.float32) z_target = A * tl.sigmoid((val_target + B) / C) total_loss += weight * (lse - z_target) tl.store(losses_ptr + row_idx, total_loss) @triton.jit def fused_softcapped_entropy_bwd_kernel( grad_input_ptr, grad_output_ptr, lse_ptr, logits_ptr, targets_ptr, mtp_weights_ptr, stride_logits_n, stride_logits_v, stride_grad_n, stride_grad_v, n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE: tl.constexpr ): row_idx = tl.program_id(0).to(tl.int64) logits_row_ptr = logits_ptr + row_idx * stride_logits_n grad_row_ptr = grad_input_ptr + row_idx * stride_grad_n lse = tl.load(lse_ptr + row_idx) grad_loss = tl.load(grad_output_ptr + row_idx) S_w = 0.0 for k in range(n_predict): if row_idx + k < n_rows: S_w += tl.load(mtp_weights_ptr + k) for off in range(0, n_cols, BLOCK_SIZE): cols = off + tl.arange(0, BLOCK_SIZE) mask = cols < n_cols val = tl.load(logits_row_ptr + cols, mask=mask, other=0.0).to(tl.float32) u = (val + B) / C sigmoid_u = tl.sigmoid(u) z = A * sigmoid_u p = tl.exp(z - lse) term1 = S_w * p term2 = tl.zeros([BLOCK_SIZE], dtype=tl.float32) for k in range(n_predict): if row_idx + k < n_rows: target = tl.load(targets_ptr + row_idx + k).to(tl.int32) weight = tl.load(mtp_weights_ptr + k) term2 += tl.where(cols == target, weight, 0.0) grad_z = grad_loss * (term1 - term2) dz_dx = (1.0 / C) * z * (1.0 - sigmoid_u) grad_x = grad_z * dz_dx tl.store(grad_row_ptr + cols, grad_x.to(tl.bfloat16), mask=mask) class FusedSoftcappedCrossEntropy(torch.autograd.Function): @staticmethod def forward(ctx, logits, targets, mtp_weights, A=23.0, B=5.0, C=7.5): n_rows, n_cols = logits.shape if mtp_weights is None: mtp_weights = torch.tensor([1.0], device=logits.device, dtype=torch.float32) n_predict = mtp_weights.shape[0] losses = torch.empty(n_rows, dtype=torch.float32, device=logits.device) lse = torch.empty(n_rows, dtype=torch.float32, device=logits.device) logits = logits.contiguous() targets = targets.contiguous() mtp_weights = mtp_weights.contiguous() grid = (n_rows,) fused_softcapped_entropy_fwd_kernel[grid]( logits, losses, lse, targets, mtp_weights, logits.stride(0), logits.stride(1), n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE=1024, num_warps=8, num_stages=4 ) ctx.save_for_backward(logits, targets, mtp_weights, lse) ctx.params = (A, B, C) return losses @staticmethod def backward(ctx, grad_output): logits, targets, mtp_weights, lse = ctx.saved_tensors A, B, C = ctx.params n_rows, n_cols = logits.shape n_predict = mtp_weights.shape[0] grad_input = torch.empty((n_rows, n_cols), dtype=torch.bfloat16, device=logits.device) grad_output = grad_output.contiguous() grid = (n_rows,) fused_softcapped_entropy_bwd_kernel[grid]( grad_input, grad_output, lse, logits, targets, mtp_weights, logits.stride(0), logits.stride(1), grad_input.stride(0), grad_input.stride(1), n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE=1024, num_warps=8, num_stages=4 ) return grad_input, None, None, None, None, None ==================================================================================================== Running Python 3.10.12 (main, May 27 2025, 17:12:29) [GCC 11.4.0] Running PyTorch 2.10.0.dev20251210+cu126 compiled for CUDA 12.6 Running Triton version 3.6.0 Mon Jan 26 02:10:09 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 31C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:62:00.0 Off | 0 | | N/A 32C P0 116W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 32C P0 115W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 31C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:6A:00.0 Off | 0 | | N/A 31C P0 115W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 33C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 33C P0 118W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 32C P0 115W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 230448 C /usr/bin/python3 1510MiB | | 1 N/A N/A 230449 C /usr/bin/python3 1510MiB | | 2 N/A N/A 230450 C /usr/bin/python3 1510MiB | | 3 N/A N/A 230451 C /usr/bin/python3 1510MiB | | 4 N/A N/A 230452 C /usr/bin/python3 1510MiB | | 5 N/A N/A 230453 C /usr/bin/python3 1510MiB | | 6 N/A N/A 230454 C /usr/bin/python3 1510MiB | | 7 N/A N/A 230455 C /usr/bin/python3 1510MiB | +-----------------------------------------------------------------------------------------+ ==================================================================================================== Compiling model and warming up kernels (~7 minutes on first execution) Sampling steps [0, 1, 2, 511, 512, 513, 1023, 1024, 1025, 1534, 1535, 1536] for warmup Resetting Model step:0/1575 val_loss:10.8326 train_time:0ms step_avg:0.04ms step:1/1575 train_time:76ms step_avg:75.56ms step:2/1575 train_time:98ms step_avg:48.85ms step:3/1575 train_time:120ms step_avg:40.15ms step:4/1575 train_time:143ms step_avg:35.74ms step:5/1575 train_time:174ms step_avg:34.70ms step:6/1575 train_time:275ms step_avg:45.88ms step:7/1575 train_time:293ms step_avg:41.86ms step:8/1575 train_time:315ms step_avg:39.36ms step:9/1575 train_time:346ms step_avg:38.41ms step:10/1575 train_time:384ms step_avg:38.39ms step:11/1575 train_time:415ms step_avg:37.72ms step:12/1575 train_time:453ms step_avg:37.75ms step:13/1575 train_time:484ms step_avg:37.24ms step:14/1575 train_time:523ms step_avg:37.36ms step:15/1575 train_time:555ms step_avg:36.97ms step:16/1575 train_time:593ms step_avg:37.06ms step:17/1575 train_time:624ms step_avg:36.72ms step:18/1575 train_time:664ms step_avg:36.87ms step:19/1575 train_time:694ms step_avg:36.54ms step:20/1575 train_time:733ms step_avg:36.63ms step:21/1575 train_time:764ms step_avg:36.38ms step:22/1575 train_time:803ms step_avg:36.50ms step:23/1575 train_time:834ms step_avg:36.27ms step:24/1575 train_time:873ms step_avg:36.37ms step:25/1575 train_time:904ms step_avg:36.14ms step:26/1575 train_time:942ms step_avg:36.23ms step:27/1575 train_time:973ms step_avg:36.04ms step:28/1575 train_time:1011ms step_avg:36.12ms step:29/1575 train_time:1043ms step_avg:35.95ms step:30/1575 train_time:1081ms step_avg:36.04ms step:31/1575 train_time:1112ms step_avg:35.88ms step:32/1575 train_time:1151ms step_avg:35.96ms step:33/1575 train_time:1182ms step_avg:35.81ms step:34/1575 train_time:1220ms step_avg:35.89ms step:35/1575 train_time:1252ms step_avg:35.76ms step:36/1575 train_time:1290ms step_avg:35.83ms step:37/1575 train_time:1321ms step_avg:35.70ms step:38/1575 train_time:1360ms step_avg:35.79ms step:39/1575 train_time:1391ms step_avg:35.67ms step:40/1575 train_time:1430ms step_avg:35.74ms step:41/1575 train_time:1461ms step_avg:35.64ms step:42/1575 train_time:1501ms step_avg:35.73ms step:43/1575 train_time:1532ms step_avg:35.62ms step:44/1575 train_time:1571ms step_avg:35.70ms step:45/1575 train_time:1602ms step_avg:35.59ms step:46/1575 train_time:1640ms step_avg:35.66ms step:47/1575 train_time:1671ms step_avg:35.55ms step:48/1575 train_time:1710ms step_avg:35.62ms step:49/1575 train_time:1741ms step_avg:35.52ms step:50/1575 train_time:1780ms step_avg:35.60ms step:51/1575 train_time:1811ms step_avg:35.51ms step:52/1575 train_time:1850ms step_avg:35.57ms step:53/1575 train_time:1881ms step_avg:35.49ms step:54/1575 train_time:1920ms step_avg:35.56ms step:55/1575 train_time:1951ms step_avg:35.48ms step:56/1575 train_time:1990ms step_avg:35.53ms step:57/1575 train_time:2021ms step_avg:35.45ms step:58/1575 train_time:2060ms step_avg:35.52ms step:59/1575 train_time:2092ms step_avg:35.46ms step:60/1575 train_time:2130ms step_avg:35.51ms step:61/1575 train_time:2161ms step_avg:35.43ms step:62/1575 train_time:2200ms step_avg:35.49ms step:63/1575 train_time:2232ms step_avg:35.43ms step:64/1575 train_time:2271ms step_avg:35.48ms step:65/1575 train_time:2302ms step_avg:35.41ms step:66/1575 train_time:2341ms step_avg:35.47ms step:67/1575 train_time:2372ms step_avg:35.40ms step:68/1575 train_time:2411ms step_avg:35.45ms step:69/1575 train_time:2442ms step_avg:35.39ms step:70/1575 train_time:2481ms step_avg:35.44ms step:71/1575 train_time:2512ms step_avg:35.38ms step:72/1575 train_time:2550ms step_avg:35.42ms step:73/1575 train_time:2582ms step_avg:35.37ms step:74/1575 train_time:2621ms step_avg:35.41ms step:75/1575 train_time:2652ms step_avg:35.35ms step:76/1575 train_time:2690ms step_avg:35.40ms step:77/1575 train_time:2722ms step_avg:35.34ms step:78/1575 train_time:2761ms step_avg:35.40ms step:79/1575 train_time:2792ms step_avg:35.34ms step:80/1575 train_time:2834ms step_avg:35.42ms step:81/1575 train_time:2861ms step_avg:35.32ms step:82/1575 train_time:2900ms step_avg:35.37ms step:83/1575 train_time:2931ms step_avg:35.32ms step:84/1575 train_time:2970ms step_avg:35.36ms step:85/1575 train_time:3001ms step_avg:35.30ms step:86/1575 train_time:3040ms step_avg:35.35ms step:87/1575 train_time:3071ms step_avg:35.30ms step:88/1575 train_time:3110ms step_avg:35.34ms step:89/1575 train_time:3141ms step_avg:35.29ms step:90/1575 train_time:3180ms step_avg:35.33ms step:91/1575 train_time:3211ms step_avg:35.28ms step:92/1575 train_time:3249ms step_avg:35.31ms step:93/1575 train_time:3280ms step_avg:35.27ms step:94/1575 train_time:3319ms step_avg:35.31ms step:95/1575 train_time:3350ms step_avg:35.26ms step:96/1575 train_time:3389ms step_avg:35.30ms step:97/1575 train_time:3420ms step_avg:35.25ms step:98/1575 train_time:3458ms step_avg:35.29ms step:99/1575 train_time:3489ms step_avg:35.24ms step:100/1575 train_time:3527ms step_avg:35.27ms step:101/1575 train_time:3558ms step_avg:35.23ms step:102/1575 train_time:3597ms step_avg:35.27ms step:103/1575 train_time:3628ms step_avg:35.23ms step:104/1575 train_time:3667ms step_avg:35.26ms step:105/1575 train_time:3698ms step_avg:35.22ms step:106/1575 train_time:3736ms step_avg:35.25ms step:107/1575 train_time:3768ms step_avg:35.21ms step:108/1575 train_time:3806ms step_avg:35.24ms step:109/1575 train_time:3837ms step_avg:35.21ms step:110/1575 train_time:3876ms step_avg:35.24ms step:111/1575 train_time:3907ms step_avg:35.20ms step:112/1575 train_time:3945ms step_avg:35.23ms step:113/1575 train_time:3977ms step_avg:35.19ms step:114/1575 train_time:4015ms step_avg:35.22ms step:115/1575 train_time:4046ms step_avg:35.18ms step:116/1575 train_time:4085ms step_avg:35.21ms step:117/1575 train_time:4115ms step_avg:35.17ms step:118/1575 train_time:4153ms step_avg:35.20ms step:119/1575 train_time:4185ms step_avg:35.17ms step:120/1575 train_time:4224ms step_avg:35.20ms step:121/1575 train_time:4255ms step_avg:35.16ms step:122/1575 train_time:4293ms step_avg:35.19ms step:123/1575 train_time:4324ms step_avg:35.15ms step:124/1575 train_time:4363ms step_avg:35.18ms step:125/1575 train_time:4393ms step_avg:35.15ms step:126/1575 train_time:4431ms step_avg:35.17ms step:127/1575 train_time:4462ms step_avg:35.14ms step:128/1575 train_time:4501ms step_avg:35.16ms step:129/1575 train_time:4532ms step_avg:35.13ms step:130/1575 train_time:4571ms step_avg:35.16ms step:131/1575 train_time:4602ms step_avg:35.13ms step:132/1575 train_time:4641ms step_avg:35.16ms step:133/1575 train_time:4672ms step_avg:35.13ms step:134/1575 train_time:4710ms step_avg:35.15ms step:135/1575 train_time:4741ms step_avg:35.12ms step:136/1575 train_time:4780ms step_avg:35.15ms step:137/1575 train_time:4811ms step_avg:35.12ms step:138/1575 train_time:4849ms step_avg:35.14ms step:139/1575 train_time:4880ms step_avg:35.11ms step:140/1575 train_time:4919ms step_avg:35.14ms step:141/1575 train_time:4950ms step_avg:35.11ms step:142/1575 train_time:4989ms step_avg:35.13ms step:143/1575 train_time:5019ms step_avg:35.10ms step:144/1575 train_time:5058ms step_avg:35.12ms step:145/1575 train_time:5089ms step_avg:35.10ms step:146/1575 train_time:5127ms step_avg:35.12ms step:147/1575 train_time:5158ms step_avg:35.09ms step:148/1575 train_time:5197ms step_avg:35.12ms step:149/1575 train_time:5228ms step_avg:35.09ms step:150/1575 train_time:5267ms step_avg:35.11ms step:151/1575 train_time:5298ms step_avg:35.08ms step:152/1575 train_time:5336ms step_avg:35.11ms step:153/1575 train_time:5368ms step_avg:35.08ms step:154/1575 train_time:5406ms step_avg:35.10ms step:155/1575 train_time:5437ms step_avg:35.08ms step:156/1575 train_time:5475ms step_avg:35.10ms step:157/1575 train_time:5506ms step_avg:35.07ms step:158/1575 train_time:5545ms step_avg:35.09ms step:159/1575 train_time:5575ms step_avg:35.07ms step:160/1575 train_time:5614ms step_avg:35.09ms step:161/1575 train_time:5645ms step_avg:35.06ms step:162/1575 train_time:5684ms step_avg:35.09ms step:163/1575 train_time:5715ms step_avg:35.06ms step:164/1575 train_time:5753ms step_avg:35.08ms step:165/1575 train_time:5784ms step_avg:35.05ms step:166/1575 train_time:5823ms step_avg:35.08ms step:167/1575 train_time:5854ms step_avg:35.05ms step:168/1575 train_time:5892ms step_avg:35.07ms step:169/1575 train_time:5923ms step_avg:35.05ms step:170/1575 train_time:5962ms step_avg:35.07ms step:171/1575 train_time:5993ms step_avg:35.05ms step:172/1575 train_time:6031ms step_avg:35.06ms step:173/1575 train_time:6062ms step_avg:35.04ms step:174/1575 train_time:6101ms step_avg:35.06ms step:175/1575 train_time:6132ms step_avg:35.04ms step:176/1575 train_time:6170ms step_avg:35.06ms step:177/1575 train_time:6201ms step_avg:35.03ms step:178/1575 train_time:6240ms step_avg:35.06ms step:179/1575 train_time:6271ms step_avg:35.03ms step:180/1575 train_time:6309ms step_avg:35.05ms step:181/1575 train_time:6340ms step_avg:35.03ms step:182/1575 train_time:6379ms step_avg:35.05ms step:183/1575 train_time:6410ms step_avg:35.03ms step:184/1575 train_time:6449ms step_avg:35.05ms step:185/1575 train_time:6480ms step_avg:35.03ms step:186/1575 train_time:6518ms step_avg:35.04ms step:187/1575 train_time:6549ms step_avg:35.02ms step:188/1575 train_time:6588ms step_avg:35.04ms step:189/1575 train_time:6619ms step_avg:35.02ms step:190/1575 train_time:6657ms step_avg:35.04ms step:191/1575 train_time:6688ms step_avg:35.02ms step:192/1575 train_time:6727ms step_avg:35.03ms step:193/1575 train_time:6757ms step_avg:35.01ms step:194/1575 train_time:6796ms step_avg:35.03ms step:195/1575 train_time:6827ms step_avg:35.01ms step:196/1575 train_time:6866ms step_avg:35.03ms step:197/1575 train_time:6896ms step_avg:35.01ms step:198/1575 train_time:6935ms step_avg:35.02ms step:199/1575 train_time:6966ms step_avg:35.00ms step:200/1575 train_time:7004ms step_avg:35.02ms step:201/1575 train_time:7035ms step_avg:35.00ms step:202/1575 train_time:7073ms step_avg:35.02ms step:203/1575 train_time:7104ms step_avg:35.00ms step:204/1575 train_time:7142ms step_avg:35.01ms step:205/1575 train_time:7173ms step_avg:34.99ms step:206/1575 train_time:7212ms step_avg:35.01ms step:207/1575 train_time:7243ms step_avg:34.99ms step:208/1575 train_time:7282ms step_avg:35.01ms step:209/1575 train_time:7314ms step_avg:34.99ms step:210/1575 train_time:7352ms step_avg:35.01ms step:211/1575 train_time:7383ms step_avg:34.99ms step:212/1575 train_time:7422ms step_avg:35.01ms step:213/1575 train_time:7453ms step_avg:34.99ms step:214/1575 train_time:7491ms step_avg:35.00ms step:215/1575 train_time:7522ms step_avg:34.99ms step:216/1575 train_time:7561ms step_avg:35.00ms step:217/1575 train_time:7592ms step_avg:34.99ms step:218/1575 train_time:7631ms step_avg:35.00ms step:219/1575 train_time:7661ms step_avg:34.98ms step:220/1575 train_time:7700ms step_avg:35.00ms step:221/1575 train_time:7732ms step_avg:34.99ms step:222/1575 train_time:7770ms step_avg:35.00ms step:223/1575 train_time:7802ms step_avg:34.98ms step:224/1575 train_time:7841ms step_avg:35.00ms step:225/1575 train_time:7871ms step_avg:34.98ms step:226/1575 train_time:7910ms step_avg:35.00ms step:227/1575 train_time:7941ms step_avg:34.98ms step:228/1575 train_time:7979ms step_avg:35.00ms step:229/1575 train_time:8010ms step_avg:34.98ms step:230/1575 train_time:8049ms step_avg:34.99ms step:231/1575 train_time:8080ms step_avg:34.98ms step:232/1575 train_time:8119ms step_avg:35.00ms step:233/1575 train_time:8150ms step_avg:34.98ms step:234/1575 train_time:8188ms step_avg:34.99ms step:235/1575 train_time:8219ms step_avg:34.97ms step:236/1575 train_time:8257ms step_avg:34.99ms step:237/1575 train_time:8288ms step_avg:34.97ms step:238/1575 train_time:8327ms step_avg:34.99ms step:239/1575 train_time:8358ms step_avg:34.97ms step:240/1575 train_time:8396ms step_avg:34.99ms step:241/1575 train_time:8427ms step_avg:34.97ms step:242/1575 train_time:8466ms step_avg:34.98ms step:243/1575 train_time:8497ms step_avg:34.97ms step:244/1575 train_time:8535ms step_avg:34.98ms step:245/1575 train_time:8566ms step_avg:34.96ms step:246/1575 train_time:8605ms step_avg:34.98ms step:247/1575 train_time:8636ms step_avg:34.96ms step:248/1575 train_time:8674ms step_avg:34.98ms step:249/1575 train_time:8705ms step_avg:34.96ms step:250/1575 train_time:8743ms step_avg:34.97ms step:250/1575 val_loss:4.5841 train_time:8792ms step_avg:35.17ms step:251/1575 train_time:8811ms step_avg:35.10ms step:252/1575 train_time:8831ms step_avg:35.04ms step:253/1575 train_time:8848ms step_avg:34.97ms step:254/1575 train_time:8885ms step_avg:34.98ms step:255/1575 train_time:8917ms step_avg:34.97ms step:256/1575 train_time:8956ms step_avg:34.98ms step:257/1575 train_time:8987ms step_avg:34.97ms step:258/1575 train_time:9026ms step_avg:34.99ms step:259/1575 train_time:9057ms step_avg:34.97ms step:260/1575 train_time:9095ms step_avg:34.98ms step:261/1575 train_time:9126ms step_avg:34.97ms step:262/1575 train_time:9166ms step_avg:34.98ms step:263/1575 train_time:9196ms step_avg:34.96ms step:264/1575 train_time:9234ms step_avg:34.98ms step:265/1575 train_time:9265ms step_avg:34.96ms step:266/1575 train_time:9303ms step_avg:34.97ms step:267/1575 train_time:9334ms step_avg:34.96ms step:268/1575 train_time:9372ms step_avg:34.97ms step:269/1575 train_time:9403ms step_avg:34.96ms step:270/1575 train_time:9442ms step_avg:34.97ms step:271/1575 train_time:9472ms step_avg:34.95ms step:272/1575 train_time:9511ms step_avg:34.97ms step:273/1575 train_time:9542ms step_avg:34.95ms step:274/1575 train_time:9581ms step_avg:34.97ms step:275/1575 train_time:9612ms step_avg:34.95ms step:276/1575 train_time:9650ms step_avg:34.96ms step:277/1575 train_time:9681ms step_avg:34.95ms step:278/1575 train_time:9720ms step_avg:34.96ms step:279/1575 train_time:9751ms step_avg:34.95ms step:280/1575 train_time:9790ms step_avg:34.96ms step:281/1575 train_time:9821ms step_avg:34.95ms step:282/1575 train_time:9860ms step_avg:34.96ms step:283/1575 train_time:9891ms step_avg:34.95ms step:284/1575 train_time:9930ms step_avg:34.96ms step:285/1575 train_time:9960ms step_avg:34.95ms step:286/1575 train_time:9999ms step_avg:34.96ms step:287/1575 train_time:10030ms step_avg:34.95ms step:288/1575 train_time:10068ms step_avg:34.96ms step:289/1575 train_time:10099ms step_avg:34.95ms step:290/1575 train_time:10138ms step_avg:34.96ms step:291/1575 train_time:10169ms step_avg:34.94ms step:292/1575 train_time:10208ms step_avg:34.96ms step:293/1575 train_time:10239ms step_avg:34.94ms step:294/1575 train_time:10277ms step_avg:34.96ms step:295/1575 train_time:10308ms step_avg:34.94ms step:296/1575 train_time:10347ms step_avg:34.95ms step:297/1575 train_time:10377ms step_avg:34.94ms step:298/1575 train_time:10416ms step_avg:34.95ms step:299/1575 train_time:10446ms step_avg:34.94ms step:300/1575 train_time:10485ms step_avg:34.95ms step:301/1575 train_time:10516ms step_avg:34.94ms step:302/1575 train_time:10554ms step_avg:34.95ms step:303/1575 train_time:10585ms step_avg:34.93ms step:304/1575 train_time:10623ms step_avg:34.94ms step:305/1575 train_time:10654ms step_avg:34.93ms step:306/1575 train_time:10693ms step_avg:34.94ms step:307/1575 train_time:10723ms step_avg:34.93ms step:308/1575 train_time:10762ms step_avg:34.94ms step:309/1575 train_time:10793ms step_avg:34.93ms step:310/1575 train_time:10831ms step_avg:34.94ms step:311/1575 train_time:10862ms step_avg:34.93ms step:312/1575 train_time:10901ms step_avg:34.94ms step:313/1575 train_time:10932ms step_avg:34.93ms step:314/1575 train_time:10971ms step_avg:34.94ms step:315/1575 train_time:11002ms step_avg:34.93ms step:316/1575 train_time:11040ms step_avg:34.94ms step:317/1575 train_time:11071ms step_avg:34.93ms step:318/1575 train_time:11110ms step_avg:34.94ms step:319/1575 train_time:11142ms step_avg:34.93ms step:320/1575 train_time:11180ms step_avg:34.94ms step:321/1575 train_time:11211ms step_avg:34.93ms step:322/1575 train_time:11250ms step_avg:34.94ms step:323/1575 train_time:11281ms step_avg:34.93ms step:324/1575 train_time:11320ms step_avg:34.94ms step:325/1575 train_time:11351ms step_avg:34.93ms step:326/1575 train_time:11389ms step_avg:34.94ms step:327/1575 train_time:11420ms step_avg:34.92ms step:328/1575 train_time:11459ms step_avg:34.94ms step:329/1575 train_time:11491ms step_avg:34.93ms step:330/1575 train_time:11529ms step_avg:34.94ms step:331/1575 train_time:11560ms step_avg:34.93ms step:332/1575 train_time:11599ms step_avg:34.94ms step:333/1575 train_time:11631ms step_avg:34.93ms step:334/1575 train_time:11669ms step_avg:34.94ms step:335/1575 train_time:11700ms step_avg:34.93ms step:336/1575 train_time:11739ms step_avg:34.94ms step:337/1575 train_time:11770ms step_avg:34.93ms step:338/1575 train_time:11808ms step_avg:34.94ms step:339/1575 train_time:11839ms step_avg:34.92ms step:340/1575 train_time:11877ms step_avg:34.93ms step:341/1575 train_time:11908ms step_avg:34.92ms step:342/1575 train_time:11947ms step_avg:34.93ms step:343/1575 train_time:11978ms step_avg:34.92ms step:344/1575 train_time:12017ms step_avg:34.93ms step:345/1575 train_time:12047ms step_avg:34.92ms step:346/1575 train_time:12086ms step_avg:34.93ms step:347/1575 train_time:12116ms step_avg:34.92ms step:348/1575 train_time:12156ms step_avg:34.93ms step:349/1575 train_time:12185ms step_avg:34.92ms step:350/1575 train_time:12224ms step_avg:34.93ms step:351/1575 train_time:12255ms step_avg:34.91ms step:352/1575 train_time:12293ms step_avg:34.92ms step:353/1575 train_time:12324ms step_avg:34.91ms step:354/1575 train_time:12362ms step_avg:34.92ms step:355/1575 train_time:12393ms step_avg:34.91ms step:356/1575 train_time:12431ms step_avg:34.92ms step:357/1575 train_time:12462ms step_avg:34.91ms step:358/1575 train_time:12501ms step_avg:34.92ms step:359/1575 train_time:12532ms step_avg:34.91ms step:360/1575 train_time:12570ms step_avg:34.92ms step:361/1575 train_time:12601ms step_avg:34.91ms step:362/1575 train_time:12640ms step_avg:34.92ms step:363/1575 train_time:12670ms step_avg:34.90ms step:364/1575 train_time:12709ms step_avg:34.91ms step:365/1575 train_time:12740ms step_avg:34.90ms step:366/1575 train_time:12778ms step_avg:34.91ms step:367/1575 train_time:12809ms step_avg:34.90ms step:368/1575 train_time:12848ms step_avg:34.91ms step:369/1575 train_time:12878ms step_avg:34.90ms step:370/1575 train_time:12917ms step_avg:34.91ms step:371/1575 train_time:12948ms step_avg:34.90ms step:372/1575 train_time:12986ms step_avg:34.91ms step:373/1575 train_time:13017ms step_avg:34.90ms step:374/1575 train_time:13055ms step_avg:34.91ms step:375/1575 train_time:13086ms step_avg:34.90ms step:376/1575 train_time:13125ms step_avg:34.91ms step:377/1575 train_time:13156ms step_avg:34.90ms step:378/1575 train_time:13194ms step_avg:34.90ms step:379/1575 train_time:13225ms step_avg:34.89ms step:380/1575 train_time:13263ms step_avg:34.90ms step:381/1575 train_time:13294ms step_avg:34.89ms step:382/1575 train_time:13332ms step_avg:34.90ms step:383/1575 train_time:13363ms step_avg:34.89ms step:384/1575 train_time:13402ms step_avg:34.90ms step:385/1575 train_time:13433ms step_avg:34.89ms step:386/1575 train_time:13472ms step_avg:34.90ms step:387/1575 train_time:13503ms step_avg:34.89ms step:388/1575 train_time:13541ms step_avg:34.90ms step:389/1575 train_time:13572ms step_avg:34.89ms step:390/1575 train_time:13610ms step_avg:34.90ms step:391/1575 train_time:13641ms step_avg:34.89ms step:392/1575 train_time:13680ms step_avg:34.90ms step:393/1575 train_time:13711ms step_avg:34.89ms step:394/1575 train_time:13749ms step_avg:34.90ms step:395/1575 train_time:13780ms step_avg:34.89ms step:396/1575 train_time:13818ms step_avg:34.89ms step:397/1575 train_time:13849ms step_avg:34.88ms step:398/1575 train_time:13887ms step_avg:34.89ms step:399/1575 train_time:13919ms step_avg:34.88ms step:400/1575 train_time:13957ms step_avg:34.89ms step:401/1575 train_time:13988ms step_avg:34.88ms step:402/1575 train_time:14027ms step_avg:34.89ms step:403/1575 train_time:14058ms step_avg:34.88ms step:404/1575 train_time:14096ms step_avg:34.89ms step:405/1575 train_time:14127ms step_avg:34.88ms step:406/1575 train_time:14166ms step_avg:34.89ms step:407/1575 train_time:14197ms step_avg:34.88ms step:408/1575 train_time:14235ms step_avg:34.89ms step:409/1575 train_time:14266ms step_avg:34.88ms step:410/1575 train_time:14304ms step_avg:34.89ms step:411/1575 train_time:14335ms step_avg:34.88ms step:412/1575 train_time:14374ms step_avg:34.89ms step:413/1575 train_time:14404ms step_avg:34.88ms step:414/1575 train_time:14443ms step_avg:34.89ms step:415/1575 train_time:14474ms step_avg:34.88ms step:416/1575 train_time:14512ms step_avg:34.89ms step:417/1575 train_time:14543ms step_avg:34.88ms step:418/1575 train_time:14581ms step_avg:34.88ms step:419/1575 train_time:14612ms step_avg:34.87ms step:420/1575 train_time:14651ms step_avg:34.88ms step:421/1575 train_time:14681ms step_avg:34.87ms step:422/1575 train_time:14720ms step_avg:34.88ms step:423/1575 train_time:14751ms step_avg:34.87ms step:424/1575 train_time:14789ms step_avg:34.88ms step:425/1575 train_time:14820ms step_avg:34.87ms step:426/1575 train_time:14859ms step_avg:34.88ms step:427/1575 train_time:14891ms step_avg:34.87ms step:428/1575 train_time:14930ms step_avg:34.88ms step:429/1575 train_time:14961ms step_avg:34.87ms step:430/1575 train_time:15000ms step_avg:34.88ms step:431/1575 train_time:15032ms step_avg:34.88ms step:432/1575 train_time:15070ms step_avg:34.88ms step:433/1575 train_time:15101ms step_avg:34.88ms step:434/1575 train_time:15140ms step_avg:34.89ms step:435/1575 train_time:15172ms step_avg:34.88ms step:436/1575 train_time:15210ms step_avg:34.89ms step:437/1575 train_time:15241ms step_avg:34.88ms step:438/1575 train_time:15280ms step_avg:34.89ms step:439/1575 train_time:15311ms step_avg:34.88ms step:440/1575 train_time:15350ms step_avg:34.89ms step:441/1575 train_time:15380ms step_avg:34.88ms step:442/1575 train_time:15419ms step_avg:34.88ms step:443/1575 train_time:15450ms step_avg:34.88ms step:444/1575 train_time:15489ms step_avg:34.88ms step:445/1575 train_time:15520ms step_avg:34.88ms step:446/1575 train_time:15558ms step_avg:34.88ms step:447/1575 train_time:15589ms step_avg:34.88ms step:448/1575 train_time:15628ms step_avg:34.88ms step:449/1575 train_time:15659ms step_avg:34.87ms step:450/1575 train_time:15697ms step_avg:34.88ms step:451/1575 train_time:15728ms step_avg:34.87ms step:452/1575 train_time:15766ms step_avg:34.88ms step:453/1575 train_time:15797ms step_avg:34.87ms step:454/1575 train_time:15835ms step_avg:34.88ms step:455/1575 train_time:15866ms step_avg:34.87ms step:456/1575 train_time:15905ms step_avg:34.88ms step:457/1575 train_time:15936ms step_avg:34.87ms step:458/1575 train_time:15974ms step_avg:34.88ms step:459/1575 train_time:16005ms step_avg:34.87ms step:460/1575 train_time:16044ms step_avg:34.88ms step:461/1575 train_time:16074ms step_avg:34.87ms step:462/1575 train_time:16113ms step_avg:34.88ms step:463/1575 train_time:16144ms step_avg:34.87ms step:464/1575 train_time:16183ms step_avg:34.88ms step:465/1575 train_time:16214ms step_avg:34.87ms step:466/1575 train_time:16252ms step_avg:34.88ms step:467/1575 train_time:16283ms step_avg:34.87ms step:468/1575 train_time:16321ms step_avg:34.87ms step:469/1575 train_time:16352ms step_avg:34.87ms step:470/1575 train_time:16391ms step_avg:34.87ms step:471/1575 train_time:16422ms step_avg:34.87ms step:472/1575 train_time:16460ms step_avg:34.87ms step:473/1575 train_time:16491ms step_avg:34.87ms step:474/1575 train_time:16529ms step_avg:34.87ms step:475/1575 train_time:16561ms step_avg:34.86ms step:476/1575 train_time:16599ms step_avg:34.87ms step:477/1575 train_time:16631ms step_avg:34.87ms step:478/1575 train_time:16670ms step_avg:34.87ms step:479/1575 train_time:16700ms step_avg:34.87ms step:480/1575 train_time:16739ms step_avg:34.87ms step:481/1575 train_time:16769ms step_avg:34.86ms step:482/1575 train_time:16808ms step_avg:34.87ms step:483/1575 train_time:16839ms step_avg:34.86ms step:484/1575 train_time:16878ms step_avg:34.87ms step:485/1575 train_time:16908ms step_avg:34.86ms step:486/1575 train_time:16947ms step_avg:34.87ms step:487/1575 train_time:16978ms step_avg:34.86ms step:488/1575 train_time:17016ms step_avg:34.87ms step:489/1575 train_time:17047ms step_avg:34.86ms step:490/1575 train_time:17086ms step_avg:34.87ms step:491/1575 train_time:17116ms step_avg:34.86ms step:492/1575 train_time:17155ms step_avg:34.87ms step:493/1575 train_time:17185ms step_avg:34.86ms step:494/1575 train_time:17224ms step_avg:34.87ms step:495/1575 train_time:17255ms step_avg:34.86ms step:496/1575 train_time:17294ms step_avg:34.87ms step:497/1575 train_time:17324ms step_avg:34.86ms step:498/1575 train_time:17363ms step_avg:34.87ms step:499/1575 train_time:17394ms step_avg:34.86ms step:500/1575 train_time:17432ms step_avg:34.86ms step:500/1575 val_loss:4.2451 train_time:17481ms step_avg:34.96ms step:501/1575 train_time:17500ms step_avg:34.93ms step:502/1575 train_time:17520ms step_avg:34.90ms step:503/1575 train_time:17537ms step_avg:34.87ms step:504/1575 train_time:17575ms step_avg:34.87ms step:505/1575 train_time:17608ms step_avg:34.87ms step:506/1575 train_time:17648ms step_avg:34.88ms step:507/1575 train_time:17680ms step_avg:34.87ms step:508/1575 train_time:17719ms step_avg:34.88ms step:509/1575 train_time:17750ms step_avg:34.87ms step:510/1575 train_time:17788ms step_avg:34.88ms step:511/1575 train_time:17819ms step_avg:34.87ms step:512/1575 train_time:17857ms step_avg:34.88ms step:513/1575 train_time:17929ms step_avg:34.95ms step:514/1575 train_time:17987ms step_avg:34.99ms step:515/1575 train_time:18048ms step_avg:35.05ms step:516/1575 train_time:18108ms step_avg:35.09ms step:517/1575 train_time:18170ms step_avg:35.15ms step:518/1575 train_time:18229ms step_avg:35.19ms step:519/1575 train_time:18291ms step_avg:35.24ms step:520/1575 train_time:18350ms step_avg:35.29ms step:521/1575 train_time:18413ms step_avg:35.34ms step:522/1575 train_time:18473ms step_avg:35.39ms step:523/1575 train_time:18538ms step_avg:35.44ms step:524/1575 train_time:18599ms step_avg:35.49ms step:525/1575 train_time:18663ms step_avg:35.55ms step:526/1575 train_time:18724ms step_avg:35.60ms step:527/1575 train_time:18787ms step_avg:35.65ms step:528/1575 train_time:18848ms step_avg:35.70ms step:529/1575 train_time:18912ms step_avg:35.75ms step:530/1575 train_time:18971ms step_avg:35.79ms step:531/1575 train_time:19033ms step_avg:35.84ms step:532/1575 train_time:19094ms step_avg:35.89ms step:533/1575 train_time:19158ms step_avg:35.94ms step:534/1575 train_time:19219ms step_avg:35.99ms step:535/1575 train_time:19280ms step_avg:36.04ms step:536/1575 train_time:19339ms step_avg:36.08ms step:537/1575 train_time:19402ms step_avg:36.13ms step:538/1575 train_time:19461ms step_avg:36.17ms step:539/1575 train_time:19525ms step_avg:36.22ms step:540/1575 train_time:19586ms step_avg:36.27ms step:541/1575 train_time:19650ms step_avg:36.32ms step:542/1575 train_time:19709ms step_avg:36.36ms step:543/1575 train_time:19773ms step_avg:36.41ms step:544/1575 train_time:19833ms step_avg:36.46ms step:545/1575 train_time:19897ms step_avg:36.51ms step:546/1575 train_time:19957ms step_avg:36.55ms step:547/1575 train_time:20020ms step_avg:36.60ms step:548/1575 train_time:20080ms step_avg:36.64ms step:549/1575 train_time:20144ms step_avg:36.69ms step:550/1575 train_time:20204ms step_avg:36.73ms step:551/1575 train_time:20266ms step_avg:36.78ms step:552/1575 train_time:20325ms step_avg:36.82ms step:553/1575 train_time:20388ms step_avg:36.87ms step:554/1575 train_time:20448ms step_avg:36.91ms step:555/1575 train_time:20512ms step_avg:36.96ms step:556/1575 train_time:20571ms step_avg:37.00ms step:557/1575 train_time:20636ms step_avg:37.05ms step:558/1575 train_time:20696ms step_avg:37.09ms step:559/1575 train_time:20760ms step_avg:37.14ms step:560/1575 train_time:20820ms step_avg:37.18ms step:561/1575 train_time:20883ms step_avg:37.22ms step:562/1575 train_time:20943ms step_avg:37.26ms step:563/1575 train_time:21007ms step_avg:37.31ms step:564/1575 train_time:21067ms step_avg:37.35ms step:565/1575 train_time:21130ms step_avg:37.40ms step:566/1575 train_time:21190ms step_avg:37.44ms step:567/1575 train_time:21253ms step_avg:37.48ms step:568/1575 train_time:21312ms step_avg:37.52ms step:569/1575 train_time:21380ms step_avg:37.57ms step:570/1575 train_time:21437ms step_avg:37.61ms step:571/1575 train_time:21502ms step_avg:37.66ms step:572/1575 train_time:21562ms step_avg:37.70ms step:573/1575 train_time:21622ms step_avg:37.73ms step:574/1575 train_time:21681ms step_avg:37.77ms step:575/1575 train_time:21746ms step_avg:37.82ms step:576/1575 train_time:21806ms step_avg:37.86ms step:577/1575 train_time:21869ms step_avg:37.90ms step:578/1575 train_time:21928ms step_avg:37.94ms step:579/1575 train_time:21991ms step_avg:37.98ms step:580/1575 train_time:22051ms step_avg:38.02ms step:581/1575 train_time:22114ms step_avg:38.06ms step:582/1575 train_time:22173ms step_avg:38.10ms step:583/1575 train_time:22237ms step_avg:38.14ms step:584/1575 train_time:22297ms step_avg:38.18ms step:585/1575 train_time:22360ms step_avg:38.22ms step:586/1575 train_time:22420ms step_avg:38.26ms step:587/1575 train_time:22483ms step_avg:38.30ms step:588/1575 train_time:22542ms step_avg:38.34ms step:589/1575 train_time:22605ms step_avg:38.38ms step:590/1575 train_time:22665ms step_avg:38.41ms step:591/1575 train_time:22729ms step_avg:38.46ms step:592/1575 train_time:22789ms step_avg:38.49ms step:593/1575 train_time:22852ms step_avg:38.54ms step:594/1575 train_time:22912ms step_avg:38.57ms step:595/1575 train_time:22974ms step_avg:38.61ms step:596/1575 train_time:23034ms step_avg:38.65ms step:597/1575 train_time:23098ms step_avg:38.69ms step:598/1575 train_time:23157ms step_avg:38.72ms step:599/1575 train_time:23220ms step_avg:38.77ms step:600/1575 train_time:23279ms step_avg:38.80ms step:601/1575 train_time:23343ms step_avg:38.84ms step:602/1575 train_time:23402ms step_avg:38.87ms step:603/1575 train_time:23465ms step_avg:38.91ms step:604/1575 train_time:23525ms step_avg:38.95ms step:605/1575 train_time:23588ms step_avg:38.99ms step:606/1575 train_time:23648ms step_avg:39.02ms step:607/1575 train_time:23711ms step_avg:39.06ms step:608/1575 train_time:23771ms step_avg:39.10ms step:609/1575 train_time:23834ms step_avg:39.14ms step:610/1575 train_time:23894ms step_avg:39.17ms step:611/1575 train_time:23957ms step_avg:39.21ms step:612/1575 train_time:24017ms step_avg:39.24ms step:613/1575 train_time:24080ms step_avg:39.28ms step:614/1575 train_time:24140ms step_avg:39.32ms step:615/1575 train_time:24202ms step_avg:39.35ms step:616/1575 train_time:24263ms step_avg:39.39ms step:617/1575 train_time:24325ms step_avg:39.42ms step:618/1575 train_time:24385ms step_avg:39.46ms step:619/1575 train_time:24448ms step_avg:39.50ms step:620/1575 train_time:24508ms step_avg:39.53ms step:621/1575 train_time:24570ms step_avg:39.57ms step:622/1575 train_time:24630ms step_avg:39.60ms step:623/1575 train_time:24693ms step_avg:39.64ms step:624/1575 train_time:24753ms step_avg:39.67ms step:625/1575 train_time:24816ms step_avg:39.71ms step:626/1575 train_time:24876ms step_avg:39.74ms step:627/1575 train_time:24939ms step_avg:39.78ms step:628/1575 train_time:24999ms step_avg:39.81ms step:629/1575 train_time:25062ms step_avg:39.84ms step:630/1575 train_time:25122ms step_avg:39.88ms step:631/1575 train_time:25185ms step_avg:39.91ms step:632/1575 train_time:25245ms step_avg:39.94ms step:633/1575 train_time:25308ms step_avg:39.98ms step:634/1575 train_time:25368ms step_avg:40.01ms step:635/1575 train_time:25430ms step_avg:40.05ms step:636/1575 train_time:25490ms step_avg:40.08ms step:637/1575 train_time:25554ms step_avg:40.12ms step:638/1575 train_time:25613ms step_avg:40.15ms step:639/1575 train_time:25677ms step_avg:40.18ms step:640/1575 train_time:25736ms step_avg:40.21ms step:641/1575 train_time:25799ms step_avg:40.25ms step:642/1575 train_time:25859ms step_avg:40.28ms step:643/1575 train_time:25922ms step_avg:40.31ms step:644/1575 train_time:25982ms step_avg:40.34ms step:645/1575 train_time:26045ms step_avg:40.38ms step:646/1575 train_time:26105ms step_avg:40.41ms step:647/1575 train_time:26168ms step_avg:40.45ms step:648/1575 train_time:26230ms step_avg:40.48ms step:649/1575 train_time:26294ms step_avg:40.51ms step:650/1575 train_time:26351ms step_avg:40.54ms step:651/1575 train_time:26415ms step_avg:40.58ms step:652/1575 train_time:26473ms step_avg:40.60ms step:653/1575 train_time:26537ms step_avg:40.64ms step:654/1575 train_time:26597ms step_avg:40.67ms step:655/1575 train_time:26662ms step_avg:40.71ms step:656/1575 train_time:26720ms step_avg:40.73ms step:657/1575 train_time:26782ms step_avg:40.76ms step:658/1575 train_time:26841ms step_avg:40.79ms step:659/1575 train_time:26904ms step_avg:40.83ms step:660/1575 train_time:26964ms step_avg:40.85ms step:661/1575 train_time:27027ms step_avg:40.89ms step:662/1575 train_time:27087ms step_avg:40.92ms step:663/1575 train_time:27150ms step_avg:40.95ms step:664/1575 train_time:27210ms step_avg:40.98ms step:665/1575 train_time:27273ms step_avg:41.01ms step:666/1575 train_time:27333ms step_avg:41.04ms step:667/1575 train_time:27396ms step_avg:41.07ms step:668/1575 train_time:27455ms step_avg:41.10ms step:669/1575 train_time:27518ms step_avg:41.13ms step:670/1575 train_time:27578ms step_avg:41.16ms step:671/1575 train_time:27641ms step_avg:41.19ms step:672/1575 train_time:27700ms step_avg:41.22ms step:673/1575 train_time:27763ms step_avg:41.25ms step:674/1575 train_time:27823ms step_avg:41.28ms step:675/1575 train_time:27886ms step_avg:41.31ms step:676/1575 train_time:27945ms step_avg:41.34ms step:677/1575 train_time:28008ms step_avg:41.37ms step:678/1575 train_time:28068ms step_avg:41.40ms step:679/1575 train_time:28132ms step_avg:41.43ms step:680/1575 train_time:28191ms step_avg:41.46ms step:681/1575 train_time:28255ms step_avg:41.49ms step:682/1575 train_time:28315ms step_avg:41.52ms step:683/1575 train_time:28379ms step_avg:41.55ms step:684/1575 train_time:28438ms step_avg:41.58ms step:685/1575 train_time:28501ms step_avg:41.61ms step:686/1575 train_time:28560ms step_avg:41.63ms step:687/1575 train_time:28626ms step_avg:41.67ms step:688/1575 train_time:28683ms step_avg:41.69ms step:689/1575 train_time:28746ms step_avg:41.72ms step:690/1575 train_time:28807ms step_avg:41.75ms step:691/1575 train_time:28870ms step_avg:41.78ms step:692/1575 train_time:28929ms step_avg:41.80ms step:693/1575 train_time:28992ms step_avg:41.84ms step:694/1575 train_time:29052ms step_avg:41.86ms step:695/1575 train_time:29115ms step_avg:41.89ms step:696/1575 train_time:29175ms step_avg:41.92ms step:697/1575 train_time:29238ms step_avg:41.95ms step:698/1575 train_time:29298ms step_avg:41.97ms step:699/1575 train_time:29361ms step_avg:42.00ms step:700/1575 train_time:29421ms step_avg:42.03ms step:701/1575 train_time:29484ms step_avg:42.06ms step:702/1575 train_time:29543ms step_avg:42.08ms step:703/1575 train_time:29606ms step_avg:42.11ms step:704/1575 train_time:29666ms step_avg:42.14ms step:705/1575 train_time:29729ms step_avg:42.17ms step:706/1575 train_time:29789ms step_avg:42.19ms step:707/1575 train_time:29852ms step_avg:42.22ms step:708/1575 train_time:29912ms step_avg:42.25ms step:709/1575 train_time:29975ms step_avg:42.28ms step:710/1575 train_time:30035ms step_avg:42.30ms step:711/1575 train_time:30098ms step_avg:42.33ms step:712/1575 train_time:30158ms step_avg:42.36ms step:713/1575 train_time:30221ms step_avg:42.39ms step:714/1575 train_time:30281ms step_avg:42.41ms step:715/1575 train_time:30344ms step_avg:42.44ms step:716/1575 train_time:30405ms step_avg:42.46ms step:717/1575 train_time:30467ms step_avg:42.49ms step:718/1575 train_time:30527ms step_avg:42.52ms step:719/1575 train_time:30589ms step_avg:42.54ms step:720/1575 train_time:30649ms step_avg:42.57ms step:721/1575 train_time:30712ms step_avg:42.60ms step:722/1575 train_time:30771ms step_avg:42.62ms step:723/1575 train_time:30834ms step_avg:42.65ms step:724/1575 train_time:30895ms step_avg:42.67ms step:725/1575 train_time:30958ms step_avg:42.70ms step:726/1575 train_time:31020ms step_avg:42.73ms step:727/1575 train_time:31083ms step_avg:42.76ms step:728/1575 train_time:31140ms step_avg:42.77ms step:729/1575 train_time:31203ms step_avg:42.80ms step:730/1575 train_time:31263ms step_avg:42.83ms step:731/1575 train_time:31326ms step_avg:42.85ms step:732/1575 train_time:31385ms step_avg:42.88ms step:733/1575 train_time:31448ms step_avg:42.90ms step:734/1575 train_time:31508ms step_avg:42.93ms step:735/1575 train_time:31571ms step_avg:42.95ms step:736/1575 train_time:31630ms step_avg:42.98ms step:737/1575 train_time:31693ms step_avg:43.00ms step:738/1575 train_time:31756ms step_avg:43.03ms step:739/1575 train_time:31819ms step_avg:43.06ms step:740/1575 train_time:31876ms step_avg:43.08ms step:741/1575 train_time:31938ms step_avg:43.10ms step:742/1575 train_time:31998ms step_avg:43.12ms step:743/1575 train_time:32061ms step_avg:43.15ms step:744/1575 train_time:32120ms step_avg:43.17ms step:745/1575 train_time:32183ms step_avg:43.20ms step:746/1575 train_time:32243ms step_avg:43.22ms step:747/1575 train_time:32306ms step_avg:43.25ms step:748/1575 train_time:32365ms step_avg:43.27ms step:749/1575 train_time:32429ms step_avg:43.30ms step:750/1575 train_time:32488ms step_avg:43.32ms step:750/1575 val_loss:3.8908 train_time:32535ms step_avg:43.38ms step:751/1575 train_time:32556ms step_avg:43.35ms step:752/1575 train_time:32615ms step_avg:43.37ms step:753/1575 train_time:32680ms step_avg:43.40ms step:754/1575 train_time:32742ms step_avg:43.42ms step:755/1575 train_time:32806ms step_avg:43.45ms step:756/1575 train_time:32865ms step_avg:43.47ms step:757/1575 train_time:32928ms step_avg:43.50ms step:758/1575 train_time:32987ms step_avg:43.52ms step:759/1575 train_time:33050ms step_avg:43.54ms step:760/1575 train_time:33109ms step_avg:43.56ms step:761/1575 train_time:33173ms step_avg:43.59ms step:762/1575 train_time:33232ms step_avg:43.61ms step:763/1575 train_time:33294ms step_avg:43.64ms step:764/1575 train_time:33354ms step_avg:43.66ms step:765/1575 train_time:33416ms step_avg:43.68ms step:766/1575 train_time:33475ms step_avg:43.70ms step:767/1575 train_time:33538ms step_avg:43.73ms step:768/1575 train_time:33598ms step_avg:43.75ms step:769/1575 train_time:33663ms step_avg:43.78ms step:770/1575 train_time:33723ms step_avg:43.80ms step:771/1575 train_time:33789ms step_avg:43.82ms step:772/1575 train_time:33848ms step_avg:43.84ms step:773/1575 train_time:33912ms step_avg:43.87ms step:774/1575 train_time:33971ms step_avg:43.89ms step:775/1575 train_time:34034ms step_avg:43.92ms step:776/1575 train_time:34094ms step_avg:43.94ms step:777/1575 train_time:34156ms step_avg:43.96ms step:778/1575 train_time:34215ms step_avg:43.98ms step:779/1575 train_time:34278ms step_avg:44.00ms step:780/1575 train_time:34337ms step_avg:44.02ms step:781/1575 train_time:34400ms step_avg:44.05ms step:782/1575 train_time:34459ms step_avg:44.07ms step:783/1575 train_time:34522ms step_avg:44.09ms step:784/1575 train_time:34583ms step_avg:44.11ms step:785/1575 train_time:34646ms step_avg:44.13ms step:786/1575 train_time:34706ms step_avg:44.16ms step:787/1575 train_time:34770ms step_avg:44.18ms step:788/1575 train_time:34830ms step_avg:44.20ms step:789/1575 train_time:34893ms step_avg:44.22ms step:790/1575 train_time:34953ms step_avg:44.24ms step:791/1575 train_time:35016ms step_avg:44.27ms step:792/1575 train_time:35075ms step_avg:44.29ms step:793/1575 train_time:35138ms step_avg:44.31ms step:794/1575 train_time:35197ms step_avg:44.33ms step:795/1575 train_time:35259ms step_avg:44.35ms step:796/1575 train_time:35319ms step_avg:44.37ms step:797/1575 train_time:35382ms step_avg:44.39ms step:798/1575 train_time:35441ms step_avg:44.41ms step:799/1575 train_time:35504ms step_avg:44.44ms step:800/1575 train_time:35564ms step_avg:44.45ms step:801/1575 train_time:35627ms step_avg:44.48ms step:802/1575 train_time:35686ms step_avg:44.50ms step:803/1575 train_time:35750ms step_avg:44.52ms step:804/1575 train_time:35809ms step_avg:44.54ms step:805/1575 train_time:35873ms step_avg:44.56ms step:806/1575 train_time:35933ms step_avg:44.58ms step:807/1575 train_time:36000ms step_avg:44.61ms step:808/1575 train_time:36060ms step_avg:44.63ms step:809/1575 train_time:36120ms step_avg:44.65ms step:810/1575 train_time:36179ms step_avg:44.67ms step:811/1575 train_time:36241ms step_avg:44.69ms step:812/1575 train_time:36301ms step_avg:44.71ms step:813/1575 train_time:36364ms step_avg:44.73ms step:814/1575 train_time:36423ms step_avg:44.75ms step:815/1575 train_time:36486ms step_avg:44.77ms step:816/1575 train_time:36546ms step_avg:44.79ms step:817/1575 train_time:36609ms step_avg:44.81ms step:818/1575 train_time:36668ms step_avg:44.83ms step:819/1575 train_time:36732ms step_avg:44.85ms step:820/1575 train_time:36792ms step_avg:44.87ms step:821/1575 train_time:36856ms step_avg:44.89ms step:822/1575 train_time:36915ms step_avg:44.91ms step:823/1575 train_time:36979ms step_avg:44.93ms step:824/1575 train_time:37038ms step_avg:44.95ms step:825/1575 train_time:37103ms step_avg:44.97ms step:826/1575 train_time:37161ms step_avg:44.99ms step:827/1575 train_time:37224ms step_avg:45.01ms step:828/1575 train_time:37284ms step_avg:45.03ms step:829/1575 train_time:37347ms step_avg:45.05ms step:830/1575 train_time:37407ms step_avg:45.07ms step:831/1575 train_time:37470ms step_avg:45.09ms step:832/1575 train_time:37529ms step_avg:45.11ms step:833/1575 train_time:37592ms step_avg:45.13ms step:834/1575 train_time:37651ms step_avg:45.15ms step:835/1575 train_time:37714ms step_avg:45.17ms step:836/1575 train_time:37774ms step_avg:45.18ms step:837/1575 train_time:37836ms step_avg:45.20ms step:838/1575 train_time:37896ms step_avg:45.22ms step:839/1575 train_time:37959ms step_avg:45.24ms step:840/1575 train_time:38019ms step_avg:45.26ms step:841/1575 train_time:38082ms step_avg:45.28ms step:842/1575 train_time:38141ms step_avg:45.30ms step:843/1575 train_time:38204ms step_avg:45.32ms step:844/1575 train_time:38264ms step_avg:45.34ms step:845/1575 train_time:38326ms step_avg:45.36ms step:846/1575 train_time:38385ms step_avg:45.37ms step:847/1575 train_time:38449ms step_avg:45.39ms step:848/1575 train_time:38508ms step_avg:45.41ms step:849/1575 train_time:38571ms step_avg:45.43ms step:850/1575 train_time:38631ms step_avg:45.45ms step:851/1575 train_time:38693ms step_avg:45.47ms step:852/1575 train_time:38753ms step_avg:45.48ms step:853/1575 train_time:38816ms step_avg:45.51ms step:854/1575 train_time:38875ms step_avg:45.52ms step:855/1575 train_time:38938ms step_avg:45.54ms step:856/1575 train_time:38998ms step_avg:45.56ms step:857/1575 train_time:39062ms step_avg:45.58ms step:858/1575 train_time:39121ms step_avg:45.60ms step:859/1575 train_time:39184ms step_avg:45.62ms step:860/1575 train_time:39244ms step_avg:45.63ms step:861/1575 train_time:39307ms step_avg:45.65ms step:862/1575 train_time:39367ms step_avg:45.67ms step:863/1575 train_time:39430ms step_avg:45.69ms step:864/1575 train_time:39489ms step_avg:45.71ms step:865/1575 train_time:39553ms step_avg:45.73ms step:866/1575 train_time:39612ms step_avg:45.74ms step:867/1575 train_time:39675ms step_avg:45.76ms step:868/1575 train_time:39735ms step_avg:45.78ms step:869/1575 train_time:39797ms step_avg:45.80ms step:870/1575 train_time:39857ms step_avg:45.81ms step:871/1575 train_time:39920ms step_avg:45.83ms step:872/1575 train_time:39980ms step_avg:45.85ms step:873/1575 train_time:40043ms step_avg:45.87ms step:874/1575 train_time:40102ms step_avg:45.88ms step:875/1575 train_time:40165ms step_avg:45.90ms step:876/1575 train_time:40225ms step_avg:45.92ms step:877/1575 train_time:40290ms step_avg:45.94ms step:878/1575 train_time:40351ms step_avg:45.96ms step:879/1575 train_time:40414ms step_avg:45.98ms step:880/1575 train_time:40472ms step_avg:45.99ms step:881/1575 train_time:40535ms step_avg:46.01ms step:882/1575 train_time:40594ms step_avg:46.03ms step:883/1575 train_time:40657ms step_avg:46.04ms step:884/1575 train_time:40716ms step_avg:46.06ms step:885/1575 train_time:40779ms step_avg:46.08ms step:886/1575 train_time:40843ms step_avg:46.10ms step:887/1575 train_time:40904ms step_avg:46.12ms step:888/1575 train_time:40964ms step_avg:46.13ms step:889/1575 train_time:41027ms step_avg:46.15ms step:890/1575 train_time:41086ms step_avg:46.16ms step:891/1575 train_time:41151ms step_avg:46.19ms step:892/1575 train_time:41212ms step_avg:46.20ms step:893/1575 train_time:41271ms step_avg:46.22ms step:894/1575 train_time:41330ms step_avg:46.23ms step:895/1575 train_time:41393ms step_avg:46.25ms step:896/1575 train_time:41452ms step_avg:46.26ms step:897/1575 train_time:41516ms step_avg:46.28ms step:898/1575 train_time:41575ms step_avg:46.30ms step:899/1575 train_time:41637ms step_avg:46.32ms step:900/1575 train_time:41697ms step_avg:46.33ms step:901/1575 train_time:41760ms step_avg:46.35ms step:902/1575 train_time:41820ms step_avg:46.36ms step:903/1575 train_time:41883ms step_avg:46.38ms step:904/1575 train_time:41943ms step_avg:46.40ms step:905/1575 train_time:42006ms step_avg:46.42ms step:906/1575 train_time:42065ms step_avg:46.43ms step:907/1575 train_time:42129ms step_avg:46.45ms step:908/1575 train_time:42190ms step_avg:46.46ms step:909/1575 train_time:42252ms step_avg:46.48ms step:910/1575 train_time:42311ms step_avg:46.50ms step:911/1575 train_time:42375ms step_avg:46.51ms step:912/1575 train_time:42434ms step_avg:46.53ms step:913/1575 train_time:42497ms step_avg:46.55ms step:914/1575 train_time:42557ms step_avg:46.56ms step:915/1575 train_time:42620ms step_avg:46.58ms step:916/1575 train_time:42679ms step_avg:46.59ms step:917/1575 train_time:42744ms step_avg:46.61ms step:918/1575 train_time:42802ms step_avg:46.62ms step:919/1575 train_time:42865ms step_avg:46.64ms step:920/1575 train_time:42925ms step_avg:46.66ms step:921/1575 train_time:42988ms step_avg:46.68ms step:922/1575 train_time:43047ms step_avg:46.69ms step:923/1575 train_time:43110ms step_avg:46.71ms step:924/1575 train_time:43170ms step_avg:46.72ms step:925/1575 train_time:43233ms step_avg:46.74ms step:926/1575 train_time:43293ms step_avg:46.75ms step:927/1575 train_time:43356ms step_avg:46.77ms step:928/1575 train_time:43416ms step_avg:46.78ms step:929/1575 train_time:43479ms step_avg:46.80ms step:930/1575 train_time:43539ms step_avg:46.82ms step:931/1575 train_time:43601ms step_avg:46.83ms step:932/1575 train_time:43660ms step_avg:46.85ms step:933/1575 train_time:43724ms step_avg:46.86ms step:934/1575 train_time:43784ms step_avg:46.88ms step:935/1575 train_time:43847ms step_avg:46.90ms step:936/1575 train_time:43906ms step_avg:46.91ms step:937/1575 train_time:43969ms step_avg:46.93ms step:938/1575 train_time:44030ms step_avg:46.94ms step:939/1575 train_time:44092ms step_avg:46.96ms step:940/1575 train_time:44151ms step_avg:46.97ms step:941/1575 train_time:44215ms step_avg:46.99ms step:942/1575 train_time:44274ms step_avg:47.00ms step:943/1575 train_time:44337ms step_avg:47.02ms step:944/1575 train_time:44396ms step_avg:47.03ms step:945/1575 train_time:44459ms step_avg:47.05ms step:946/1575 train_time:44519ms step_avg:47.06ms step:947/1575 train_time:44582ms step_avg:47.08ms step:948/1575 train_time:44642ms step_avg:47.09ms step:949/1575 train_time:44705ms step_avg:47.11ms step:950/1575 train_time:44764ms step_avg:47.12ms step:951/1575 train_time:44829ms step_avg:47.14ms step:952/1575 train_time:44888ms step_avg:47.15ms step:953/1575 train_time:44951ms step_avg:47.17ms step:954/1575 train_time:45010ms step_avg:47.18ms step:955/1575 train_time:45073ms step_avg:47.20ms step:956/1575 train_time:45135ms step_avg:47.21ms step:957/1575 train_time:45196ms step_avg:47.23ms step:958/1575 train_time:45255ms step_avg:47.24ms step:959/1575 train_time:45318ms step_avg:47.26ms step:960/1575 train_time:45377ms step_avg:47.27ms step:961/1575 train_time:45441ms step_avg:47.29ms step:962/1575 train_time:45503ms step_avg:47.30ms step:963/1575 train_time:45564ms step_avg:47.31ms step:964/1575 train_time:45625ms step_avg:47.33ms step:965/1575 train_time:45688ms step_avg:47.34ms step:966/1575 train_time:45747ms step_avg:47.36ms step:967/1575 train_time:45810ms step_avg:47.37ms step:968/1575 train_time:45869ms step_avg:47.39ms step:969/1575 train_time:45932ms step_avg:47.40ms step:970/1575 train_time:45992ms step_avg:47.41ms step:971/1575 train_time:46054ms step_avg:47.43ms step:972/1575 train_time:46114ms step_avg:47.44ms step:973/1575 train_time:46177ms step_avg:47.46ms step:974/1575 train_time:46238ms step_avg:47.47ms step:975/1575 train_time:46300ms step_avg:47.49ms step:976/1575 train_time:46359ms step_avg:47.50ms step:977/1575 train_time:46425ms step_avg:47.52ms step:978/1575 train_time:46483ms step_avg:47.53ms step:979/1575 train_time:46546ms step_avg:47.54ms step:980/1575 train_time:46605ms step_avg:47.56ms step:981/1575 train_time:46668ms step_avg:47.57ms step:982/1575 train_time:46728ms step_avg:47.58ms step:983/1575 train_time:46791ms step_avg:47.60ms step:984/1575 train_time:46850ms step_avg:47.61ms step:985/1575 train_time:46914ms step_avg:47.63ms step:986/1575 train_time:46973ms step_avg:47.64ms step:987/1575 train_time:47037ms step_avg:47.66ms step:988/1575 train_time:47096ms step_avg:47.67ms step:989/1575 train_time:47159ms step_avg:47.68ms step:990/1575 train_time:47219ms step_avg:47.70ms step:991/1575 train_time:47281ms step_avg:47.71ms step:992/1575 train_time:47340ms step_avg:47.72ms step:993/1575 train_time:47404ms step_avg:47.74ms step:994/1575 train_time:47463ms step_avg:47.75ms step:995/1575 train_time:47526ms step_avg:47.77ms step:996/1575 train_time:47587ms step_avg:47.78ms step:997/1575 train_time:47650ms step_avg:47.79ms step:998/1575 train_time:47709ms step_avg:47.80ms step:999/1575 train_time:47773ms step_avg:47.82ms step:1000/1575 train_time:47833ms step_avg:47.83ms step:1000/1575 val_loss:3.5834 train_time:47879ms step_avg:47.88ms step:1001/1575 train_time:47899ms step_avg:47.85ms step:1002/1575 train_time:47960ms step_avg:47.86ms step:1003/1575 train_time:48025ms step_avg:47.88ms step:1004/1575 train_time:48086ms step_avg:47.89ms step:1005/1575 train_time:48150ms step_avg:47.91ms step:1006/1575 train_time:48209ms step_avg:47.92ms step:1007/1575 train_time:48273ms step_avg:47.94ms step:1008/1575 train_time:48331ms step_avg:47.95ms step:1009/1575 train_time:48394ms step_avg:47.96ms step:1010/1575 train_time:48455ms step_avg:47.97ms step:1011/1575 train_time:48518ms step_avg:47.99ms step:1012/1575 train_time:48577ms step_avg:48.00ms step:1013/1575 train_time:48639ms step_avg:48.02ms step:1014/1575 train_time:48698ms step_avg:48.03ms step:1015/1575 train_time:48760ms step_avg:48.04ms step:1016/1575 train_time:48820ms step_avg:48.05ms step:1017/1575 train_time:48884ms step_avg:48.07ms step:1018/1575 train_time:48945ms step_avg:48.08ms step:1019/1575 train_time:49009ms step_avg:48.09ms step:1020/1575 train_time:49069ms step_avg:48.11ms step:1021/1575 train_time:49133ms step_avg:48.12ms step:1022/1575 train_time:49193ms step_avg:48.13ms step:1023/1575 train_time:49256ms step_avg:48.15ms step:1024/1575 train_time:49315ms step_avg:48.16ms step:1025/1575 train_time:49388ms step_avg:48.18ms step:1026/1575 train_time:49470ms step_avg:48.22ms step:1027/1575 train_time:49561ms step_avg:48.26ms step:1028/1575 train_time:49646ms step_avg:48.29ms step:1029/1575 train_time:49735ms step_avg:48.33ms step:1030/1575 train_time:49821ms step_avg:48.37ms step:1031/1575 train_time:49911ms step_avg:48.41ms step:1032/1575 train_time:49997ms step_avg:48.45ms step:1033/1575 train_time:50089ms step_avg:48.49ms step:1034/1575 train_time:50175ms step_avg:48.52ms step:1035/1575 train_time:50265ms step_avg:48.57ms step:1036/1575 train_time:50351ms step_avg:48.60ms step:1037/1575 train_time:50440ms step_avg:48.64ms step:1038/1575 train_time:50529ms step_avg:48.68ms step:1039/1575 train_time:50615ms step_avg:48.72ms step:1040/1575 train_time:50700ms step_avg:48.75ms step:1041/1575 train_time:50790ms step_avg:48.79ms step:1042/1575 train_time:50876ms step_avg:48.83ms step:1043/1575 train_time:50966ms step_avg:48.86ms step:1044/1575 train_time:51052ms step_avg:48.90ms step:1045/1575 train_time:51142ms step_avg:48.94ms step:1046/1575 train_time:51228ms step_avg:48.97ms step:1047/1575 train_time:51317ms step_avg:49.01ms step:1048/1575 train_time:51403ms step_avg:49.05ms step:1049/1575 train_time:51493ms step_avg:49.09ms step:1050/1575 train_time:51578ms step_avg:49.12ms step:1051/1575 train_time:51667ms step_avg:49.16ms step:1052/1575 train_time:51753ms step_avg:49.19ms step:1053/1575 train_time:51842ms step_avg:49.23ms step:1054/1575 train_time:51928ms step_avg:49.27ms step:1055/1575 train_time:52017ms step_avg:49.31ms step:1056/1575 train_time:52104ms step_avg:49.34ms step:1057/1575 train_time:52194ms step_avg:49.38ms step:1058/1575 train_time:52280ms step_avg:49.41ms step:1059/1575 train_time:52370ms step_avg:49.45ms step:1060/1575 train_time:52456ms step_avg:49.49ms step:1061/1575 train_time:52545ms step_avg:49.52ms step:1062/1575 train_time:52631ms step_avg:49.56ms step:1063/1575 train_time:52720ms step_avg:49.60ms step:1064/1575 train_time:52806ms step_avg:49.63ms step:1065/1575 train_time:52896ms step_avg:49.67ms step:1066/1575 train_time:52982ms step_avg:49.70ms step:1067/1575 train_time:53072ms step_avg:49.74ms step:1068/1575 train_time:53158ms step_avg:49.77ms step:1069/1575 train_time:53247ms step_avg:49.81ms step:1070/1575 train_time:53333ms step_avg:49.84ms step:1071/1575 train_time:53422ms step_avg:49.88ms step:1072/1575 train_time:53510ms step_avg:49.92ms step:1073/1575 train_time:53597ms step_avg:49.95ms step:1074/1575 train_time:53683ms step_avg:49.98ms step:1075/1575 train_time:53773ms step_avg:50.02ms step:1076/1575 train_time:53859ms step_avg:50.06ms step:1077/1575 train_time:53948ms step_avg:50.09ms step:1078/1575 train_time:54034ms step_avg:50.12ms step:1079/1575 train_time:54123ms step_avg:50.16ms step:1080/1575 train_time:54209ms step_avg:50.19ms step:1081/1575 train_time:54298ms step_avg:50.23ms step:1082/1575 train_time:54385ms step_avg:50.26ms step:1083/1575 train_time:54474ms step_avg:50.30ms step:1084/1575 train_time:54565ms step_avg:50.34ms step:1085/1575 train_time:54651ms step_avg:50.37ms step:1086/1575 train_time:54737ms step_avg:50.40ms step:1087/1575 train_time:54826ms step_avg:50.44ms step:1088/1575 train_time:54912ms step_avg:50.47ms step:1089/1575 train_time:55002ms step_avg:50.51ms step:1090/1575 train_time:55088ms step_avg:50.54ms step:1091/1575 train_time:55177ms step_avg:50.57ms step:1092/1575 train_time:55264ms step_avg:50.61ms step:1093/1575 train_time:55353ms step_avg:50.64ms step:1094/1575 train_time:55438ms step_avg:50.67ms step:1095/1575 train_time:55528ms step_avg:50.71ms step:1096/1575 train_time:55617ms step_avg:50.75ms step:1097/1575 train_time:55704ms step_avg:50.78ms step:1098/1575 train_time:55789ms step_avg:50.81ms step:1099/1575 train_time:55878ms step_avg:50.84ms step:1100/1575 train_time:55964ms step_avg:50.88ms step:1101/1575 train_time:56054ms step_avg:50.91ms step:1102/1575 train_time:56140ms step_avg:50.94ms step:1103/1575 train_time:56231ms step_avg:50.98ms step:1104/1575 train_time:56317ms step_avg:51.01ms step:1105/1575 train_time:56406ms step_avg:51.05ms step:1106/1575 train_time:56492ms step_avg:51.08ms step:1107/1575 train_time:56582ms step_avg:51.11ms step:1108/1575 train_time:56668ms step_avg:51.14ms step:1109/1575 train_time:56757ms step_avg:51.18ms step:1110/1575 train_time:56843ms step_avg:51.21ms step:1111/1575 train_time:56932ms step_avg:51.24ms step:1112/1575 train_time:57019ms step_avg:51.28ms step:1113/1575 train_time:57108ms step_avg:51.31ms step:1114/1575 train_time:57193ms step_avg:51.34ms step:1115/1575 train_time:57283ms step_avg:51.37ms step:1116/1575 train_time:57368ms step_avg:51.41ms step:1117/1575 train_time:57457ms step_avg:51.44ms step:1118/1575 train_time:57543ms step_avg:51.47ms step:1119/1575 train_time:57633ms step_avg:51.50ms step:1120/1575 train_time:57719ms step_avg:51.53ms step:1121/1575 train_time:57808ms step_avg:51.57ms step:1122/1575 train_time:57894ms step_avg:51.60ms step:1123/1575 train_time:57984ms step_avg:51.63ms step:1124/1575 train_time:58069ms step_avg:51.66ms step:1125/1575 train_time:58159ms step_avg:51.70ms step:1126/1575 train_time:58245ms step_avg:51.73ms step:1127/1575 train_time:58335ms step_avg:51.76ms step:1128/1575 train_time:58420ms step_avg:51.79ms step:1129/1575 train_time:58510ms step_avg:51.82ms step:1130/1575 train_time:58598ms step_avg:51.86ms step:1131/1575 train_time:58687ms step_avg:51.89ms step:1132/1575 train_time:58783ms step_avg:51.93ms step:1133/1575 train_time:58868ms step_avg:51.96ms step:1134/1575 train_time:58953ms step_avg:51.99ms step:1135/1575 train_time:59041ms step_avg:52.02ms step:1136/1575 train_time:59127ms step_avg:52.05ms step:1137/1575 train_time:59216ms step_avg:52.08ms step:1138/1575 train_time:59303ms step_avg:52.11ms step:1139/1575 train_time:59394ms step_avg:52.15ms step:1140/1575 train_time:59473ms step_avg:52.17ms step:1141/1575 train_time:59561ms step_avg:52.20ms step:1142/1575 train_time:59646ms step_avg:52.23ms step:1143/1575 train_time:59736ms step_avg:52.26ms step:1144/1575 train_time:59822ms step_avg:52.29ms step:1145/1575 train_time:59911ms step_avg:52.32ms step:1146/1575 train_time:60001ms step_avg:52.36ms step:1147/1575 train_time:60089ms step_avg:52.39ms step:1148/1575 train_time:60174ms step_avg:52.42ms step:1149/1575 train_time:60264ms step_avg:52.45ms step:1150/1575 train_time:60350ms step_avg:52.48ms step:1151/1575 train_time:60439ms step_avg:52.51ms step:1152/1575 train_time:60526ms step_avg:52.54ms step:1153/1575 train_time:60614ms step_avg:52.57ms step:1154/1575 train_time:60701ms step_avg:52.60ms step:1155/1575 train_time:60792ms step_avg:52.63ms step:1156/1575 train_time:60875ms step_avg:52.66ms step:1157/1575 train_time:60965ms step_avg:52.69ms step:1158/1575 train_time:61051ms step_avg:52.72ms step:1159/1575 train_time:61140ms step_avg:52.75ms step:1160/1575 train_time:61226ms step_avg:52.78ms step:1161/1575 train_time:61316ms step_avg:52.81ms step:1162/1575 train_time:61402ms step_avg:52.84ms step:1163/1575 train_time:61492ms step_avg:52.87ms step:1164/1575 train_time:61577ms step_avg:52.90ms step:1165/1575 train_time:61667ms step_avg:52.93ms step:1166/1575 train_time:61752ms step_avg:52.96ms step:1167/1575 train_time:61842ms step_avg:52.99ms step:1168/1575 train_time:61927ms step_avg:53.02ms step:1169/1575 train_time:62017ms step_avg:53.05ms step:1170/1575 train_time:62103ms step_avg:53.08ms step:1171/1575 train_time:62193ms step_avg:53.11ms step:1172/1575 train_time:62279ms step_avg:53.14ms step:1173/1575 train_time:62369ms step_avg:53.17ms step:1174/1575 train_time:62455ms step_avg:53.20ms step:1175/1575 train_time:62544ms step_avg:53.23ms step:1176/1575 train_time:62629ms step_avg:53.26ms step:1177/1575 train_time:62718ms step_avg:53.29ms step:1178/1575 train_time:62804ms step_avg:53.31ms step:1179/1575 train_time:62894ms step_avg:53.34ms step:1180/1575 train_time:62980ms step_avg:53.37ms step:1181/1575 train_time:63069ms step_avg:53.40ms step:1182/1575 train_time:63155ms step_avg:53.43ms step:1183/1575 train_time:63245ms step_avg:53.46ms step:1184/1575 train_time:63330ms step_avg:53.49ms step:1185/1575 train_time:63420ms step_avg:53.52ms step:1186/1575 train_time:63505ms step_avg:53.55ms step:1187/1575 train_time:63595ms step_avg:53.58ms step:1188/1575 train_time:63680ms step_avg:53.60ms step:1189/1575 train_time:63771ms step_avg:53.63ms step:1190/1575 train_time:63856ms step_avg:53.66ms step:1191/1575 train_time:63945ms step_avg:53.69ms step:1192/1575 train_time:64031ms step_avg:53.72ms step:1193/1575 train_time:64120ms step_avg:53.75ms step:1194/1575 train_time:64206ms step_avg:53.77ms step:1195/1575 train_time:64296ms step_avg:53.80ms step:1196/1575 train_time:64382ms step_avg:53.83ms step:1197/1575 train_time:64472ms step_avg:53.86ms step:1198/1575 train_time:64558ms step_avg:53.89ms step:1199/1575 train_time:64648ms step_avg:53.92ms step:1200/1575 train_time:64733ms step_avg:53.94ms step:1201/1575 train_time:64822ms step_avg:53.97ms step:1202/1575 train_time:64908ms step_avg:54.00ms step:1203/1575 train_time:64998ms step_avg:54.03ms step:1204/1575 train_time:65084ms step_avg:54.06ms step:1205/1575 train_time:65173ms step_avg:54.09ms step:1206/1575 train_time:65259ms step_avg:54.11ms step:1207/1575 train_time:65349ms step_avg:54.14ms step:1208/1575 train_time:65435ms step_avg:54.17ms step:1209/1575 train_time:65525ms step_avg:54.20ms step:1210/1575 train_time:65611ms step_avg:54.22ms step:1211/1575 train_time:65701ms step_avg:54.25ms step:1212/1575 train_time:65787ms step_avg:54.28ms step:1213/1575 train_time:65879ms step_avg:54.31ms step:1214/1575 train_time:65962ms step_avg:54.33ms step:1215/1575 train_time:66051ms step_avg:54.36ms step:1216/1575 train_time:66138ms step_avg:54.39ms step:1217/1575 train_time:66227ms step_avg:54.42ms step:1218/1575 train_time:66313ms step_avg:54.44ms step:1219/1575 train_time:66402ms step_avg:54.47ms step:1220/1575 train_time:66488ms step_avg:54.50ms step:1221/1575 train_time:66577ms step_avg:54.53ms step:1222/1575 train_time:66663ms step_avg:54.55ms step:1223/1575 train_time:66753ms step_avg:54.58ms step:1224/1575 train_time:66838ms step_avg:54.61ms step:1225/1575 train_time:66929ms step_avg:54.64ms step:1226/1575 train_time:67014ms step_avg:54.66ms step:1227/1575 train_time:67105ms step_avg:54.69ms step:1228/1575 train_time:67190ms step_avg:54.72ms step:1229/1575 train_time:67279ms step_avg:54.74ms step:1230/1575 train_time:67365ms step_avg:54.77ms step:1231/1575 train_time:67454ms step_avg:54.80ms step:1232/1575 train_time:67540ms step_avg:54.82ms step:1233/1575 train_time:67631ms step_avg:54.85ms step:1234/1575 train_time:67717ms step_avg:54.88ms step:1235/1575 train_time:67807ms step_avg:54.90ms step:1236/1575 train_time:67892ms step_avg:54.93ms step:1237/1575 train_time:67981ms step_avg:54.96ms step:1238/1575 train_time:68067ms step_avg:54.98ms step:1239/1575 train_time:68157ms step_avg:55.01ms step:1240/1575 train_time:68242ms step_avg:55.03ms step:1241/1575 train_time:68332ms step_avg:55.06ms step:1242/1575 train_time:68418ms step_avg:55.09ms step:1243/1575 train_time:68508ms step_avg:55.11ms step:1244/1575 train_time:68594ms step_avg:55.14ms step:1245/1575 train_time:68683ms step_avg:55.17ms step:1246/1575 train_time:68769ms step_avg:55.19ms step:1247/1575 train_time:68858ms step_avg:55.22ms step:1248/1575 train_time:68944ms step_avg:55.24ms step:1249/1575 train_time:69034ms step_avg:55.27ms step:1250/1575 train_time:69120ms step_avg:55.30ms step:1250/1575 val_loss:3.4073 train_time:69194ms step_avg:55.36ms step:1251/1575 train_time:69214ms step_avg:55.33ms step:1252/1575 train_time:69301ms step_avg:55.35ms step:1253/1575 train_time:69393ms step_avg:55.38ms step:1254/1575 train_time:69482ms step_avg:55.41ms step:1255/1575 train_time:69569ms step_avg:55.43ms step:1256/1575 train_time:69654ms step_avg:55.46ms step:1257/1575 train_time:69742ms step_avg:55.48ms step:1258/1575 train_time:69827ms step_avg:55.51ms step:1259/1575 train_time:69916ms step_avg:55.53ms step:1260/1575 train_time:70000ms step_avg:55.56ms step:1261/1575 train_time:70091ms step_avg:55.58ms step:1262/1575 train_time:70177ms step_avg:55.61ms step:1263/1575 train_time:70267ms step_avg:55.63ms step:1264/1575 train_time:70354ms step_avg:55.66ms step:1265/1575 train_time:70446ms step_avg:55.69ms step:1266/1575 train_time:70531ms step_avg:55.71ms step:1267/1575 train_time:70620ms step_avg:55.74ms step:1268/1575 train_time:70705ms step_avg:55.76ms step:1269/1575 train_time:70794ms step_avg:55.79ms step:1270/1575 train_time:70878ms step_avg:55.81ms step:1271/1575 train_time:70967ms step_avg:55.84ms step:1272/1575 train_time:71052ms step_avg:55.86ms step:1273/1575 train_time:71142ms step_avg:55.89ms step:1274/1575 train_time:71229ms step_avg:55.91ms step:1275/1575 train_time:71320ms step_avg:55.94ms step:1276/1575 train_time:71406ms step_avg:55.96ms step:1277/1575 train_time:71496ms step_avg:55.99ms step:1278/1575 train_time:71582ms step_avg:56.01ms step:1279/1575 train_time:71671ms step_avg:56.04ms step:1280/1575 train_time:71757ms step_avg:56.06ms step:1281/1575 train_time:71847ms step_avg:56.09ms step:1282/1575 train_time:71931ms step_avg:56.11ms step:1283/1575 train_time:72019ms step_avg:56.13ms step:1284/1575 train_time:72105ms step_avg:56.16ms step:1285/1575 train_time:72194ms step_avg:56.18ms step:1286/1575 train_time:72281ms step_avg:56.21ms step:1287/1575 train_time:72372ms step_avg:56.23ms step:1288/1575 train_time:72458ms step_avg:56.26ms step:1289/1575 train_time:72549ms step_avg:56.28ms step:1290/1575 train_time:72634ms step_avg:56.31ms step:1291/1575 train_time:72723ms step_avg:56.33ms step:1292/1575 train_time:72808ms step_avg:56.35ms step:1293/1575 train_time:72897ms step_avg:56.38ms step:1294/1575 train_time:72982ms step_avg:56.40ms step:1295/1575 train_time:73071ms step_avg:56.43ms step:1296/1575 train_time:73157ms step_avg:56.45ms step:1297/1575 train_time:73248ms step_avg:56.47ms step:1298/1575 train_time:73335ms step_avg:56.50ms step:1299/1575 train_time:73428ms step_avg:56.53ms step:1300/1575 train_time:73512ms step_avg:56.55ms step:1301/1575 train_time:73601ms step_avg:56.57ms step:1302/1575 train_time:73687ms step_avg:56.59ms step:1303/1575 train_time:73776ms step_avg:56.62ms step:1304/1575 train_time:73861ms step_avg:56.64ms step:1305/1575 train_time:73950ms step_avg:56.67ms step:1306/1575 train_time:74036ms step_avg:56.69ms step:1307/1575 train_time:74125ms step_avg:56.71ms step:1308/1575 train_time:74211ms step_avg:56.74ms step:1309/1575 train_time:74300ms step_avg:56.76ms step:1310/1575 train_time:74387ms step_avg:56.78ms step:1311/1575 train_time:74477ms step_avg:56.81ms step:1312/1575 train_time:74562ms step_avg:56.83ms step:1313/1575 train_time:74652ms step_avg:56.86ms step:1314/1575 train_time:74738ms step_avg:56.88ms step:1315/1575 train_time:74828ms step_avg:56.90ms step:1316/1575 train_time:74913ms step_avg:56.92ms step:1317/1575 train_time:75002ms step_avg:56.95ms step:1318/1575 train_time:75087ms step_avg:56.97ms step:1319/1575 train_time:75178ms step_avg:57.00ms step:1320/1575 train_time:75266ms step_avg:57.02ms step:1321/1575 train_time:75353ms step_avg:57.04ms step:1322/1575 train_time:75438ms step_avg:57.06ms step:1323/1575 train_time:75529ms step_avg:57.09ms step:1324/1575 train_time:75615ms step_avg:57.11ms step:1325/1575 train_time:75705ms step_avg:57.14ms step:1326/1575 train_time:75792ms step_avg:57.16ms step:1327/1575 train_time:75880ms step_avg:57.18ms step:1328/1575 train_time:75965ms step_avg:57.20ms step:1329/1575 train_time:76054ms step_avg:57.23ms step:1330/1575 train_time:76139ms step_avg:57.25ms step:1331/1575 train_time:76228ms step_avg:57.27ms step:1332/1575 train_time:76315ms step_avg:57.29ms step:1333/1575 train_time:76404ms step_avg:57.32ms step:1334/1575 train_time:76490ms step_avg:57.34ms step:1335/1575 train_time:76580ms step_avg:57.36ms step:1336/1575 train_time:76666ms step_avg:57.38ms step:1337/1575 train_time:76755ms step_avg:57.41ms step:1338/1575 train_time:76841ms step_avg:57.43ms step:1339/1575 train_time:76930ms step_avg:57.45ms step:1340/1575 train_time:77015ms step_avg:57.47ms step:1341/1575 train_time:77105ms step_avg:57.50ms step:1342/1575 train_time:77190ms step_avg:57.52ms step:1343/1575 train_time:77280ms step_avg:57.54ms step:1344/1575 train_time:77366ms step_avg:57.56ms step:1345/1575 train_time:77456ms step_avg:57.59ms step:1346/1575 train_time:77542ms step_avg:57.61ms step:1347/1575 train_time:77632ms step_avg:57.63ms step:1348/1575 train_time:77718ms step_avg:57.65ms step:1349/1575 train_time:77807ms step_avg:57.68ms step:1350/1575 train_time:77894ms step_avg:57.70ms step:1351/1575 train_time:77983ms step_avg:57.72ms step:1352/1575 train_time:78068ms step_avg:57.74ms step:1353/1575 train_time:78157ms step_avg:57.77ms step:1354/1575 train_time:78245ms step_avg:57.79ms step:1355/1575 train_time:78334ms step_avg:57.81ms step:1356/1575 train_time:78419ms step_avg:57.83ms step:1357/1575 train_time:78510ms step_avg:57.86ms step:1358/1575 train_time:78599ms step_avg:57.88ms step:1359/1575 train_time:78688ms step_avg:57.90ms step:1360/1575 train_time:78773ms step_avg:57.92ms step:1361/1575 train_time:78862ms step_avg:57.94ms step:1362/1575 train_time:78947ms step_avg:57.96ms step:1363/1575 train_time:79035ms step_avg:57.99ms step:1364/1575 train_time:79122ms step_avg:58.01ms step:1365/1575 train_time:79211ms step_avg:58.03ms step:1366/1575 train_time:79296ms step_avg:58.05ms step:1367/1575 train_time:79387ms step_avg:58.07ms step:1368/1575 train_time:79472ms step_avg:58.09ms step:1369/1575 train_time:79562ms step_avg:58.12ms step:1370/1575 train_time:79648ms step_avg:58.14ms step:1371/1575 train_time:79738ms step_avg:58.16ms step:1372/1575 train_time:79824ms step_avg:58.18ms step:1373/1575 train_time:79913ms step_avg:58.20ms step:1374/1575 train_time:79999ms step_avg:58.22ms step:1375/1575 train_time:80089ms step_avg:58.25ms step:1376/1575 train_time:80176ms step_avg:58.27ms step:1377/1575 train_time:80264ms step_avg:58.29ms step:1378/1575 train_time:80353ms step_avg:58.31ms step:1379/1575 train_time:80442ms step_avg:58.33ms step:1380/1575 train_time:80526ms step_avg:58.35ms step:1381/1575 train_time:80614ms step_avg:58.37ms step:1382/1575 train_time:80701ms step_avg:58.39ms step:1383/1575 train_time:80790ms step_avg:58.42ms step:1384/1575 train_time:80877ms step_avg:58.44ms step:1385/1575 train_time:80966ms step_avg:58.46ms step:1386/1575 train_time:81052ms step_avg:58.48ms step:1387/1575 train_time:81141ms step_avg:58.50ms step:1388/1575 train_time:81227ms step_avg:58.52ms step:1389/1575 train_time:81316ms step_avg:58.54ms step:1390/1575 train_time:81402ms step_avg:58.56ms step:1391/1575 train_time:81492ms step_avg:58.59ms step:1392/1575 train_time:81578ms step_avg:58.61ms step:1393/1575 train_time:81668ms step_avg:58.63ms step:1394/1575 train_time:81754ms step_avg:58.65ms step:1395/1575 train_time:81844ms step_avg:58.67ms step:1396/1575 train_time:81930ms step_avg:58.69ms step:1397/1575 train_time:82020ms step_avg:58.71ms step:1398/1575 train_time:82105ms step_avg:58.73ms step:1399/1575 train_time:82195ms step_avg:58.75ms step:1400/1575 train_time:82281ms step_avg:58.77ms step:1401/1575 train_time:82370ms step_avg:58.79ms step:1402/1575 train_time:82457ms step_avg:58.81ms step:1403/1575 train_time:82546ms step_avg:58.84ms step:1404/1575 train_time:82632ms step_avg:58.85ms step:1405/1575 train_time:82721ms step_avg:58.88ms step:1406/1575 train_time:82807ms step_avg:58.90ms step:1407/1575 train_time:82896ms step_avg:58.92ms step:1408/1575 train_time:82983ms step_avg:58.94ms step:1409/1575 train_time:83071ms step_avg:58.96ms step:1410/1575 train_time:83157ms step_avg:58.98ms step:1411/1575 train_time:83247ms step_avg:59.00ms step:1412/1575 train_time:83334ms step_avg:59.02ms step:1413/1575 train_time:83423ms step_avg:59.04ms step:1414/1575 train_time:83509ms step_avg:59.06ms step:1415/1575 train_time:83600ms step_avg:59.08ms step:1416/1575 train_time:83686ms step_avg:59.10ms step:1417/1575 train_time:83775ms step_avg:59.12ms step:1418/1575 train_time:83861ms step_avg:59.14ms step:1419/1575 train_time:83949ms step_avg:59.16ms step:1420/1575 train_time:84036ms step_avg:59.18ms step:1421/1575 train_time:84125ms step_avg:59.20ms step:1422/1575 train_time:84211ms step_avg:59.22ms step:1423/1575 train_time:84300ms step_avg:59.24ms step:1424/1575 train_time:84385ms step_avg:59.26ms step:1425/1575 train_time:84474ms step_avg:59.28ms step:1426/1575 train_time:84559ms step_avg:59.30ms step:1427/1575 train_time:84650ms step_avg:59.32ms step:1428/1575 train_time:84736ms step_avg:59.34ms step:1429/1575 train_time:84825ms step_avg:59.36ms step:1430/1575 train_time:84911ms step_avg:59.38ms step:1431/1575 train_time:85000ms step_avg:59.40ms step:1432/1575 train_time:85085ms step_avg:59.42ms step:1433/1575 train_time:85175ms step_avg:59.44ms step:1434/1575 train_time:85262ms step_avg:59.46ms step:1435/1575 train_time:85351ms step_avg:59.48ms step:1436/1575 train_time:85440ms step_avg:59.50ms step:1437/1575 train_time:85528ms step_avg:59.52ms step:1438/1575 train_time:85613ms step_avg:59.54ms step:1439/1575 train_time:85703ms step_avg:59.56ms step:1440/1575 train_time:85788ms step_avg:59.57ms step:1441/1575 train_time:85877ms step_avg:59.60ms step:1442/1575 train_time:85963ms step_avg:59.61ms step:1443/1575 train_time:86052ms step_avg:59.63ms step:1444/1575 train_time:86138ms step_avg:59.65ms step:1445/1575 train_time:86227ms step_avg:59.67ms step:1446/1575 train_time:86313ms step_avg:59.69ms step:1447/1575 train_time:86401ms step_avg:59.71ms step:1448/1575 train_time:86487ms step_avg:59.73ms step:1449/1575 train_time:86576ms step_avg:59.75ms step:1450/1575 train_time:86663ms step_avg:59.77ms step:1451/1575 train_time:86752ms step_avg:59.79ms step:1452/1575 train_time:86838ms step_avg:59.81ms step:1453/1575 train_time:86928ms step_avg:59.83ms step:1454/1575 train_time:87014ms step_avg:59.84ms step:1455/1575 train_time:87103ms step_avg:59.86ms step:1456/1575 train_time:87189ms step_avg:59.88ms step:1457/1575 train_time:87278ms step_avg:59.90ms step:1458/1575 train_time:87364ms step_avg:59.92ms step:1459/1575 train_time:87454ms step_avg:59.94ms step:1460/1575 train_time:87541ms step_avg:59.96ms step:1461/1575 train_time:87630ms step_avg:59.98ms step:1462/1575 train_time:87716ms step_avg:60.00ms step:1463/1575 train_time:87806ms step_avg:60.02ms step:1464/1575 train_time:87891ms step_avg:60.04ms step:1465/1575 train_time:87981ms step_avg:60.06ms step:1466/1575 train_time:88066ms step_avg:60.07ms step:1467/1575 train_time:88155ms step_avg:60.09ms step:1468/1575 train_time:88242ms step_avg:60.11ms step:1469/1575 train_time:88331ms step_avg:60.13ms step:1470/1575 train_time:88420ms step_avg:60.15ms step:1471/1575 train_time:88508ms step_avg:60.17ms step:1472/1575 train_time:88593ms step_avg:60.19ms step:1473/1575 train_time:88682ms step_avg:60.20ms step:1474/1575 train_time:88767ms step_avg:60.22ms step:1475/1575 train_time:88856ms step_avg:60.24ms step:1476/1575 train_time:88943ms step_avg:60.26ms step:1477/1575 train_time:89033ms step_avg:60.28ms step:1478/1575 train_time:89118ms step_avg:60.30ms step:1479/1575 train_time:89208ms step_avg:60.32ms step:1480/1575 train_time:89293ms step_avg:60.33ms step:1481/1575 train_time:89383ms step_avg:60.35ms step:1482/1575 train_time:89468ms step_avg:60.37ms step:1483/1575 train_time:89557ms step_avg:60.39ms step:1484/1575 train_time:89644ms step_avg:60.41ms step:1485/1575 train_time:89733ms step_avg:60.43ms step:1486/1575 train_time:89818ms step_avg:60.44ms step:1487/1575 train_time:89908ms step_avg:60.46ms step:1488/1575 train_time:89994ms step_avg:60.48ms step:1489/1575 train_time:90083ms step_avg:60.50ms step:1490/1575 train_time:90169ms step_avg:60.52ms step:1491/1575 train_time:90259ms step_avg:60.54ms step:1492/1575 train_time:90345ms step_avg:60.55ms step:1493/1575 train_time:90435ms step_avg:60.57ms step:1494/1575 train_time:90521ms step_avg:60.59ms step:1495/1575 train_time:90613ms step_avg:60.61ms step:1496/1575 train_time:90696ms step_avg:60.63ms step:1497/1575 train_time:90786ms step_avg:60.65ms step:1498/1575 train_time:90872ms step_avg:60.66ms step:1499/1575 train_time:90961ms step_avg:60.68ms step:1500/1575 train_time:91047ms step_avg:60.70ms step:1500/1575 val_loss:3.3006 train_time:91119ms step_avg:60.75ms step:1501/1575 train_time:91139ms step_avg:60.72ms step:1502/1575 train_time:91226ms step_avg:60.74ms step:1503/1575 train_time:91318ms step_avg:60.76ms step:1504/1575 train_time:91404ms step_avg:60.77ms step:1505/1575 train_time:91493ms step_avg:60.79ms step:1506/1575 train_time:91577ms step_avg:60.81ms step:1507/1575 train_time:91665ms step_avg:60.83ms step:1508/1575 train_time:91750ms step_avg:60.84ms step:1509/1575 train_time:91839ms step_avg:60.86ms step:1510/1575 train_time:91925ms step_avg:60.88ms step:1511/1575 train_time:92014ms step_avg:60.90ms step:1512/1575 train_time:92100ms step_avg:60.91ms step:1513/1575 train_time:92191ms step_avg:60.93ms step:1514/1575 train_time:92279ms step_avg:60.95ms step:1515/1575 train_time:92369ms step_avg:60.97ms step:1516/1575 train_time:92456ms step_avg:60.99ms step:1517/1575 train_time:92546ms step_avg:61.01ms step:1518/1575 train_time:92630ms step_avg:61.02ms step:1519/1575 train_time:92719ms step_avg:61.04ms step:1520/1575 train_time:92804ms step_avg:61.06ms step:1521/1575 train_time:92892ms step_avg:61.07ms step:1522/1575 train_time:92977ms step_avg:61.09ms step:1523/1575 train_time:93067ms step_avg:61.11ms step:1524/1575 train_time:93153ms step_avg:61.12ms step:1525/1575 train_time:93245ms step_avg:61.14ms step:1526/1575 train_time:93331ms step_avg:61.16ms step:1527/1575 train_time:93421ms step_avg:61.18ms step:1528/1575 train_time:93507ms step_avg:61.20ms step:1529/1575 train_time:93596ms step_avg:61.21ms step:1530/1575 train_time:93682ms step_avg:61.23ms step:1531/1575 train_time:93770ms step_avg:61.25ms step:1532/1575 train_time:93856ms step_avg:61.26ms step:1533/1575 train_time:93945ms step_avg:61.28ms step:1534/1575 train_time:94031ms step_avg:61.30ms step:1535/1575 train_time:94120ms step_avg:61.32ms step:1536/1575 train_time:94214ms step_avg:61.34ms step:1537/1575 train_time:94303ms step_avg:61.36ms step:1538/1575 train_time:94388ms step_avg:61.37ms step:1539/1575 train_time:94478ms step_avg:61.39ms step:1540/1575 train_time:94564ms step_avg:61.41ms step:1541/1575 train_time:94653ms step_avg:61.42ms step:1542/1575 train_time:94739ms step_avg:61.44ms step:1543/1575 train_time:94831ms step_avg:61.46ms step:1544/1575 train_time:94915ms step_avg:61.47ms step:1545/1575 train_time:95005ms step_avg:61.49ms step:1546/1575 train_time:95091ms step_avg:61.51ms step:1547/1575 train_time:95181ms step_avg:61.53ms step:1548/1575 train_time:95268ms step_avg:61.54ms step:1549/1575 train_time:95357ms step_avg:61.56ms step:1550/1575 train_time:95444ms step_avg:61.58ms step:1551/1575 train_time:95534ms step_avg:61.59ms step:1552/1575 train_time:95620ms step_avg:61.61ms step:1553/1575 train_time:95709ms step_avg:61.63ms step:1554/1575 train_time:95795ms step_avg:61.64ms step:1555/1575 train_time:95885ms step_avg:61.66ms step:1556/1575 train_time:95970ms step_avg:61.68ms step:1557/1575 train_time:96060ms step_avg:61.70ms step:1558/1575 train_time:96146ms step_avg:61.71ms step:1559/1575 train_time:96236ms step_avg:61.73ms step:1560/1575 train_time:96323ms step_avg:61.75ms step:1561/1575 train_time:96413ms step_avg:61.76ms step:1562/1575 train_time:96499ms step_avg:61.78ms step:1563/1575 train_time:96589ms step_avg:61.80ms step:1564/1575 train_time:96674ms step_avg:61.81ms step:1565/1575 train_time:96764ms step_avg:61.83ms step:1566/1575 train_time:96850ms step_avg:61.85ms step:1567/1575 train_time:96942ms step_avg:61.86ms step:1568/1575 train_time:97027ms step_avg:61.88ms step:1569/1575 train_time:97121ms step_avg:61.90ms step:1570/1575 train_time:97205ms step_avg:61.91ms step:1571/1575 train_time:97295ms step_avg:61.93ms step:1572/1575 train_time:97381ms step_avg:61.95ms step:1573/1575 train_time:97470ms step_avg:61.96ms step:1574/1575 train_time:97557ms step_avg:61.98ms step:1575/1575 train_time:97647ms step_avg:62.00ms step:1575/1575 val_loss:3.2784 train_time:97715ms step_avg:62.04ms peak memory allocated: 31016 MiB reserved: 46998 MiB