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:13:29 2026 +-----------------------------------------------------------------------------------------+ | NVIDIA-SMI 570.148.08 Driver Version: 570.148.08 CUDA Version: 12.8 | |-----------------------------------------+------------------------+----------------------+ | GPU Name Persistence-M | Bus-Id Disp.A | Volatile Uncorr. ECC | | Fan Temp Perf Pwr:Usage/Cap | Memory-Usage | GPU-Util Compute M. | | | | MIG M. | |=========================================+========================+======================| | 0 NVIDIA H100 80GB HBM3 On | 00000000:61:00.0 Off | 0 | | N/A 35C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:62:00.0 Off | 0 | | N/A 39C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 41C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 35C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:6A:00.0 Off | 0 | | N/A 35C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 42C P0 129W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 39C P0 124W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 36C P0 118W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 233856 C /usr/bin/python3 1510MiB | | 1 N/A N/A 233857 C /usr/bin/python3 1510MiB | | 2 N/A N/A 233858 C /usr/bin/python3 1510MiB | | 3 N/A N/A 233859 C /usr/bin/python3 1510MiB | | 4 N/A N/A 233860 C /usr/bin/python3 1510MiB | | 5 N/A N/A 233861 C /usr/bin/python3 1510MiB | | 6 N/A N/A 233862 C /usr/bin/python3 1510MiB | | 7 N/A N/A 233863 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.8299 train_time:0ms step_avg:0.04ms step:1/1575 train_time:89ms step_avg:88.62ms step:2/1575 train_time:114ms step_avg:56.77ms step:3/1575 train_time:135ms step_avg:44.87ms step:4/1575 train_time:162ms step_avg:40.55ms step:5/1575 train_time:193ms step_avg:38.51ms step:6/1575 train_time:289ms step_avg:48.08ms step:7/1575 train_time:307ms step_avg:43.87ms step:8/1575 train_time:336ms step_avg:42.04ms step:9/1575 train_time:367ms step_avg:40.75ms step:10/1575 train_time:405ms step_avg:40.53ms step:11/1575 train_time:436ms step_avg:39.64ms step:12/1575 train_time:475ms step_avg:39.57ms step:13/1575 train_time:506ms step_avg:38.92ms step:14/1575 train_time:545ms step_avg:38.91ms step:15/1575 train_time:576ms step_avg:38.39ms step:16/1575 train_time:614ms step_avg:38.40ms step:17/1575 train_time:646ms step_avg:37.98ms step:18/1575 train_time:685ms step_avg:38.03ms step:19/1575 train_time:716ms step_avg:37.66ms step:20/1575 train_time:754ms step_avg:37.72ms step:21/1575 train_time:786ms step_avg:37.41ms step:22/1575 train_time:825ms step_avg:37.49ms step:23/1575 train_time:856ms step_avg:37.21ms step:24/1575 train_time:895ms step_avg:37.28ms step:25/1575 train_time:925ms step_avg:37.02ms step:26/1575 train_time:964ms step_avg:37.08ms step:27/1575 train_time:995ms step_avg:36.86ms step:28/1575 train_time:1034ms step_avg:36.92ms step:29/1575 train_time:1065ms step_avg:36.72ms step:30/1575 train_time:1104ms step_avg:36.80ms step:31/1575 train_time:1135ms step_avg:36.61ms step:32/1575 train_time:1174ms step_avg:36.68ms step:33/1575 train_time:1204ms step_avg:36.50ms step:34/1575 train_time:1243ms step_avg:36.57ms step:35/1575 train_time:1274ms step_avg:36.41ms step:36/1575 train_time:1313ms step_avg:36.47ms step:37/1575 train_time:1344ms step_avg:36.32ms step:38/1575 train_time:1382ms step_avg:36.38ms step:39/1575 train_time:1414ms step_avg:36.25ms step:40/1575 train_time:1452ms step_avg:36.31ms step:41/1575 train_time:1483ms step_avg:36.18ms step:42/1575 train_time:1522ms step_avg:36.25ms step:43/1575 train_time:1553ms step_avg:36.12ms step:44/1575 train_time:1592ms step_avg:36.18ms step:45/1575 train_time:1623ms step_avg:36.07ms step:46/1575 train_time:1663ms step_avg:36.15ms step:47/1575 train_time:1693ms step_avg:36.01ms step:48/1575 train_time:1731ms step_avg:36.06ms step:49/1575 train_time:1762ms step_avg:35.97ms step:50/1575 train_time:1801ms step_avg:36.02ms step:51/1575 train_time:1832ms step_avg:35.91ms step:52/1575 train_time:1870ms step_avg:35.96ms step:53/1575 train_time:1901ms step_avg:35.87ms step:54/1575 train_time:1939ms step_avg:35.92ms step:55/1575 train_time:1970ms step_avg:35.83ms step:56/1575 train_time:2009ms step_avg:35.88ms step:57/1575 train_time:2040ms step_avg:35.80ms step:58/1575 train_time:2079ms step_avg:35.85ms step:59/1575 train_time:2110ms step_avg:35.77ms step:60/1575 train_time:2149ms step_avg:35.82ms step:61/1575 train_time:2180ms step_avg:35.74ms step:62/1575 train_time:2218ms step_avg:35.78ms step:63/1575 train_time:2249ms step_avg:35.70ms step:64/1575 train_time:2288ms step_avg:35.76ms step:65/1575 train_time:2319ms step_avg:35.68ms step:66/1575 train_time:2357ms step_avg:35.72ms step:67/1575 train_time:2388ms step_avg:35.65ms step:68/1575 train_time:2427ms step_avg:35.69ms step:69/1575 train_time:2458ms step_avg:35.63ms step:70/1575 train_time:2497ms step_avg:35.68ms step:71/1575 train_time:2528ms step_avg:35.61ms step:72/1575 train_time:2567ms step_avg:35.65ms step:73/1575 train_time:2597ms step_avg:35.58ms step:74/1575 train_time:2636ms step_avg:35.63ms step:75/1575 train_time:2667ms step_avg:35.56ms step:76/1575 train_time:2707ms step_avg:35.61ms step:77/1575 train_time:2738ms step_avg:35.55ms step:78/1575 train_time:2776ms step_avg:35.59ms step:79/1575 train_time:2807ms step_avg:35.54ms step:80/1575 train_time:2846ms step_avg:35.58ms step:81/1575 train_time:2878ms step_avg:35.53ms step:82/1575 train_time:2916ms step_avg:35.56ms step:83/1575 train_time:2947ms step_avg:35.51ms step:84/1575 train_time:2987ms step_avg:35.56ms step:85/1575 train_time:3018ms step_avg:35.50ms step:86/1575 train_time:3056ms step_avg:35.54ms step:87/1575 train_time:3087ms step_avg:35.49ms step:88/1575 train_time:3126ms step_avg:35.53ms step:89/1575 train_time:3157ms step_avg:35.48ms step:90/1575 train_time:3196ms step_avg:35.51ms step:91/1575 train_time:3227ms step_avg:35.46ms step:92/1575 train_time:3268ms step_avg:35.52ms step:93/1575 train_time:3296ms step_avg:35.45ms step:94/1575 train_time:3336ms step_avg:35.48ms step:95/1575 train_time:3366ms step_avg:35.43ms step:96/1575 train_time:3405ms step_avg:35.46ms step:97/1575 train_time:3435ms step_avg:35.41ms step:98/1575 train_time:3474ms step_avg:35.45ms step:99/1575 train_time:3505ms step_avg:35.41ms step:100/1575 train_time:3544ms step_avg:35.44ms step:101/1575 train_time:3575ms step_avg:35.39ms step:102/1575 train_time:3613ms step_avg:35.42ms step:103/1575 train_time:3644ms step_avg:35.38ms step:104/1575 train_time:3683ms step_avg:35.42ms step:105/1575 train_time:3714ms step_avg:35.37ms step:106/1575 train_time:3753ms step_avg:35.40ms step:107/1575 train_time:3783ms step_avg:35.36ms step:108/1575 train_time:3823ms step_avg:35.39ms step:109/1575 train_time:3853ms step_avg:35.35ms step:110/1575 train_time:3892ms step_avg:35.38ms step:111/1575 train_time:3922ms step_avg:35.34ms step:112/1575 train_time:3961ms step_avg:35.37ms step:113/1575 train_time:3992ms step_avg:35.33ms step:114/1575 train_time:4031ms step_avg:35.36ms step:115/1575 train_time:4062ms step_avg:35.32ms step:116/1575 train_time:4101ms step_avg:35.35ms step:117/1575 train_time:4132ms step_avg:35.31ms step:118/1575 train_time:4170ms step_avg:35.34ms step:119/1575 train_time:4201ms step_avg:35.30ms step:120/1575 train_time:4240ms step_avg:35.33ms step:121/1575 train_time:4270ms step_avg:35.29ms step:122/1575 train_time:4309ms step_avg:35.32ms step:123/1575 train_time:4340ms step_avg:35.28ms step:124/1575 train_time:4378ms step_avg:35.31ms step:125/1575 train_time:4409ms step_avg:35.27ms step:126/1575 train_time:4448ms step_avg:35.30ms step:127/1575 train_time:4479ms step_avg:35.27ms step:128/1575 train_time:4517ms step_avg:35.29ms step:129/1575 train_time:4548ms step_avg:35.26ms step:130/1575 train_time:4587ms step_avg:35.29ms step:131/1575 train_time:4618ms step_avg:35.25ms step:132/1575 train_time:4657ms step_avg:35.28ms step:133/1575 train_time:4687ms step_avg:35.24ms step:134/1575 train_time:4726ms step_avg:35.27ms step:135/1575 train_time:4757ms step_avg:35.24ms step:136/1575 train_time:4796ms step_avg:35.26ms step:137/1575 train_time:4827ms step_avg:35.23ms step:138/1575 train_time:4865ms step_avg:35.26ms step:139/1575 train_time:4897ms step_avg:35.23ms step:140/1575 train_time:4936ms step_avg:35.25ms step:141/1575 train_time:4966ms step_avg:35.22ms step:142/1575 train_time:5005ms step_avg:35.25ms step:143/1575 train_time:5036ms step_avg:35.22ms step:144/1575 train_time:5074ms step_avg:35.24ms step:145/1575 train_time:5105ms step_avg:35.21ms step:146/1575 train_time:5145ms step_avg:35.24ms step:147/1575 train_time:5176ms step_avg:35.21ms step:148/1575 train_time:5214ms step_avg:35.23ms step:149/1575 train_time:5246ms step_avg:35.21ms step:150/1575 train_time:5284ms step_avg:35.23ms step:151/1575 train_time:5315ms step_avg:35.20ms step:152/1575 train_time:5354ms step_avg:35.22ms step:153/1575 train_time:5385ms step_avg:35.20ms step:154/1575 train_time:5423ms step_avg:35.22ms step:155/1575 train_time:5454ms step_avg:35.19ms step:156/1575 train_time:5493ms step_avg:35.21ms step:157/1575 train_time:5524ms step_avg:35.18ms step:158/1575 train_time:5563ms step_avg:35.21ms step:159/1575 train_time:5594ms step_avg:35.18ms step:160/1575 train_time:5634ms step_avg:35.21ms step:161/1575 train_time:5664ms step_avg:35.18ms step:162/1575 train_time:5703ms step_avg:35.20ms step:163/1575 train_time:5734ms step_avg:35.18ms step:164/1575 train_time:5772ms step_avg:35.20ms step:165/1575 train_time:5803ms step_avg:35.17ms step:166/1575 train_time:5841ms step_avg:35.19ms step:167/1575 train_time:5872ms step_avg:35.16ms step:168/1575 train_time:5911ms step_avg:35.19ms step:169/1575 train_time:5942ms step_avg:35.16ms step:170/1575 train_time:5980ms step_avg:35.18ms step:171/1575 train_time:6011ms step_avg:35.15ms step:172/1575 train_time:6050ms step_avg:35.17ms step:173/1575 train_time:6080ms step_avg:35.15ms step:174/1575 train_time:6119ms step_avg:35.16ms step:175/1575 train_time:6150ms step_avg:35.14ms step:176/1575 train_time:6189ms step_avg:35.17ms step:177/1575 train_time:6220ms step_avg:35.14ms step:178/1575 train_time:6258ms step_avg:35.16ms step:179/1575 train_time:6289ms step_avg:35.13ms step:180/1575 train_time:6327ms step_avg:35.15ms step:181/1575 train_time:6358ms step_avg:35.13ms step:182/1575 train_time:6396ms step_avg:35.14ms step:183/1575 train_time:6427ms step_avg:35.12ms step:184/1575 train_time:6466ms step_avg:35.14ms step:185/1575 train_time:6496ms step_avg:35.11ms step:186/1575 train_time:6535ms step_avg:35.13ms step:187/1575 train_time:6566ms step_avg:35.11ms step:188/1575 train_time:6605ms step_avg:35.13ms step:189/1575 train_time:6635ms step_avg:35.11ms step:190/1575 train_time:6674ms step_avg:35.13ms step:191/1575 train_time:6705ms step_avg:35.10ms step:192/1575 train_time:6743ms step_avg:35.12ms step:193/1575 train_time:6774ms step_avg:35.10ms step:194/1575 train_time:6813ms step_avg:35.12ms step:195/1575 train_time:6843ms step_avg:35.09ms step:196/1575 train_time:6882ms step_avg:35.11ms step:197/1575 train_time:6913ms step_avg:35.09ms step:198/1575 train_time:6951ms step_avg:35.11ms step:199/1575 train_time:6982ms step_avg:35.09ms step:200/1575 train_time:7021ms step_avg:35.10ms step:201/1575 train_time:7051ms step_avg:35.08ms step:202/1575 train_time:7090ms step_avg:35.10ms step:203/1575 train_time:7120ms step_avg:35.08ms step:204/1575 train_time:7159ms step_avg:35.09ms step:205/1575 train_time:7190ms step_avg:35.07ms step:206/1575 train_time:7229ms step_avg:35.09ms step:207/1575 train_time:7260ms step_avg:35.07ms step:208/1575 train_time:7298ms step_avg:35.09ms step:209/1575 train_time:7329ms step_avg:35.07ms step:210/1575 train_time:7367ms step_avg:35.08ms step:211/1575 train_time:7398ms step_avg:35.06ms step:212/1575 train_time:7437ms step_avg:35.08ms step:213/1575 train_time:7468ms step_avg:35.06ms step:214/1575 train_time:7506ms step_avg:35.08ms step:215/1575 train_time:7537ms step_avg:35.06ms step:216/1575 train_time:7576ms step_avg:35.07ms step:217/1575 train_time:7607ms step_avg:35.05ms step:218/1575 train_time:7645ms step_avg:35.07ms step:219/1575 train_time:7676ms step_avg:35.05ms step:220/1575 train_time:7715ms step_avg:35.07ms step:221/1575 train_time:7746ms step_avg:35.05ms step:222/1575 train_time:7785ms step_avg:35.07ms step:223/1575 train_time:7816ms step_avg:35.05ms step:224/1575 train_time:7855ms step_avg:35.06ms step:225/1575 train_time:7885ms step_avg:35.05ms step:226/1575 train_time:7924ms step_avg:35.06ms step:227/1575 train_time:7955ms step_avg:35.04ms step:228/1575 train_time:7994ms step_avg:35.06ms step:229/1575 train_time:8024ms step_avg:35.04ms step:230/1575 train_time:8063ms step_avg:35.06ms step:231/1575 train_time:8094ms step_avg:35.04ms step:232/1575 train_time:8132ms step_avg:35.05ms step:233/1575 train_time:8163ms step_avg:35.04ms step:234/1575 train_time:8202ms step_avg:35.05ms step:235/1575 train_time:8232ms step_avg:35.03ms step:236/1575 train_time:8272ms step_avg:35.05ms step:237/1575 train_time:8303ms step_avg:35.03ms step:238/1575 train_time:8342ms step_avg:35.05ms step:239/1575 train_time:8372ms step_avg:35.03ms step:240/1575 train_time:8411ms step_avg:35.05ms step:241/1575 train_time:8441ms step_avg:35.03ms step:242/1575 train_time:8480ms step_avg:35.04ms step:243/1575 train_time:8511ms step_avg:35.02ms step:244/1575 train_time:8550ms step_avg:35.04ms step:245/1575 train_time:8581ms step_avg:35.02ms step:246/1575 train_time:8619ms step_avg:35.04ms step:247/1575 train_time:8650ms step_avg:35.02ms step:248/1575 train_time:8689ms step_avg:35.03ms step:249/1575 train_time:8719ms step_avg:35.02ms step:250/1575 train_time:8758ms step_avg:35.03ms step:250/1575 val_loss:4.5772 train_time:8806ms step_avg:35.22ms step:251/1575 train_time:8827ms step_avg:35.17ms step:252/1575 train_time:8848ms step_avg:35.11ms step:253/1575 train_time:8865ms step_avg:35.04ms step:254/1575 train_time:8900ms step_avg:35.04ms step:255/1575 train_time:8933ms step_avg:35.03ms step:256/1575 train_time:8973ms step_avg:35.05ms step:257/1575 train_time:9004ms step_avg:35.04ms step:258/1575 train_time:9043ms step_avg:35.05ms step:259/1575 train_time:9074ms step_avg:35.04ms step:260/1575 train_time:9113ms step_avg:35.05ms step:261/1575 train_time:9144ms step_avg:35.04ms step:262/1575 train_time:9183ms step_avg:35.05ms step:263/1575 train_time:9213ms step_avg:35.03ms step:264/1575 train_time:9252ms step_avg:35.05ms step:265/1575 train_time:9283ms step_avg:35.03ms step:266/1575 train_time:9322ms step_avg:35.05ms step:267/1575 train_time:9353ms step_avg:35.03ms step:268/1575 train_time:9392ms step_avg:35.04ms step:269/1575 train_time:9422ms step_avg:35.03ms step:270/1575 train_time:9461ms step_avg:35.04ms step:271/1575 train_time:9492ms step_avg:35.02ms step:272/1575 train_time:9530ms step_avg:35.04ms step:273/1575 train_time:9561ms step_avg:35.02ms step:274/1575 train_time:9599ms step_avg:35.03ms step:275/1575 train_time:9630ms step_avg:35.02ms step:276/1575 train_time:9669ms step_avg:35.03ms step:277/1575 train_time:9700ms step_avg:35.02ms step:278/1575 train_time:9738ms step_avg:35.03ms step:279/1575 train_time:9769ms step_avg:35.01ms step:280/1575 train_time:9807ms step_avg:35.03ms step:281/1575 train_time:9838ms step_avg:35.01ms step:282/1575 train_time:9877ms step_avg:35.02ms step:283/1575 train_time:9908ms step_avg:35.01ms step:284/1575 train_time:9946ms step_avg:35.02ms step:285/1575 train_time:9977ms step_avg:35.01ms step:286/1575 train_time:10016ms step_avg:35.02ms step:287/1575 train_time:10048ms step_avg:35.01ms step:288/1575 train_time:10086ms step_avg:35.02ms step:289/1575 train_time:10117ms step_avg:35.01ms step:290/1575 train_time:10156ms step_avg:35.02ms step:291/1575 train_time:10187ms step_avg:35.01ms step:292/1575 train_time:10225ms step_avg:35.02ms step:293/1575 train_time:10256ms step_avg:35.00ms step:294/1575 train_time:10294ms step_avg:35.01ms step:295/1575 train_time:10325ms step_avg:35.00ms step:296/1575 train_time:10363ms step_avg:35.01ms step:297/1575 train_time:10394ms step_avg:35.00ms step:298/1575 train_time:10433ms step_avg:35.01ms step:299/1575 train_time:10463ms step_avg:34.99ms step:300/1575 train_time:10502ms step_avg:35.01ms step:301/1575 train_time:10532ms step_avg:34.99ms step:302/1575 train_time:10571ms step_avg:35.00ms step:303/1575 train_time:10602ms step_avg:34.99ms step:304/1575 train_time:10640ms step_avg:35.00ms step:305/1575 train_time:10671ms step_avg:34.99ms step:306/1575 train_time:10710ms step_avg:35.00ms step:307/1575 train_time:10741ms step_avg:34.99ms step:308/1575 train_time:10779ms step_avg:35.00ms step:309/1575 train_time:10810ms step_avg:34.99ms step:310/1575 train_time:10849ms step_avg:35.00ms step:311/1575 train_time:10879ms step_avg:34.98ms step:312/1575 train_time:10918ms step_avg:34.99ms step:313/1575 train_time:10948ms step_avg:34.98ms step:314/1575 train_time:10987ms step_avg:34.99ms step:315/1575 train_time:11017ms step_avg:34.98ms step:316/1575 train_time:11056ms step_avg:34.99ms step:317/1575 train_time:11087ms step_avg:34.97ms step:318/1575 train_time:11125ms step_avg:34.99ms step:319/1575 train_time:11156ms step_avg:34.97ms step:320/1575 train_time:11195ms step_avg:34.98ms step:321/1575 train_time:11225ms step_avg:34.97ms step:322/1575 train_time:11264ms step_avg:34.98ms step:323/1575 train_time:11294ms step_avg:34.97ms step:324/1575 train_time:11333ms step_avg:34.98ms step:325/1575 train_time:11364ms step_avg:34.97ms step:326/1575 train_time:11403ms step_avg:34.98ms step:327/1575 train_time:11434ms step_avg:34.96ms step:328/1575 train_time:11472ms step_avg:34.98ms step:329/1575 train_time:11503ms step_avg:34.96ms step:330/1575 train_time:11542ms step_avg:34.97ms step:331/1575 train_time:11572ms step_avg:34.96ms step:332/1575 train_time:11611ms step_avg:34.97ms step:333/1575 train_time:11642ms step_avg:34.96ms step:334/1575 train_time:11680ms step_avg:34.97ms step:335/1575 train_time:11711ms step_avg:34.96ms step:336/1575 train_time:11750ms step_avg:34.97ms step:337/1575 train_time:11780ms step_avg:34.96ms step:338/1575 train_time:11819ms step_avg:34.97ms step:339/1575 train_time:11850ms step_avg:34.95ms step:340/1575 train_time:11888ms step_avg:34.96ms step:341/1575 train_time:11919ms step_avg:34.95ms step:342/1575 train_time:11957ms step_avg:34.96ms step:343/1575 train_time:11988ms step_avg:34.95ms step:344/1575 train_time:12026ms step_avg:34.96ms step:345/1575 train_time:12057ms step_avg:34.95ms step:346/1575 train_time:12096ms step_avg:34.96ms step:347/1575 train_time:12127ms step_avg:34.95ms step:348/1575 train_time:12165ms step_avg:34.96ms step:349/1575 train_time:12196ms step_avg:34.95ms step:350/1575 train_time:12234ms step_avg:34.96ms step:351/1575 train_time:12265ms step_avg:34.94ms step:352/1575 train_time:12303ms step_avg:34.95ms step:353/1575 train_time:12334ms step_avg:34.94ms step:354/1575 train_time:12373ms step_avg:34.95ms step:355/1575 train_time:12404ms step_avg:34.94ms step:356/1575 train_time:12443ms step_avg:34.95ms step:357/1575 train_time:12473ms step_avg:34.94ms step:358/1575 train_time:12512ms step_avg:34.95ms step:359/1575 train_time:12543ms step_avg:34.94ms step:360/1575 train_time:12582ms step_avg:34.95ms step:361/1575 train_time:12612ms step_avg:34.94ms step:362/1575 train_time:12652ms step_avg:34.95ms step:363/1575 train_time:12682ms step_avg:34.94ms step:364/1575 train_time:12721ms step_avg:34.95ms step:365/1575 train_time:12751ms step_avg:34.94ms step:366/1575 train_time:12790ms step_avg:34.95ms step:367/1575 train_time:12821ms step_avg:34.93ms step:368/1575 train_time:12859ms step_avg:34.94ms step:369/1575 train_time:12890ms step_avg:34.93ms step:370/1575 train_time:12928ms step_avg:34.94ms step:371/1575 train_time:12959ms step_avg:34.93ms step:372/1575 train_time:12998ms step_avg:34.94ms step:373/1575 train_time:13029ms step_avg:34.93ms step:374/1575 train_time:13067ms step_avg:34.94ms step:375/1575 train_time:13098ms step_avg:34.93ms step:376/1575 train_time:13137ms step_avg:34.94ms step:377/1575 train_time:13167ms step_avg:34.93ms step:378/1575 train_time:13205ms step_avg:34.93ms step:379/1575 train_time:13236ms step_avg:34.92ms step:380/1575 train_time:13275ms step_avg:34.93ms step:381/1575 train_time:13306ms step_avg:34.92ms step:382/1575 train_time:13344ms step_avg:34.93ms step:383/1575 train_time:13375ms step_avg:34.92ms step:384/1575 train_time:13414ms step_avg:34.93ms step:385/1575 train_time:13445ms step_avg:34.92ms step:386/1575 train_time:13484ms step_avg:34.93ms step:387/1575 train_time:13514ms step_avg:34.92ms step:388/1575 train_time:13553ms step_avg:34.93ms step:389/1575 train_time:13584ms step_avg:34.92ms step:390/1575 train_time:13623ms step_avg:34.93ms step:391/1575 train_time:13653ms step_avg:34.92ms step:392/1575 train_time:13692ms step_avg:34.93ms step:393/1575 train_time:13723ms step_avg:34.92ms step:394/1575 train_time:13762ms step_avg:34.93ms step:395/1575 train_time:13792ms step_avg:34.92ms step:396/1575 train_time:13831ms step_avg:34.93ms step:397/1575 train_time:13862ms step_avg:34.92ms step:398/1575 train_time:13901ms step_avg:34.93ms step:399/1575 train_time:13931ms step_avg:34.92ms step:400/1575 train_time:13970ms step_avg:34.93ms step:401/1575 train_time:14001ms step_avg:34.92ms step:402/1575 train_time:14040ms step_avg:34.92ms step:403/1575 train_time:14071ms step_avg:34.91ms step:404/1575 train_time:14109ms step_avg:34.92ms step:405/1575 train_time:14140ms step_avg:34.91ms step:406/1575 train_time:14179ms step_avg:34.92ms step:407/1575 train_time:14209ms step_avg:34.91ms step:408/1575 train_time:14248ms step_avg:34.92ms step:409/1575 train_time:14279ms step_avg:34.91ms step:410/1575 train_time:14318ms step_avg:34.92ms step:411/1575 train_time:14348ms step_avg:34.91ms step:412/1575 train_time:14387ms step_avg:34.92ms step:413/1575 train_time:14417ms step_avg:34.91ms step:414/1575 train_time:14456ms step_avg:34.92ms step:415/1575 train_time:14487ms step_avg:34.91ms step:416/1575 train_time:14525ms step_avg:34.92ms step:417/1575 train_time:14556ms step_avg:34.91ms step:418/1575 train_time:14595ms step_avg:34.92ms step:419/1575 train_time:14625ms step_avg:34.90ms step:420/1575 train_time:14664ms step_avg:34.91ms step:421/1575 train_time:14695ms step_avg:34.90ms step:422/1575 train_time:14733ms step_avg:34.91ms step:423/1575 train_time:14764ms step_avg:34.90ms step:424/1575 train_time:14802ms step_avg:34.91ms step:425/1575 train_time:14833ms step_avg:34.90ms step:426/1575 train_time:14873ms step_avg:34.91ms step:427/1575 train_time:14903ms step_avg:34.90ms step:428/1575 train_time:14942ms step_avg:34.91ms step:429/1575 train_time:14972ms step_avg:34.90ms step:430/1575 train_time:15011ms step_avg:34.91ms step:431/1575 train_time:15042ms step_avg:34.90ms step:432/1575 train_time:15082ms step_avg:34.91ms step:433/1575 train_time:15112ms step_avg:34.90ms step:434/1575 train_time:15151ms step_avg:34.91ms step:435/1575 train_time:15181ms step_avg:34.90ms step:436/1575 train_time:15220ms step_avg:34.91ms step:437/1575 train_time:15251ms step_avg:34.90ms step:438/1575 train_time:15290ms step_avg:34.91ms step:439/1575 train_time:15320ms step_avg:34.90ms step:440/1575 train_time:15358ms step_avg:34.91ms step:441/1575 train_time:15389ms step_avg:34.90ms step:442/1575 train_time:15428ms step_avg:34.90ms step:443/1575 train_time:15458ms step_avg:34.89ms step:444/1575 train_time:15497ms step_avg:34.90ms step:445/1575 train_time:15528ms step_avg:34.89ms step:446/1575 train_time:15566ms step_avg:34.90ms step:447/1575 train_time:15597ms step_avg:34.89ms step:448/1575 train_time:15635ms step_avg:34.90ms step:449/1575 train_time:15666ms step_avg:34.89ms step:450/1575 train_time:15705ms step_avg:34.90ms step:451/1575 train_time:15736ms step_avg:34.89ms step:452/1575 train_time:15775ms step_avg:34.90ms step:453/1575 train_time:15805ms step_avg:34.89ms step:454/1575 train_time:15843ms step_avg:34.90ms step:455/1575 train_time:15874ms step_avg:34.89ms step:456/1575 train_time:15913ms step_avg:34.90ms step:457/1575 train_time:15944ms step_avg:34.89ms step:458/1575 train_time:15982ms step_avg:34.90ms step:459/1575 train_time:16013ms step_avg:34.89ms step:460/1575 train_time:16052ms step_avg:34.89ms step:461/1575 train_time:16083ms step_avg:34.89ms step:462/1575 train_time:16121ms step_avg:34.89ms step:463/1575 train_time:16152ms step_avg:34.89ms step:464/1575 train_time:16191ms step_avg:34.89ms step:465/1575 train_time:16222ms step_avg:34.89ms step:466/1575 train_time:16260ms step_avg:34.89ms step:467/1575 train_time:16291ms step_avg:34.88ms step:468/1575 train_time:16329ms step_avg:34.89ms step:469/1575 train_time:16360ms step_avg:34.88ms step:470/1575 train_time:16399ms step_avg:34.89ms step:471/1575 train_time:16430ms step_avg:34.88ms step:472/1575 train_time:16468ms step_avg:34.89ms step:473/1575 train_time:16499ms step_avg:34.88ms step:474/1575 train_time:16537ms step_avg:34.89ms step:475/1575 train_time:16568ms step_avg:34.88ms step:476/1575 train_time:16606ms step_avg:34.89ms step:477/1575 train_time:16637ms step_avg:34.88ms step:478/1575 train_time:16676ms step_avg:34.89ms step:479/1575 train_time:16706ms step_avg:34.88ms step:480/1575 train_time:16744ms step_avg:34.88ms step:481/1575 train_time:16775ms step_avg:34.88ms step:482/1575 train_time:16814ms step_avg:34.88ms step:483/1575 train_time:16845ms step_avg:34.88ms step:484/1575 train_time:16883ms step_avg:34.88ms step:485/1575 train_time:16914ms step_avg:34.87ms step:486/1575 train_time:16953ms step_avg:34.88ms step:487/1575 train_time:16984ms step_avg:34.88ms step:488/1575 train_time:17023ms step_avg:34.88ms step:489/1575 train_time:17054ms step_avg:34.87ms step:490/1575 train_time:17092ms step_avg:34.88ms step:491/1575 train_time:17123ms step_avg:34.87ms step:492/1575 train_time:17162ms step_avg:34.88ms step:493/1575 train_time:17192ms step_avg:34.87ms step:494/1575 train_time:17230ms step_avg:34.88ms step:495/1575 train_time:17261ms step_avg:34.87ms step:496/1575 train_time:17300ms step_avg:34.88ms step:497/1575 train_time:17330ms step_avg:34.87ms step:498/1575 train_time:17369ms step_avg:34.88ms step:499/1575 train_time:17400ms step_avg:34.87ms step:500/1575 train_time:17438ms step_avg:34.88ms step:500/1575 val_loss:4.2311 train_time:17486ms step_avg:34.97ms step:501/1575 train_time:17508ms step_avg:34.95ms step:502/1575 train_time:17528ms step_avg:34.92ms step:503/1575 train_time:17546ms step_avg:34.88ms step:504/1575 train_time:17580ms step_avg:34.88ms step:505/1575 train_time:17612ms step_avg:34.87ms step:506/1575 train_time:17651ms step_avg:34.88ms step:507/1575 train_time:17683ms step_avg:34.88ms step:508/1575 train_time:17722ms step_avg:34.88ms step:509/1575 train_time:17752ms step_avg:34.88ms step:510/1575 train_time:17791ms step_avg:34.88ms step:511/1575 train_time:17822ms step_avg:34.88ms step:512/1575 train_time:17860ms step_avg:34.88ms step:513/1575 train_time:17937ms step_avg:34.97ms step:514/1575 train_time:17990ms step_avg:35.00ms step:515/1575 train_time:18052ms step_avg:35.05ms step:516/1575 train_time:18111ms step_avg:35.10ms step:517/1575 train_time:18173ms step_avg:35.15ms step:518/1575 train_time:18232ms step_avg:35.20ms step:519/1575 train_time:18294ms step_avg:35.25ms step:520/1575 train_time:18353ms step_avg:35.29ms step:521/1575 train_time:18416ms step_avg:35.35ms step:522/1575 train_time:18475ms step_avg:35.39ms step:523/1575 train_time:18539ms step_avg:35.45ms step:524/1575 train_time:18600ms step_avg:35.50ms step:525/1575 train_time:18664ms step_avg:35.55ms step:526/1575 train_time:18725ms step_avg:35.60ms step:527/1575 train_time:18790ms step_avg:35.65ms step:528/1575 train_time:18857ms step_avg:35.71ms step:529/1575 train_time:18916ms step_avg:35.76ms step:530/1575 train_time:18974ms step_avg:35.80ms step:531/1575 train_time:19037ms step_avg:35.85ms step:532/1575 train_time:19096ms step_avg:35.89ms step:533/1575 train_time:19159ms step_avg:35.95ms step:534/1575 train_time:19219ms step_avg:35.99ms step:535/1575 train_time:19283ms step_avg:36.04ms step:536/1575 train_time:19342ms step_avg:36.09ms step:537/1575 train_time:19406ms step_avg:36.14ms step:538/1575 train_time:19465ms step_avg:36.18ms step:539/1575 train_time:19529ms step_avg:36.23ms step:540/1575 train_time:19589ms step_avg:36.28ms step:541/1575 train_time:19652ms step_avg:36.33ms step:542/1575 train_time:19713ms step_avg:36.37ms step:543/1575 train_time:19777ms step_avg:36.42ms step:544/1575 train_time:19837ms step_avg:36.46ms step:545/1575 train_time:19901ms step_avg:36.52ms step:546/1575 train_time:19960ms step_avg:36.56ms step:547/1575 train_time:20024ms step_avg:36.61ms step:548/1575 train_time:20083ms step_avg:36.65ms step:549/1575 train_time:20146ms step_avg:36.70ms step:550/1575 train_time:20205ms step_avg:36.74ms step:551/1575 train_time:20269ms step_avg:36.79ms step:552/1575 train_time:20328ms step_avg:36.83ms step:553/1575 train_time:20392ms step_avg:36.87ms step:554/1575 train_time:20451ms step_avg:36.91ms step:555/1575 train_time:20515ms step_avg:36.96ms step:556/1575 train_time:20573ms step_avg:37.00ms step:557/1575 train_time:20637ms step_avg:37.05ms step:558/1575 train_time:20695ms step_avg:37.09ms step:559/1575 train_time:20758ms step_avg:37.13ms step:560/1575 train_time:20818ms step_avg:37.18ms step:561/1575 train_time:20882ms step_avg:37.22ms step:562/1575 train_time:20941ms step_avg:37.26ms step:563/1575 train_time:21005ms step_avg:37.31ms step:564/1575 train_time:21064ms step_avg:37.35ms step:565/1575 train_time:21127ms step_avg:37.39ms step:566/1575 train_time:21187ms step_avg:37.43ms step:567/1575 train_time:21250ms step_avg:37.48ms step:568/1575 train_time:21310ms step_avg:37.52ms step:569/1575 train_time:21378ms step_avg:37.57ms step:570/1575 train_time:21434ms step_avg:37.60ms step:571/1575 train_time:21497ms step_avg:37.65ms step:572/1575 train_time:21556ms step_avg:37.69ms step:573/1575 train_time:21618ms step_avg:37.73ms step:574/1575 train_time:21677ms step_avg:37.76ms step:575/1575 train_time:21741ms step_avg:37.81ms step:576/1575 train_time:21800ms step_avg:37.85ms step:577/1575 train_time:21864ms step_avg:37.89ms step:578/1575 train_time:21923ms step_avg:37.93ms step:579/1575 train_time:21989ms step_avg:37.98ms step:580/1575 train_time:22048ms step_avg:38.01ms step:581/1575 train_time:22113ms step_avg:38.06ms step:582/1575 train_time:22169ms step_avg:38.09ms step:583/1575 train_time:22232ms step_avg:38.13ms step:584/1575 train_time:22291ms step_avg:38.17ms step:585/1575 train_time:22354ms step_avg:38.21ms step:586/1575 train_time:22413ms step_avg:38.25ms step:587/1575 train_time:22476ms step_avg:38.29ms step:588/1575 train_time:22536ms step_avg:38.33ms step:589/1575 train_time:22599ms step_avg:38.37ms step:590/1575 train_time:22658ms step_avg:38.40ms step:591/1575 train_time:22720ms step_avg:38.44ms step:592/1575 train_time:22780ms step_avg:38.48ms step:593/1575 train_time:22843ms step_avg:38.52ms step:594/1575 train_time:22903ms step_avg:38.56ms step:595/1575 train_time:22965ms step_avg:38.60ms step:596/1575 train_time:23026ms step_avg:38.63ms step:597/1575 train_time:23089ms step_avg:38.68ms step:598/1575 train_time:23148ms step_avg:38.71ms step:599/1575 train_time:23212ms step_avg:38.75ms step:600/1575 train_time:23271ms step_avg:38.79ms step:601/1575 train_time:23335ms step_avg:38.83ms step:602/1575 train_time:23394ms step_avg:38.86ms step:603/1575 train_time:23457ms step_avg:38.90ms step:604/1575 train_time:23517ms step_avg:38.94ms step:605/1575 train_time:23579ms step_avg:38.97ms step:606/1575 train_time:23639ms step_avg:39.01ms step:607/1575 train_time:23701ms step_avg:39.05ms step:608/1575 train_time:23760ms step_avg:39.08ms step:609/1575 train_time:23824ms step_avg:39.12ms step:610/1575 train_time:23884ms step_avg:39.15ms step:611/1575 train_time:23947ms step_avg:39.19ms step:612/1575 train_time:24007ms step_avg:39.23ms step:613/1575 train_time:24070ms step_avg:39.27ms step:614/1575 train_time:24130ms step_avg:39.30ms step:615/1575 train_time:24193ms step_avg:39.34ms step:616/1575 train_time:24253ms step_avg:39.37ms step:617/1575 train_time:24315ms step_avg:39.41ms step:618/1575 train_time:24375ms step_avg:39.44ms step:619/1575 train_time:24440ms step_avg:39.48ms step:620/1575 train_time:24498ms step_avg:39.51ms step:621/1575 train_time:24560ms step_avg:39.55ms step:622/1575 train_time:24620ms step_avg:39.58ms step:623/1575 train_time:24683ms step_avg:39.62ms step:624/1575 train_time:24742ms step_avg:39.65ms step:625/1575 train_time:24806ms step_avg:39.69ms step:626/1575 train_time:24866ms step_avg:39.72ms step:627/1575 train_time:24929ms step_avg:39.76ms step:628/1575 train_time:24988ms step_avg:39.79ms step:629/1575 train_time:25051ms step_avg:39.83ms step:630/1575 train_time:25110ms step_avg:39.86ms step:631/1575 train_time:25174ms step_avg:39.89ms step:632/1575 train_time:25233ms step_avg:39.93ms step:633/1575 train_time:25297ms step_avg:39.96ms step:634/1575 train_time:25356ms step_avg:39.99ms step:635/1575 train_time:25419ms step_avg:40.03ms step:636/1575 train_time:25479ms step_avg:40.06ms step:637/1575 train_time:25542ms step_avg:40.10ms step:638/1575 train_time:25602ms step_avg:40.13ms step:639/1575 train_time:25664ms step_avg:40.16ms step:640/1575 train_time:25724ms step_avg:40.19ms step:641/1575 train_time:25787ms step_avg:40.23ms step:642/1575 train_time:25847ms step_avg:40.26ms step:643/1575 train_time:25909ms step_avg:40.29ms step:644/1575 train_time:25969ms step_avg:40.32ms step:645/1575 train_time:26033ms step_avg:40.36ms step:646/1575 train_time:26092ms step_avg:40.39ms step:647/1575 train_time:26155ms step_avg:40.42ms step:648/1575 train_time:26214ms step_avg:40.45ms step:649/1575 train_time:26277ms step_avg:40.49ms step:650/1575 train_time:26336ms step_avg:40.52ms step:651/1575 train_time:26399ms step_avg:40.55ms step:652/1575 train_time:26459ms step_avg:40.58ms step:653/1575 train_time:26522ms step_avg:40.62ms step:654/1575 train_time:26581ms step_avg:40.64ms step:655/1575 train_time:26644ms step_avg:40.68ms step:656/1575 train_time:26703ms step_avg:40.71ms step:657/1575 train_time:26766ms step_avg:40.74ms step:658/1575 train_time:26825ms step_avg:40.77ms step:659/1575 train_time:26891ms step_avg:40.81ms step:660/1575 train_time:26950ms step_avg:40.83ms step:661/1575 train_time:27013ms step_avg:40.87ms step:662/1575 train_time:27071ms step_avg:40.89ms step:663/1575 train_time:27137ms step_avg:40.93ms step:664/1575 train_time:27196ms step_avg:40.96ms step:665/1575 train_time:27258ms step_avg:40.99ms step:666/1575 train_time:27317ms step_avg:41.02ms step:667/1575 train_time:27380ms step_avg:41.05ms step:668/1575 train_time:27440ms step_avg:41.08ms step:669/1575 train_time:27503ms step_avg:41.11ms step:670/1575 train_time:27562ms step_avg:41.14ms step:671/1575 train_time:27625ms step_avg:41.17ms step:672/1575 train_time:27685ms step_avg:41.20ms step:673/1575 train_time:27748ms step_avg:41.23ms step:674/1575 train_time:27807ms step_avg:41.26ms step:675/1575 train_time:27871ms step_avg:41.29ms step:676/1575 train_time:27930ms step_avg:41.32ms step:677/1575 train_time:27993ms step_avg:41.35ms step:678/1575 train_time:28053ms step_avg:41.38ms step:679/1575 train_time:28117ms step_avg:41.41ms step:680/1575 train_time:28176ms step_avg:41.43ms step:681/1575 train_time:28239ms step_avg:41.47ms step:682/1575 train_time:28298ms step_avg:41.49ms step:683/1575 train_time:28361ms step_avg:41.52ms step:684/1575 train_time:28420ms step_avg:41.55ms step:685/1575 train_time:28484ms step_avg:41.58ms step:686/1575 train_time:28543ms step_avg:41.61ms step:687/1575 train_time:28606ms step_avg:41.64ms step:688/1575 train_time:28666ms step_avg:41.67ms step:689/1575 train_time:28729ms step_avg:41.70ms step:690/1575 train_time:28789ms step_avg:41.72ms step:691/1575 train_time:28853ms step_avg:41.75ms step:692/1575 train_time:28912ms step_avg:41.78ms step:693/1575 train_time:28976ms step_avg:41.81ms step:694/1575 train_time:29035ms step_avg:41.84ms step:695/1575 train_time:29098ms step_avg:41.87ms step:696/1575 train_time:29157ms step_avg:41.89ms step:697/1575 train_time:29220ms step_avg:41.92ms step:698/1575 train_time:29280ms step_avg:41.95ms step:699/1575 train_time:29343ms step_avg:41.98ms step:700/1575 train_time:29401ms step_avg:42.00ms step:701/1575 train_time:29465ms step_avg:42.03ms step:702/1575 train_time:29525ms step_avg:42.06ms step:703/1575 train_time:29588ms step_avg:42.09ms step:704/1575 train_time:29647ms step_avg:42.11ms step:705/1575 train_time:29710ms step_avg:42.14ms step:706/1575 train_time:29769ms step_avg:42.17ms step:707/1575 train_time:29832ms step_avg:42.20ms step:708/1575 train_time:29892ms step_avg:42.22ms step:709/1575 train_time:29955ms step_avg:42.25ms step:710/1575 train_time:30015ms step_avg:42.28ms step:711/1575 train_time:30078ms step_avg:42.30ms step:712/1575 train_time:30138ms step_avg:42.33ms step:713/1575 train_time:30200ms step_avg:42.36ms step:714/1575 train_time:30259ms step_avg:42.38ms step:715/1575 train_time:30323ms step_avg:42.41ms step:716/1575 train_time:30382ms step_avg:42.43ms step:717/1575 train_time:30445ms step_avg:42.46ms step:718/1575 train_time:30504ms step_avg:42.48ms step:719/1575 train_time:30568ms step_avg:42.51ms step:720/1575 train_time:30628ms step_avg:42.54ms step:721/1575 train_time:30692ms step_avg:42.57ms step:722/1575 train_time:30750ms step_avg:42.59ms step:723/1575 train_time:30813ms step_avg:42.62ms step:724/1575 train_time:30873ms step_avg:42.64ms step:725/1575 train_time:30937ms step_avg:42.67ms step:726/1575 train_time:30996ms step_avg:42.69ms step:727/1575 train_time:31059ms step_avg:42.72ms step:728/1575 train_time:31118ms step_avg:42.74ms step:729/1575 train_time:31181ms step_avg:42.77ms step:730/1575 train_time:31240ms step_avg:42.79ms step:731/1575 train_time:31303ms step_avg:42.82ms step:732/1575 train_time:31362ms step_avg:42.84ms step:733/1575 train_time:31426ms step_avg:42.87ms step:734/1575 train_time:31485ms step_avg:42.89ms step:735/1575 train_time:31548ms step_avg:42.92ms step:736/1575 train_time:31609ms step_avg:42.95ms step:737/1575 train_time:31672ms step_avg:42.97ms step:738/1575 train_time:31731ms step_avg:43.00ms step:739/1575 train_time:31794ms step_avg:43.02ms step:740/1575 train_time:31853ms step_avg:43.05ms step:741/1575 train_time:31917ms step_avg:43.07ms step:742/1575 train_time:31976ms step_avg:43.09ms step:743/1575 train_time:32040ms step_avg:43.12ms step:744/1575 train_time:32099ms step_avg:43.14ms step:745/1575 train_time:32162ms step_avg:43.17ms step:746/1575 train_time:32225ms step_avg:43.20ms step:747/1575 train_time:32287ms step_avg:43.22ms step:748/1575 train_time:32346ms step_avg:43.24ms step:749/1575 train_time:32408ms step_avg:43.27ms step:750/1575 train_time:32467ms step_avg:43.29ms step:750/1575 val_loss:3.8816 train_time:32513ms step_avg:43.35ms step:751/1575 train_time:32534ms step_avg:43.32ms step:752/1575 train_time:32591ms step_avg:43.34ms step:753/1575 train_time:32658ms step_avg:43.37ms step:754/1575 train_time:32718ms step_avg:43.39ms step:755/1575 train_time:32782ms step_avg:43.42ms step:756/1575 train_time:32841ms step_avg:43.44ms step:757/1575 train_time:32904ms step_avg:43.47ms step:758/1575 train_time:32963ms step_avg:43.49ms step:759/1575 train_time:33026ms step_avg:43.51ms step:760/1575 train_time:33085ms step_avg:43.53ms step:761/1575 train_time:33150ms step_avg:43.56ms step:762/1575 train_time:33207ms step_avg:43.58ms step:763/1575 train_time:33270ms step_avg:43.60ms step:764/1575 train_time:33329ms step_avg:43.62ms step:765/1575 train_time:33392ms step_avg:43.65ms step:766/1575 train_time:33452ms step_avg:43.67ms step:767/1575 train_time:33516ms step_avg:43.70ms step:768/1575 train_time:33576ms step_avg:43.72ms step:769/1575 train_time:33640ms step_avg:43.75ms step:770/1575 train_time:33700ms step_avg:43.77ms step:771/1575 train_time:33764ms step_avg:43.79ms step:772/1575 train_time:33823ms step_avg:43.81ms step:773/1575 train_time:33887ms step_avg:43.84ms step:774/1575 train_time:33946ms step_avg:43.86ms step:775/1575 train_time:34010ms step_avg:43.88ms step:776/1575 train_time:34069ms step_avg:43.90ms step:777/1575 train_time:34132ms step_avg:43.93ms step:778/1575 train_time:34191ms step_avg:43.95ms step:779/1575 train_time:34254ms step_avg:43.97ms step:780/1575 train_time:34313ms step_avg:43.99ms step:781/1575 train_time:34376ms step_avg:44.01ms step:782/1575 train_time:34438ms step_avg:44.04ms step:783/1575 train_time:34499ms step_avg:44.06ms step:784/1575 train_time:34558ms step_avg:44.08ms step:785/1575 train_time:34621ms step_avg:44.10ms step:786/1575 train_time:34681ms step_avg:44.12ms step:787/1575 train_time:34745ms step_avg:44.15ms step:788/1575 train_time:34804ms step_avg:44.17ms step:789/1575 train_time:34868ms step_avg:44.19ms step:790/1575 train_time:34927ms step_avg:44.21ms step:791/1575 train_time:34990ms step_avg:44.24ms step:792/1575 train_time:35049ms step_avg:44.25ms step:793/1575 train_time:35113ms step_avg:44.28ms step:794/1575 train_time:35173ms step_avg:44.30ms step:795/1575 train_time:35236ms step_avg:44.32ms step:796/1575 train_time:35295ms step_avg:44.34ms step:797/1575 train_time:35357ms step_avg:44.36ms step:798/1575 train_time:35417ms step_avg:44.38ms step:799/1575 train_time:35480ms step_avg:44.41ms step:800/1575 train_time:35539ms step_avg:44.42ms step:801/1575 train_time:35603ms step_avg:44.45ms step:802/1575 train_time:35663ms step_avg:44.47ms step:803/1575 train_time:35726ms step_avg:44.49ms step:804/1575 train_time:35785ms step_avg:44.51ms step:805/1575 train_time:35849ms step_avg:44.53ms step:806/1575 train_time:35908ms step_avg:44.55ms step:807/1575 train_time:35971ms step_avg:44.57ms step:808/1575 train_time:36032ms step_avg:44.59ms step:809/1575 train_time:36095ms step_avg:44.62ms step:810/1575 train_time:36157ms step_avg:44.64ms step:811/1575 train_time:36220ms step_avg:44.66ms step:812/1575 train_time:36277ms step_avg:44.68ms step:813/1575 train_time:36340ms step_avg:44.70ms step:814/1575 train_time:36399ms step_avg:44.72ms step:815/1575 train_time:36464ms step_avg:44.74ms step:816/1575 train_time:36524ms step_avg:44.76ms step:817/1575 train_time:36586ms step_avg:44.78ms step:818/1575 train_time:36646ms step_avg:44.80ms step:819/1575 train_time:36709ms step_avg:44.82ms step:820/1575 train_time:36768ms step_avg:44.84ms step:821/1575 train_time:36833ms step_avg:44.86ms step:822/1575 train_time:36892ms step_avg:44.88ms step:823/1575 train_time:36955ms step_avg:44.90ms step:824/1575 train_time:37014ms step_avg:44.92ms step:825/1575 train_time:37078ms step_avg:44.94ms step:826/1575 train_time:37138ms step_avg:44.96ms step:827/1575 train_time:37201ms step_avg:44.98ms step:828/1575 train_time:37260ms step_avg:45.00ms step:829/1575 train_time:37324ms step_avg:45.02ms step:830/1575 train_time:37383ms step_avg:45.04ms step:831/1575 train_time:37447ms step_avg:45.06ms step:832/1575 train_time:37506ms step_avg:45.08ms step:833/1575 train_time:37569ms step_avg:45.10ms step:834/1575 train_time:37629ms step_avg:45.12ms step:835/1575 train_time:37692ms step_avg:45.14ms step:836/1575 train_time:37752ms step_avg:45.16ms step:837/1575 train_time:37815ms step_avg:45.18ms step:838/1575 train_time:37874ms step_avg:45.20ms step:839/1575 train_time:37937ms step_avg:45.22ms step:840/1575 train_time:37996ms step_avg:45.23ms step:841/1575 train_time:38060ms step_avg:45.26ms step:842/1575 train_time:38119ms step_avg:45.27ms step:843/1575 train_time:38183ms step_avg:45.29ms step:844/1575 train_time:38242ms step_avg:45.31ms step:845/1575 train_time:38306ms step_avg:45.33ms step:846/1575 train_time:38365ms step_avg:45.35ms step:847/1575 train_time:38429ms step_avg:45.37ms step:848/1575 train_time:38488ms step_avg:45.39ms step:849/1575 train_time:38553ms step_avg:45.41ms step:850/1575 train_time:38611ms step_avg:45.42ms step:851/1575 train_time:38675ms step_avg:45.45ms step:852/1575 train_time:38734ms step_avg:45.46ms step:853/1575 train_time:38798ms step_avg:45.48ms step:854/1575 train_time:38857ms step_avg:45.50ms step:855/1575 train_time:38920ms step_avg:45.52ms step:856/1575 train_time:38979ms step_avg:45.54ms step:857/1575 train_time:39042ms step_avg:45.56ms step:858/1575 train_time:39101ms step_avg:45.57ms step:859/1575 train_time:39165ms step_avg:45.59ms step:860/1575 train_time:39225ms step_avg:45.61ms step:861/1575 train_time:39288ms step_avg:45.63ms step:862/1575 train_time:39348ms step_avg:45.65ms step:863/1575 train_time:39411ms step_avg:45.67ms step:864/1575 train_time:39470ms step_avg:45.68ms step:865/1575 train_time:39534ms step_avg:45.70ms step:866/1575 train_time:39593ms step_avg:45.72ms step:867/1575 train_time:39658ms step_avg:45.74ms step:868/1575 train_time:39715ms step_avg:45.76ms step:869/1575 train_time:39779ms step_avg:45.78ms step:870/1575 train_time:39838ms step_avg:45.79ms step:871/1575 train_time:39902ms step_avg:45.81ms step:872/1575 train_time:39961ms step_avg:45.83ms step:873/1575 train_time:40025ms step_avg:45.85ms step:874/1575 train_time:40084ms step_avg:45.86ms step:875/1575 train_time:40147ms step_avg:45.88ms step:876/1575 train_time:40207ms step_avg:45.90ms step:877/1575 train_time:40270ms step_avg:45.92ms step:878/1575 train_time:40329ms step_avg:45.93ms step:879/1575 train_time:40393ms step_avg:45.95ms step:880/1575 train_time:40452ms step_avg:45.97ms step:881/1575 train_time:40515ms step_avg:45.99ms step:882/1575 train_time:40575ms step_avg:46.00ms step:883/1575 train_time:40638ms step_avg:46.02ms step:884/1575 train_time:40697ms step_avg:46.04ms step:885/1575 train_time:40760ms step_avg:46.06ms step:886/1575 train_time:40825ms step_avg:46.08ms step:887/1575 train_time:40885ms step_avg:46.09ms step:888/1575 train_time:40945ms step_avg:46.11ms step:889/1575 train_time:41009ms step_avg:46.13ms step:890/1575 train_time:41066ms step_avg:46.14ms step:891/1575 train_time:41129ms step_avg:46.16ms step:892/1575 train_time:41189ms step_avg:46.18ms step:893/1575 train_time:41252ms step_avg:46.19ms step:894/1575 train_time:41311ms step_avg:46.21ms step:895/1575 train_time:41375ms step_avg:46.23ms step:896/1575 train_time:41434ms step_avg:46.24ms step:897/1575 train_time:41497ms step_avg:46.26ms step:898/1575 train_time:41557ms step_avg:46.28ms step:899/1575 train_time:41622ms step_avg:46.30ms step:900/1575 train_time:41681ms step_avg:46.31ms step:901/1575 train_time:41743ms step_avg:46.33ms step:902/1575 train_time:41802ms step_avg:46.34ms step:903/1575 train_time:41865ms step_avg:46.36ms step:904/1575 train_time:41925ms step_avg:46.38ms step:905/1575 train_time:41988ms step_avg:46.40ms step:906/1575 train_time:42049ms step_avg:46.41ms step:907/1575 train_time:42112ms step_avg:46.43ms step:908/1575 train_time:42171ms step_avg:46.44ms step:909/1575 train_time:42234ms step_avg:46.46ms step:910/1575 train_time:42294ms step_avg:46.48ms step:911/1575 train_time:42357ms step_avg:46.50ms step:912/1575 train_time:42416ms step_avg:46.51ms step:913/1575 train_time:42479ms step_avg:46.53ms step:914/1575 train_time:42539ms step_avg:46.54ms step:915/1575 train_time:42603ms step_avg:46.56ms step:916/1575 train_time:42664ms step_avg:46.58ms step:917/1575 train_time:42726ms step_avg:46.59ms step:918/1575 train_time:42785ms step_avg:46.61ms step:919/1575 train_time:42849ms step_avg:46.63ms step:920/1575 train_time:42908ms step_avg:46.64ms step:921/1575 train_time:42972ms step_avg:46.66ms step:922/1575 train_time:43032ms step_avg:46.67ms step:923/1575 train_time:43096ms step_avg:46.69ms step:924/1575 train_time:43154ms step_avg:46.70ms step:925/1575 train_time:43217ms step_avg:46.72ms step:926/1575 train_time:43276ms step_avg:46.73ms step:927/1575 train_time:43340ms step_avg:46.75ms step:928/1575 train_time:43399ms step_avg:46.77ms step:929/1575 train_time:43462ms step_avg:46.78ms step:930/1575 train_time:43522ms step_avg:46.80ms step:931/1575 train_time:43585ms step_avg:46.82ms step:932/1575 train_time:43645ms step_avg:46.83ms step:933/1575 train_time:43708ms step_avg:46.85ms step:934/1575 train_time:43767ms step_avg:46.86ms step:935/1575 train_time:43831ms step_avg:46.88ms step:936/1575 train_time:43890ms step_avg:46.89ms step:937/1575 train_time:43954ms step_avg:46.91ms step:938/1575 train_time:44013ms step_avg:46.92ms step:939/1575 train_time:44077ms step_avg:46.94ms step:940/1575 train_time:44136ms step_avg:46.95ms step:941/1575 train_time:44199ms step_avg:46.97ms step:942/1575 train_time:44258ms step_avg:46.98ms step:943/1575 train_time:44323ms step_avg:47.00ms step:944/1575 train_time:44382ms step_avg:47.01ms step:945/1575 train_time:44445ms step_avg:47.03ms step:946/1575 train_time:44505ms step_avg:47.05ms step:947/1575 train_time:44567ms step_avg:47.06ms step:948/1575 train_time:44627ms step_avg:47.07ms step:949/1575 train_time:44690ms step_avg:47.09ms step:950/1575 train_time:44749ms step_avg:47.10ms step:951/1575 train_time:44812ms step_avg:47.12ms step:952/1575 train_time:44872ms step_avg:47.13ms step:953/1575 train_time:44936ms step_avg:47.15ms step:954/1575 train_time:44995ms step_avg:47.16ms step:955/1575 train_time:45058ms step_avg:47.18ms step:956/1575 train_time:45117ms step_avg:47.19ms step:957/1575 train_time:45181ms step_avg:47.21ms step:958/1575 train_time:45240ms step_avg:47.22ms step:959/1575 train_time:45303ms step_avg:47.24ms step:960/1575 train_time:45362ms step_avg:47.25ms step:961/1575 train_time:45426ms step_avg:47.27ms step:962/1575 train_time:45486ms step_avg:47.28ms step:963/1575 train_time:45549ms step_avg:47.30ms step:964/1575 train_time:45608ms step_avg:47.31ms step:965/1575 train_time:45672ms step_avg:47.33ms step:966/1575 train_time:45731ms step_avg:47.34ms step:967/1575 train_time:45794ms step_avg:47.36ms step:968/1575 train_time:45857ms step_avg:47.37ms step:969/1575 train_time:45918ms step_avg:47.39ms step:970/1575 train_time:45977ms step_avg:47.40ms step:971/1575 train_time:46041ms step_avg:47.42ms step:972/1575 train_time:46099ms step_avg:47.43ms step:973/1575 train_time:46164ms step_avg:47.44ms step:974/1575 train_time:46222ms step_avg:47.46ms step:975/1575 train_time:46286ms step_avg:47.47ms step:976/1575 train_time:46346ms step_avg:47.49ms step:977/1575 train_time:46409ms step_avg:47.50ms step:978/1575 train_time:46468ms step_avg:47.51ms step:979/1575 train_time:46531ms step_avg:47.53ms step:980/1575 train_time:46590ms step_avg:47.54ms step:981/1575 train_time:46653ms step_avg:47.56ms step:982/1575 train_time:46715ms step_avg:47.57ms step:983/1575 train_time:46777ms step_avg:47.59ms step:984/1575 train_time:46836ms step_avg:47.60ms step:985/1575 train_time:46899ms step_avg:47.61ms step:986/1575 train_time:46958ms step_avg:47.62ms step:987/1575 train_time:47021ms step_avg:47.64ms step:988/1575 train_time:47080ms step_avg:47.65ms step:989/1575 train_time:47143ms step_avg:47.67ms step:990/1575 train_time:47203ms step_avg:47.68ms step:991/1575 train_time:47266ms step_avg:47.70ms step:992/1575 train_time:47326ms step_avg:47.71ms step:993/1575 train_time:47389ms step_avg:47.72ms step:994/1575 train_time:47448ms step_avg:47.73ms step:995/1575 train_time:47512ms step_avg:47.75ms step:996/1575 train_time:47571ms step_avg:47.76ms step:997/1575 train_time:47634ms step_avg:47.78ms step:998/1575 train_time:47693ms step_avg:47.79ms step:999/1575 train_time:47757ms step_avg:47.80ms step:1000/1575 train_time:47816ms step_avg:47.82ms step:1000/1575 val_loss:3.5830 train_time:47862ms step_avg:47.86ms step:1001/1575 train_time:47884ms step_avg:47.84ms step:1002/1575 train_time:47944ms step_avg:47.85ms step:1003/1575 train_time:48012ms step_avg:47.87ms step:1004/1575 train_time:48075ms step_avg:47.88ms step:1005/1575 train_time:48138ms step_avg:47.90ms step:1006/1575 train_time:48198ms step_avg:47.91ms step:1007/1575 train_time:48262ms step_avg:47.93ms step:1008/1575 train_time:48322ms step_avg:47.94ms step:1009/1575 train_time:48385ms step_avg:47.95ms step:1010/1575 train_time:48444ms step_avg:47.96ms step:1011/1575 train_time:48507ms step_avg:47.98ms step:1012/1575 train_time:48566ms step_avg:47.99ms step:1013/1575 train_time:48630ms step_avg:48.01ms step:1014/1575 train_time:48688ms step_avg:48.02ms step:1015/1575 train_time:48750ms step_avg:48.03ms step:1016/1575 train_time:48810ms step_avg:48.04ms step:1017/1575 train_time:48873ms step_avg:48.06ms step:1018/1575 train_time:48934ms step_avg:48.07ms step:1019/1575 train_time:48999ms step_avg:48.09ms step:1020/1575 train_time:49058ms step_avg:48.10ms step:1021/1575 train_time:49123ms step_avg:48.11ms step:1022/1575 train_time:49183ms step_avg:48.12ms step:1023/1575 train_time:49245ms step_avg:48.14ms step:1024/1575 train_time:49305ms step_avg:48.15ms step:1025/1575 train_time:49380ms step_avg:48.18ms step:1026/1575 train_time:49459ms step_avg:48.21ms step:1027/1575 train_time:49549ms step_avg:48.25ms step:1028/1575 train_time:49636ms step_avg:48.28ms step:1029/1575 train_time:49724ms step_avg:48.32ms step:1030/1575 train_time:49811ms step_avg:48.36ms step:1031/1575 train_time:49900ms step_avg:48.40ms step:1032/1575 train_time:49986ms step_avg:48.44ms step:1033/1575 train_time:50078ms step_avg:48.48ms step:1034/1575 train_time:50164ms step_avg:48.51ms step:1035/1575 train_time:50254ms step_avg:48.56ms step:1036/1575 train_time:50341ms step_avg:48.59ms step:1037/1575 train_time:50432ms step_avg:48.63ms step:1038/1575 train_time:50516ms step_avg:48.67ms step:1039/1575 train_time:50605ms step_avg:48.71ms step:1040/1575 train_time:50690ms step_avg:48.74ms step:1041/1575 train_time:50780ms step_avg:48.78ms step:1042/1575 train_time:50868ms step_avg:48.82ms step:1043/1575 train_time:50961ms step_avg:48.86ms step:1044/1575 train_time:51046ms step_avg:48.89ms step:1045/1575 train_time:51134ms step_avg:48.93ms step:1046/1575 train_time:51220ms step_avg:48.97ms step:1047/1575 train_time:51310ms step_avg:49.01ms step:1048/1575 train_time:51395ms step_avg:49.04ms step:1049/1575 train_time:51485ms step_avg:49.08ms step:1050/1575 train_time:51570ms step_avg:49.11ms step:1051/1575 train_time:51661ms step_avg:49.15ms step:1052/1575 train_time:51746ms step_avg:49.19ms step:1053/1575 train_time:51835ms step_avg:49.23ms step:1054/1575 train_time:51922ms step_avg:49.26ms step:1055/1575 train_time:52011ms step_avg:49.30ms step:1056/1575 train_time:52098ms step_avg:49.33ms step:1057/1575 train_time:52187ms step_avg:49.37ms step:1058/1575 train_time:52273ms step_avg:49.41ms step:1059/1575 train_time:52362ms step_avg:49.44ms step:1060/1575 train_time:52450ms step_avg:49.48ms step:1061/1575 train_time:52538ms step_avg:49.52ms step:1062/1575 train_time:52624ms step_avg:49.55ms step:1063/1575 train_time:52714ms step_avg:49.59ms step:1064/1575 train_time:52800ms step_avg:49.62ms step:1065/1575 train_time:52889ms step_avg:49.66ms step:1066/1575 train_time:52976ms step_avg:49.70ms step:1067/1575 train_time:53065ms step_avg:49.73ms step:1068/1575 train_time:53152ms step_avg:49.77ms step:1069/1575 train_time:53241ms step_avg:49.80ms step:1070/1575 train_time:53327ms step_avg:49.84ms step:1071/1575 train_time:53419ms step_avg:49.88ms step:1072/1575 train_time:53503ms step_avg:49.91ms step:1073/1575 train_time:53592ms step_avg:49.95ms step:1074/1575 train_time:53679ms step_avg:49.98ms step:1075/1575 train_time:53768ms step_avg:50.02ms step:1076/1575 train_time:53854ms step_avg:50.05ms step:1077/1575 train_time:53943ms step_avg:50.09ms step:1078/1575 train_time:54029ms step_avg:50.12ms step:1079/1575 train_time:54119ms step_avg:50.16ms step:1080/1575 train_time:54204ms step_avg:50.19ms step:1081/1575 train_time:54294ms step_avg:50.23ms step:1082/1575 train_time:54380ms step_avg:50.26ms step:1083/1575 train_time:54469ms step_avg:50.29ms step:1084/1575 train_time:54555ms step_avg:50.33ms step:1085/1575 train_time:54645ms step_avg:50.36ms step:1086/1575 train_time:54731ms step_avg:50.40ms step:1087/1575 train_time:54820ms step_avg:50.43ms step:1088/1575 train_time:54906ms step_avg:50.46ms step:1089/1575 train_time:54996ms step_avg:50.50ms step:1090/1575 train_time:55082ms step_avg:50.53ms step:1091/1575 train_time:55172ms step_avg:50.57ms step:1092/1575 train_time:55257ms step_avg:50.60ms step:1093/1575 train_time:55349ms step_avg:50.64ms step:1094/1575 train_time:55433ms step_avg:50.67ms step:1095/1575 train_time:55523ms step_avg:50.71ms step:1096/1575 train_time:55608ms step_avg:50.74ms step:1097/1575 train_time:55699ms step_avg:50.77ms step:1098/1575 train_time:55784ms step_avg:50.80ms step:1099/1575 train_time:55873ms step_avg:50.84ms step:1100/1575 train_time:55959ms step_avg:50.87ms step:1101/1575 train_time:56051ms step_avg:50.91ms step:1102/1575 train_time:56137ms step_avg:50.94ms step:1103/1575 train_time:56224ms step_avg:50.97ms step:1104/1575 train_time:56310ms step_avg:51.01ms step:1105/1575 train_time:56399ms step_avg:51.04ms step:1106/1575 train_time:56485ms step_avg:51.07ms step:1107/1575 train_time:56574ms step_avg:51.11ms step:1108/1575 train_time:56660ms step_avg:51.14ms step:1109/1575 train_time:56749ms step_avg:51.17ms step:1110/1575 train_time:56835ms step_avg:51.20ms step:1111/1575 train_time:56924ms step_avg:51.24ms step:1112/1575 train_time:57011ms step_avg:51.27ms step:1113/1575 train_time:57100ms step_avg:51.30ms step:1114/1575 train_time:57186ms step_avg:51.33ms step:1115/1575 train_time:57276ms step_avg:51.37ms step:1116/1575 train_time:57362ms step_avg:51.40ms step:1117/1575 train_time:57452ms step_avg:51.43ms step:1118/1575 train_time:57536ms step_avg:51.46ms step:1119/1575 train_time:57627ms step_avg:51.50ms step:1120/1575 train_time:57712ms step_avg:51.53ms step:1121/1575 train_time:57802ms step_avg:51.56ms step:1122/1575 train_time:57887ms step_avg:51.59ms step:1123/1575 train_time:57977ms step_avg:51.63ms step:1124/1575 train_time:58062ms step_avg:51.66ms step:1125/1575 train_time:58152ms step_avg:51.69ms step:1126/1575 train_time:58238ms step_avg:51.72ms step:1127/1575 train_time:58328ms step_avg:51.76ms step:1128/1575 train_time:58421ms step_avg:51.79ms step:1129/1575 train_time:58505ms step_avg:51.82ms step:1130/1575 train_time:58590ms step_avg:51.85ms step:1131/1575 train_time:58679ms step_avg:51.88ms step:1132/1575 train_time:58767ms step_avg:51.91ms step:1133/1575 train_time:58855ms step_avg:51.95ms step:1134/1575 train_time:58941ms step_avg:51.98ms step:1135/1575 train_time:59030ms step_avg:52.01ms step:1136/1575 train_time:59116ms step_avg:52.04ms step:1137/1575 train_time:59205ms step_avg:52.07ms step:1138/1575 train_time:59291ms step_avg:52.10ms step:1139/1575 train_time:59381ms step_avg:52.13ms step:1140/1575 train_time:59467ms step_avg:52.16ms step:1141/1575 train_time:59560ms step_avg:52.20ms step:1142/1575 train_time:59642ms step_avg:52.23ms step:1143/1575 train_time:59732ms step_avg:52.26ms step:1144/1575 train_time:59817ms step_avg:52.29ms step:1145/1575 train_time:59905ms step_avg:52.32ms step:1146/1575 train_time:59996ms step_avg:52.35ms step:1147/1575 train_time:60083ms step_avg:52.38ms step:1148/1575 train_time:60169ms step_avg:52.41ms step:1149/1575 train_time:60260ms step_avg:52.45ms step:1150/1575 train_time:60345ms step_avg:52.47ms step:1151/1575 train_time:60434ms step_avg:52.51ms step:1152/1575 train_time:60520ms step_avg:52.53ms step:1153/1575 train_time:60609ms step_avg:52.57ms step:1154/1575 train_time:60695ms step_avg:52.60ms step:1155/1575 train_time:60784ms step_avg:52.63ms step:1156/1575 train_time:60871ms step_avg:52.66ms step:1157/1575 train_time:60960ms step_avg:52.69ms step:1158/1575 train_time:61046ms step_avg:52.72ms step:1159/1575 train_time:61136ms step_avg:52.75ms step:1160/1575 train_time:61223ms step_avg:52.78ms step:1161/1575 train_time:61313ms step_avg:52.81ms step:1162/1575 train_time:61397ms step_avg:52.84ms step:1163/1575 train_time:61486ms step_avg:52.87ms step:1164/1575 train_time:61572ms step_avg:52.90ms step:1165/1575 train_time:61662ms step_avg:52.93ms step:1166/1575 train_time:61747ms step_avg:52.96ms step:1167/1575 train_time:61837ms step_avg:52.99ms step:1168/1575 train_time:61923ms step_avg:53.02ms step:1169/1575 train_time:62013ms step_avg:53.05ms step:1170/1575 train_time:62098ms step_avg:53.08ms step:1171/1575 train_time:62188ms step_avg:53.11ms step:1172/1575 train_time:62274ms step_avg:53.13ms step:1173/1575 train_time:62364ms step_avg:53.17ms step:1174/1575 train_time:62449ms step_avg:53.19ms step:1175/1575 train_time:62540ms step_avg:53.23ms step:1176/1575 train_time:62625ms step_avg:53.25ms step:1177/1575 train_time:62714ms step_avg:53.28ms step:1178/1575 train_time:62801ms step_avg:53.31ms step:1179/1575 train_time:62890ms step_avg:53.34ms step:1180/1575 train_time:62976ms step_avg:53.37ms step:1181/1575 train_time:63065ms step_avg:53.40ms step:1182/1575 train_time:63152ms step_avg:53.43ms step:1183/1575 train_time:63242ms step_avg:53.46ms step:1184/1575 train_time:63329ms step_avg:53.49ms step:1185/1575 train_time:63418ms step_avg:53.52ms step:1186/1575 train_time:63504ms step_avg:53.54ms step:1187/1575 train_time:63593ms step_avg:53.57ms step:1188/1575 train_time:63679ms step_avg:53.60ms step:1189/1575 train_time:63768ms step_avg:53.63ms step:1190/1575 train_time:63854ms step_avg:53.66ms step:1191/1575 train_time:63943ms step_avg:53.69ms step:1192/1575 train_time:64029ms step_avg:53.72ms step:1193/1575 train_time:64119ms step_avg:53.75ms step:1194/1575 train_time:64205ms step_avg:53.77ms step:1195/1575 train_time:64295ms step_avg:53.80ms step:1196/1575 train_time:64380ms step_avg:53.83ms step:1197/1575 train_time:64470ms step_avg:53.86ms step:1198/1575 train_time:64556ms step_avg:53.89ms step:1199/1575 train_time:64645ms step_avg:53.92ms step:1200/1575 train_time:64732ms step_avg:53.94ms step:1201/1575 train_time:64822ms step_avg:53.97ms step:1202/1575 train_time:64906ms step_avg:54.00ms step:1203/1575 train_time:64996ms step_avg:54.03ms step:1204/1575 train_time:65082ms step_avg:54.05ms step:1205/1575 train_time:65172ms step_avg:54.08ms step:1206/1575 train_time:65257ms step_avg:54.11ms step:1207/1575 train_time:65346ms step_avg:54.14ms step:1208/1575 train_time:65433ms step_avg:54.17ms step:1209/1575 train_time:65522ms step_avg:54.20ms step:1210/1575 train_time:65608ms step_avg:54.22ms step:1211/1575 train_time:65698ms step_avg:54.25ms step:1212/1575 train_time:65784ms step_avg:54.28ms step:1213/1575 train_time:65873ms step_avg:54.31ms step:1214/1575 train_time:65959ms step_avg:54.33ms step:1215/1575 train_time:66048ms step_avg:54.36ms step:1216/1575 train_time:66135ms step_avg:54.39ms step:1217/1575 train_time:66224ms step_avg:54.42ms step:1218/1575 train_time:66311ms step_avg:54.44ms step:1219/1575 train_time:66401ms step_avg:54.47ms step:1220/1575 train_time:66486ms step_avg:54.50ms step:1221/1575 train_time:66575ms step_avg:54.53ms step:1222/1575 train_time:66661ms step_avg:54.55ms step:1223/1575 train_time:66751ms step_avg:54.58ms step:1224/1575 train_time:66838ms step_avg:54.61ms step:1225/1575 train_time:66926ms step_avg:54.63ms step:1226/1575 train_time:67011ms step_avg:54.66ms step:1227/1575 train_time:67101ms step_avg:54.69ms step:1228/1575 train_time:67187ms step_avg:54.71ms step:1229/1575 train_time:67277ms step_avg:54.74ms step:1230/1575 train_time:67362ms step_avg:54.77ms step:1231/1575 train_time:67452ms step_avg:54.79ms step:1232/1575 train_time:67537ms step_avg:54.82ms step:1233/1575 train_time:67626ms step_avg:54.85ms step:1234/1575 train_time:67713ms step_avg:54.87ms step:1235/1575 train_time:67802ms step_avg:54.90ms step:1236/1575 train_time:67888ms step_avg:54.93ms step:1237/1575 train_time:67978ms step_avg:54.95ms step:1238/1575 train_time:68062ms step_avg:54.98ms step:1239/1575 train_time:68152ms step_avg:55.01ms step:1240/1575 train_time:68237ms step_avg:55.03ms step:1241/1575 train_time:68328ms step_avg:55.06ms step:1242/1575 train_time:68415ms step_avg:55.08ms step:1243/1575 train_time:68503ms step_avg:55.11ms step:1244/1575 train_time:68589ms step_avg:55.14ms step:1245/1575 train_time:68679ms step_avg:55.16ms step:1246/1575 train_time:68764ms step_avg:55.19ms step:1247/1575 train_time:68854ms step_avg:55.22ms step:1248/1575 train_time:68940ms step_avg:55.24ms step:1249/1575 train_time:69029ms step_avg:55.27ms step:1250/1575 train_time:69115ms step_avg:55.29ms step:1250/1575 val_loss:3.4050 train_time:69187ms step_avg:55.35ms step:1251/1575 train_time:69209ms step_avg:55.32ms step:1252/1575 train_time:69296ms step_avg:55.35ms step:1253/1575 train_time:69388ms step_avg:55.38ms step:1254/1575 train_time:69476ms step_avg:55.40ms step:1255/1575 train_time:69563ms step_avg:55.43ms step:1256/1575 train_time:69649ms step_avg:55.45ms step:1257/1575 train_time:69736ms step_avg:55.48ms step:1258/1575 train_time:69821ms step_avg:55.50ms step:1259/1575 train_time:69910ms step_avg:55.53ms step:1260/1575 train_time:69994ms step_avg:55.55ms step:1261/1575 train_time:70082ms step_avg:55.58ms step:1262/1575 train_time:70168ms step_avg:55.60ms step:1263/1575 train_time:70261ms step_avg:55.63ms step:1264/1575 train_time:70349ms step_avg:55.66ms step:1265/1575 train_time:70440ms step_avg:55.68ms step:1266/1575 train_time:70527ms step_avg:55.71ms step:1267/1575 train_time:70617ms step_avg:55.74ms step:1268/1575 train_time:70701ms step_avg:55.76ms step:1269/1575 train_time:70790ms step_avg:55.78ms step:1270/1575 train_time:70875ms step_avg:55.81ms step:1271/1575 train_time:70964ms step_avg:55.83ms step:1272/1575 train_time:71048ms step_avg:55.86ms step:1273/1575 train_time:71137ms step_avg:55.88ms step:1274/1575 train_time:71224ms step_avg:55.91ms step:1275/1575 train_time:71315ms step_avg:55.93ms step:1276/1575 train_time:71401ms step_avg:55.96ms step:1277/1575 train_time:71491ms step_avg:55.98ms step:1278/1575 train_time:71578ms step_avg:56.01ms step:1279/1575 train_time:71667ms step_avg:56.03ms step:1280/1575 train_time:71751ms step_avg:56.06ms step:1281/1575 train_time:71841ms step_avg:56.08ms step:1282/1575 train_time:71926ms step_avg:56.10ms step:1283/1575 train_time:72014ms step_avg:56.13ms step:1284/1575 train_time:72099ms step_avg:56.15ms step:1285/1575 train_time:72188ms step_avg:56.18ms step:1286/1575 train_time:72275ms step_avg:56.20ms step:1287/1575 train_time:72365ms step_avg:56.23ms step:1288/1575 train_time:72452ms step_avg:56.25ms step:1289/1575 train_time:72541ms step_avg:56.28ms step:1290/1575 train_time:72627ms step_avg:56.30ms step:1291/1575 train_time:72716ms step_avg:56.33ms step:1292/1575 train_time:72801ms step_avg:56.35ms step:1293/1575 train_time:72889ms step_avg:56.37ms step:1294/1575 train_time:72976ms step_avg:56.40ms step:1295/1575 train_time:73065ms step_avg:56.42ms step:1296/1575 train_time:73151ms step_avg:56.44ms step:1297/1575 train_time:73241ms step_avg:56.47ms step:1298/1575 train_time:73327ms step_avg:56.49ms step:1299/1575 train_time:73418ms step_avg:56.52ms step:1300/1575 train_time:73504ms step_avg:56.54ms step:1301/1575 train_time:73594ms step_avg:56.57ms step:1302/1575 train_time:73679ms step_avg:56.59ms step:1303/1575 train_time:73769ms step_avg:56.61ms step:1304/1575 train_time:73854ms step_avg:56.64ms step:1305/1575 train_time:73942ms step_avg:56.66ms step:1306/1575 train_time:74028ms step_avg:56.68ms step:1307/1575 train_time:74119ms step_avg:56.71ms step:1308/1575 train_time:74204ms step_avg:56.73ms step:1309/1575 train_time:74294ms step_avg:56.76ms step:1310/1575 train_time:74380ms step_avg:56.78ms step:1311/1575 train_time:74470ms step_avg:56.80ms step:1312/1575 train_time:74556ms step_avg:56.83ms step:1313/1575 train_time:74645ms step_avg:56.85ms step:1314/1575 train_time:74732ms step_avg:56.87ms step:1315/1575 train_time:74822ms step_avg:56.90ms step:1316/1575 train_time:74907ms step_avg:56.92ms step:1317/1575 train_time:74996ms step_avg:56.94ms step:1318/1575 train_time:75082ms step_avg:56.97ms step:1319/1575 train_time:75171ms step_avg:56.99ms step:1320/1575 train_time:75256ms step_avg:57.01ms step:1321/1575 train_time:75345ms step_avg:57.04ms step:1322/1575 train_time:75431ms step_avg:57.06ms step:1323/1575 train_time:75522ms step_avg:57.08ms step:1324/1575 train_time:75607ms step_avg:57.11ms step:1325/1575 train_time:75698ms step_avg:57.13ms step:1326/1575 train_time:75785ms step_avg:57.15ms step:1327/1575 train_time:75873ms step_avg:57.18ms step:1328/1575 train_time:75958ms step_avg:57.20ms step:1329/1575 train_time:76046ms step_avg:57.22ms step:1330/1575 train_time:76132ms step_avg:57.24ms step:1331/1575 train_time:76222ms step_avg:57.27ms step:1332/1575 train_time:76310ms step_avg:57.29ms step:1333/1575 train_time:76397ms step_avg:57.31ms step:1334/1575 train_time:76483ms step_avg:57.33ms step:1335/1575 train_time:76573ms step_avg:57.36ms step:1336/1575 train_time:76658ms step_avg:57.38ms step:1337/1575 train_time:76748ms step_avg:57.40ms step:1338/1575 train_time:76834ms step_avg:57.42ms step:1339/1575 train_time:76923ms step_avg:57.45ms step:1340/1575 train_time:77009ms step_avg:57.47ms step:1341/1575 train_time:77098ms step_avg:57.49ms step:1342/1575 train_time:77183ms step_avg:57.51ms step:1343/1575 train_time:77273ms step_avg:57.54ms step:1344/1575 train_time:77358ms step_avg:57.56ms step:1345/1575 train_time:77448ms step_avg:57.58ms step:1346/1575 train_time:77534ms step_avg:57.60ms step:1347/1575 train_time:77624ms step_avg:57.63ms step:1348/1575 train_time:77710ms step_avg:57.65ms step:1349/1575 train_time:77800ms step_avg:57.67ms step:1350/1575 train_time:77886ms step_avg:57.69ms step:1351/1575 train_time:77976ms step_avg:57.72ms step:1352/1575 train_time:78060ms step_avg:57.74ms step:1353/1575 train_time:78149ms step_avg:57.76ms step:1354/1575 train_time:78236ms step_avg:57.78ms step:1355/1575 train_time:78325ms step_avg:57.80ms step:1356/1575 train_time:78411ms step_avg:57.83ms step:1357/1575 train_time:78500ms step_avg:57.85ms step:1358/1575 train_time:78591ms step_avg:57.87ms step:1359/1575 train_time:78679ms step_avg:57.89ms step:1360/1575 train_time:78765ms step_avg:57.92ms step:1361/1575 train_time:78853ms step_avg:57.94ms step:1362/1575 train_time:78938ms step_avg:57.96ms step:1363/1575 train_time:79028ms step_avg:57.98ms step:1364/1575 train_time:79113ms step_avg:58.00ms step:1365/1575 train_time:79203ms step_avg:58.02ms step:1366/1575 train_time:79289ms step_avg:58.04ms step:1367/1575 train_time:79379ms step_avg:58.07ms step:1368/1575 train_time:79466ms step_avg:58.09ms step:1369/1575 train_time:79555ms step_avg:58.11ms step:1370/1575 train_time:79640ms step_avg:58.13ms step:1371/1575 train_time:79730ms step_avg:58.15ms step:1372/1575 train_time:79817ms step_avg:58.18ms step:1373/1575 train_time:79906ms step_avg:58.20ms step:1374/1575 train_time:79991ms step_avg:58.22ms step:1375/1575 train_time:80082ms step_avg:58.24ms step:1376/1575 train_time:80168ms step_avg:58.26ms step:1377/1575 train_time:80257ms step_avg:58.28ms step:1378/1575 train_time:80342ms step_avg:58.30ms step:1379/1575 train_time:80431ms step_avg:58.33ms step:1380/1575 train_time:80517ms step_avg:58.35ms step:1381/1575 train_time:80606ms step_avg:58.37ms step:1382/1575 train_time:80692ms step_avg:58.39ms step:1383/1575 train_time:80782ms step_avg:58.41ms step:1384/1575 train_time:80870ms step_avg:58.43ms step:1385/1575 train_time:80958ms step_avg:58.45ms step:1386/1575 train_time:81044ms step_avg:58.47ms step:1387/1575 train_time:81136ms step_avg:58.50ms step:1388/1575 train_time:81220ms step_avg:58.52ms step:1389/1575 train_time:81308ms step_avg:58.54ms step:1390/1575 train_time:81395ms step_avg:58.56ms step:1391/1575 train_time:81484ms step_avg:58.58ms step:1392/1575 train_time:81570ms step_avg:58.60ms step:1393/1575 train_time:81660ms step_avg:58.62ms step:1394/1575 train_time:81745ms step_avg:58.64ms step:1395/1575 train_time:81835ms step_avg:58.66ms step:1396/1575 train_time:81920ms step_avg:58.68ms step:1397/1575 train_time:82009ms step_avg:58.70ms step:1398/1575 train_time:82095ms step_avg:58.72ms step:1399/1575 train_time:82184ms step_avg:58.74ms step:1400/1575 train_time:82269ms step_avg:58.76ms step:1401/1575 train_time:82360ms step_avg:58.79ms step:1402/1575 train_time:82446ms step_avg:58.81ms step:1403/1575 train_time:82536ms step_avg:58.83ms step:1404/1575 train_time:82621ms step_avg:58.85ms step:1405/1575 train_time:82711ms step_avg:58.87ms step:1406/1575 train_time:82796ms step_avg:58.89ms step:1407/1575 train_time:82886ms step_avg:58.91ms step:1408/1575 train_time:82973ms step_avg:58.93ms step:1409/1575 train_time:83061ms step_avg:58.95ms step:1410/1575 train_time:83147ms step_avg:58.97ms step:1411/1575 train_time:83236ms step_avg:58.99ms step:1412/1575 train_time:83322ms step_avg:59.01ms step:1413/1575 train_time:83412ms step_avg:59.03ms step:1414/1575 train_time:83497ms step_avg:59.05ms step:1415/1575 train_time:83587ms step_avg:59.07ms step:1416/1575 train_time:83674ms step_avg:59.09ms step:1417/1575 train_time:83763ms step_avg:59.11ms step:1418/1575 train_time:83849ms step_avg:59.13ms step:1419/1575 train_time:83939ms step_avg:59.15ms step:1420/1575 train_time:84025ms step_avg:59.17ms step:1421/1575 train_time:84114ms step_avg:59.19ms step:1422/1575 train_time:84199ms step_avg:59.21ms step:1423/1575 train_time:84291ms step_avg:59.23ms step:1424/1575 train_time:84377ms step_avg:59.25ms step:1425/1575 train_time:84466ms step_avg:59.27ms step:1426/1575 train_time:84551ms step_avg:59.29ms step:1427/1575 train_time:84641ms step_avg:59.31ms step:1428/1575 train_time:84726ms step_avg:59.33ms step:1429/1575 train_time:84816ms step_avg:59.35ms step:1430/1575 train_time:84901ms step_avg:59.37ms step:1431/1575 train_time:84992ms step_avg:59.39ms step:1432/1575 train_time:85076ms step_avg:59.41ms step:1433/1575 train_time:85166ms step_avg:59.43ms step:1434/1575 train_time:85252ms step_avg:59.45ms step:1435/1575 train_time:85342ms step_avg:59.47ms step:1436/1575 train_time:85428ms step_avg:59.49ms step:1437/1575 train_time:85518ms step_avg:59.51ms step:1438/1575 train_time:85603ms step_avg:59.53ms step:1439/1575 train_time:85693ms step_avg:59.55ms step:1440/1575 train_time:85778ms step_avg:59.57ms step:1441/1575 train_time:85867ms step_avg:59.59ms step:1442/1575 train_time:85955ms step_avg:59.61ms step:1443/1575 train_time:86043ms step_avg:59.63ms step:1444/1575 train_time:86128ms step_avg:59.65ms step:1445/1575 train_time:86217ms step_avg:59.67ms step:1446/1575 train_time:86303ms step_avg:59.68ms step:1447/1575 train_time:86393ms step_avg:59.70ms step:1448/1575 train_time:86479ms step_avg:59.72ms step:1449/1575 train_time:86568ms step_avg:59.74ms step:1450/1575 train_time:86655ms step_avg:59.76ms step:1451/1575 train_time:86744ms step_avg:59.78ms step:1452/1575 train_time:86830ms step_avg:59.80ms step:1453/1575 train_time:86920ms step_avg:59.82ms step:1454/1575 train_time:87005ms step_avg:59.84ms step:1455/1575 train_time:87095ms step_avg:59.86ms step:1456/1575 train_time:87180ms step_avg:59.88ms step:1457/1575 train_time:87270ms step_avg:59.90ms step:1458/1575 train_time:87355ms step_avg:59.91ms step:1459/1575 train_time:87444ms step_avg:59.93ms step:1460/1575 train_time:87530ms step_avg:59.95ms step:1461/1575 train_time:87620ms step_avg:59.97ms step:1462/1575 train_time:87706ms step_avg:59.99ms step:1463/1575 train_time:87794ms step_avg:60.01ms step:1464/1575 train_time:87880ms step_avg:60.03ms step:1465/1575 train_time:87969ms step_avg:60.05ms step:1466/1575 train_time:88054ms step_avg:60.06ms step:1467/1575 train_time:88144ms step_avg:60.08ms step:1468/1575 train_time:88229ms step_avg:60.10ms step:1469/1575 train_time:88319ms step_avg:60.12ms step:1470/1575 train_time:88405ms step_avg:60.14ms step:1471/1575 train_time:88494ms step_avg:60.16ms step:1472/1575 train_time:88580ms step_avg:60.18ms step:1473/1575 train_time:88670ms step_avg:60.20ms step:1474/1575 train_time:88756ms step_avg:60.21ms step:1475/1575 train_time:88845ms step_avg:60.23ms step:1476/1575 train_time:88930ms step_avg:60.25ms step:1477/1575 train_time:89020ms step_avg:60.27ms step:1478/1575 train_time:89106ms step_avg:60.29ms step:1479/1575 train_time:89196ms step_avg:60.31ms step:1480/1575 train_time:89281ms step_avg:60.33ms step:1481/1575 train_time:89371ms step_avg:60.34ms step:1482/1575 train_time:89456ms step_avg:60.36ms step:1483/1575 train_time:89546ms step_avg:60.38ms step:1484/1575 train_time:89631ms step_avg:60.40ms step:1485/1575 train_time:89723ms step_avg:60.42ms step:1486/1575 train_time:89808ms step_avg:60.44ms step:1487/1575 train_time:89898ms step_avg:60.46ms step:1488/1575 train_time:89985ms step_avg:60.47ms step:1489/1575 train_time:90073ms step_avg:60.49ms step:1490/1575 train_time:90158ms step_avg:60.51ms step:1491/1575 train_time:90247ms step_avg:60.53ms step:1492/1575 train_time:90333ms step_avg:60.55ms step:1493/1575 train_time:90423ms step_avg:60.56ms step:1494/1575 train_time:90509ms step_avg:60.58ms step:1495/1575 train_time:90598ms step_avg:60.60ms step:1496/1575 train_time:90684ms step_avg:60.62ms step:1497/1575 train_time:90773ms step_avg:60.64ms step:1498/1575 train_time:90859ms step_avg:60.65ms step:1499/1575 train_time:90948ms step_avg:60.67ms step:1500/1575 train_time:91037ms step_avg:60.69ms step:1500/1575 val_loss:3.2993 train_time:91107ms step_avg:60.74ms step:1501/1575 train_time:91128ms step_avg:60.71ms step:1502/1575 train_time:91216ms step_avg:60.73ms step:1503/1575 train_time:91310ms step_avg:60.75ms step:1504/1575 train_time:91398ms step_avg:60.77ms step:1505/1575 train_time:91486ms step_avg:60.79ms step:1506/1575 train_time:91572ms step_avg:60.80ms step:1507/1575 train_time:91660ms step_avg:60.82ms step:1508/1575 train_time:91744ms step_avg:60.84ms step:1509/1575 train_time:91832ms step_avg:60.86ms step:1510/1575 train_time:91917ms step_avg:60.87ms step:1511/1575 train_time:92005ms step_avg:60.89ms step:1512/1575 train_time:92090ms step_avg:60.91ms step:1513/1575 train_time:92182ms step_avg:60.93ms step:1514/1575 train_time:92270ms step_avg:60.94ms step:1515/1575 train_time:92361ms step_avg:60.96ms step:1516/1575 train_time:92448ms step_avg:60.98ms step:1517/1575 train_time:92537ms step_avg:61.00ms step:1518/1575 train_time:92622ms step_avg:61.02ms step:1519/1575 train_time:92710ms step_avg:61.03ms step:1520/1575 train_time:92796ms step_avg:61.05ms step:1521/1575 train_time:92884ms step_avg:61.07ms step:1522/1575 train_time:92968ms step_avg:61.08ms step:1523/1575 train_time:93059ms step_avg:61.10ms step:1524/1575 train_time:93147ms step_avg:61.12ms step:1525/1575 train_time:93235ms step_avg:61.14ms step:1526/1575 train_time:93321ms step_avg:61.15ms step:1527/1575 train_time:93412ms step_avg:61.17ms step:1528/1575 train_time:93500ms step_avg:61.19ms step:1529/1575 train_time:93589ms step_avg:61.21ms step:1530/1575 train_time:93674ms step_avg:61.22ms step:1531/1575 train_time:93762ms step_avg:61.24ms step:1532/1575 train_time:93848ms step_avg:61.26ms step:1533/1575 train_time:93937ms step_avg:61.28ms step:1534/1575 train_time:94022ms step_avg:61.29ms step:1535/1575 train_time:94113ms step_avg:61.31ms step:1536/1575 train_time:94208ms step_avg:61.33ms step:1537/1575 train_time:94294ms step_avg:61.35ms step:1538/1575 train_time:94380ms step_avg:61.37ms step:1539/1575 train_time:94470ms step_avg:61.38ms step:1540/1575 train_time:94557ms step_avg:61.40ms step:1541/1575 train_time:94646ms step_avg:61.42ms step:1542/1575 train_time:94732ms step_avg:61.43ms step:1543/1575 train_time:94821ms step_avg:61.45ms step:1544/1575 train_time:94907ms step_avg:61.47ms step:1545/1575 train_time:94998ms step_avg:61.49ms step:1546/1575 train_time:95083ms step_avg:61.50ms step:1547/1575 train_time:95174ms step_avg:61.52ms step:1548/1575 train_time:95260ms step_avg:61.54ms step:1549/1575 train_time:95351ms step_avg:61.56ms step:1550/1575 train_time:95437ms step_avg:61.57ms step:1551/1575 train_time:95527ms step_avg:61.59ms step:1552/1575 train_time:95614ms step_avg:61.61ms step:1553/1575 train_time:95702ms step_avg:61.62ms step:1554/1575 train_time:95787ms step_avg:61.64ms step:1555/1575 train_time:95877ms step_avg:61.66ms step:1556/1575 train_time:95962ms step_avg:61.67ms step:1557/1575 train_time:96052ms step_avg:61.69ms step:1558/1575 train_time:96138ms step_avg:61.71ms step:1559/1575 train_time:96228ms step_avg:61.72ms step:1560/1575 train_time:96314ms step_avg:61.74ms step:1561/1575 train_time:96404ms step_avg:61.76ms step:1562/1575 train_time:96491ms step_avg:61.77ms step:1563/1575 train_time:96581ms step_avg:61.79ms step:1564/1575 train_time:96666ms step_avg:61.81ms step:1565/1575 train_time:96757ms step_avg:61.83ms step:1566/1575 train_time:96843ms step_avg:61.84ms step:1567/1575 train_time:96932ms step_avg:61.86ms step:1568/1575 train_time:97018ms step_avg:61.87ms step:1569/1575 train_time:97114ms step_avg:61.90ms step:1570/1575 train_time:97198ms step_avg:61.91ms step:1571/1575 train_time:97286ms step_avg:61.93ms step:1572/1575 train_time:97372ms step_avg:61.94ms step:1573/1575 train_time:97464ms step_avg:61.96ms step:1574/1575 train_time:97548ms step_avg:61.97ms step:1575/1575 train_time:97638ms step_avg:61.99ms step:1575/1575 val_loss:3.2778 train_time:97704ms step_avg:62.03ms peak memory allocated: 30933 MiB reserved: 47000 MiB