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:33:33 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 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 41C P0 122W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 35C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:6A:00.0 Off | 0 | | N/A 35C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 42C P0 131W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 40C P0 124W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 36C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 254119 C /usr/bin/python3 1510MiB | | 1 N/A N/A 254120 C /usr/bin/python3 1510MiB | | 2 N/A N/A 254121 C /usr/bin/python3 1510MiB | | 3 N/A N/A 254122 C /usr/bin/python3 1510MiB | | 4 N/A N/A 254123 C /usr/bin/python3 1510MiB | | 5 N/A N/A 254124 C /usr/bin/python3 1510MiB | | 6 N/A N/A 254125 C /usr/bin/python3 1510MiB | | 7 N/A N/A 254126 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.8295 train_time:0ms step_avg:0.03ms step:1/1575 train_time:91ms step_avg:90.81ms step:2/1575 train_time:115ms step_avg:57.59ms step:3/1575 train_time:136ms step_avg:45.35ms step:4/1575 train_time:167ms step_avg:41.70ms step:5/1575 train_time:197ms step_avg:39.40ms step:6/1575 train_time:270ms step_avg:45.08ms step:7/1575 train_time:289ms step_avg:41.27ms step:8/1575 train_time:440ms step_avg:55.05ms step:9/1575 train_time:471ms step_avg:52.30ms step:10/1575 train_time:510ms step_avg:50.98ms step:11/1575 train_time:540ms step_avg:49.10ms step:12/1575 train_time:579ms step_avg:48.23ms step:13/1575 train_time:610ms step_avg:46.91ms step:14/1575 train_time:649ms step_avg:46.33ms step:15/1575 train_time:680ms step_avg:45.31ms step:16/1575 train_time:719ms step_avg:44.93ms step:17/1575 train_time:749ms step_avg:44.09ms step:18/1575 train_time:788ms step_avg:43.79ms step:19/1575 train_time:819ms step_avg:43.11ms step:20/1575 train_time:858ms step_avg:42.90ms step:21/1575 train_time:889ms step_avg:42.32ms step:22/1575 train_time:928ms step_avg:42.16ms step:23/1575 train_time:959ms step_avg:41.68ms step:24/1575 train_time:998ms step_avg:41.57ms step:25/1575 train_time:1029ms step_avg:41.14ms step:26/1575 train_time:1067ms step_avg:41.05ms step:27/1575 train_time:1098ms step_avg:40.68ms step:28/1575 train_time:1137ms step_avg:40.61ms step:29/1575 train_time:1168ms step_avg:40.26ms step:30/1575 train_time:1207ms step_avg:40.24ms step:31/1575 train_time:1238ms step_avg:39.92ms step:32/1575 train_time:1276ms step_avg:39.89ms step:33/1575 train_time:1307ms step_avg:39.61ms step:34/1575 train_time:1346ms step_avg:39.58ms step:35/1575 train_time:1377ms step_avg:39.34ms step:36/1575 train_time:1416ms step_avg:39.33ms step:37/1575 train_time:1447ms step_avg:39.10ms step:38/1575 train_time:1485ms step_avg:39.09ms step:39/1575 train_time:1516ms step_avg:38.88ms step:40/1575 train_time:1555ms step_avg:38.88ms step:41/1575 train_time:1586ms step_avg:38.69ms step:42/1575 train_time:1625ms step_avg:38.70ms step:43/1575 train_time:1656ms step_avg:38.52ms step:44/1575 train_time:1695ms step_avg:38.52ms step:45/1575 train_time:1726ms step_avg:38.35ms step:46/1575 train_time:1765ms step_avg:38.36ms step:47/1575 train_time:1796ms step_avg:38.22ms step:48/1575 train_time:1835ms step_avg:38.23ms step:49/1575 train_time:1866ms step_avg:38.08ms step:50/1575 train_time:1905ms step_avg:38.10ms step:51/1575 train_time:1936ms step_avg:37.96ms step:52/1575 train_time:1974ms step_avg:37.97ms step:53/1575 train_time:2006ms step_avg:37.84ms step:54/1575 train_time:2045ms step_avg:37.87ms step:55/1575 train_time:2076ms step_avg:37.75ms step:56/1575 train_time:2115ms step_avg:37.77ms step:57/1575 train_time:2146ms step_avg:37.64ms step:58/1575 train_time:2185ms step_avg:37.67ms step:59/1575 train_time:2216ms step_avg:37.56ms step:60/1575 train_time:2255ms step_avg:37.59ms step:61/1575 train_time:2286ms step_avg:37.48ms step:62/1575 train_time:2326ms step_avg:37.51ms step:63/1575 train_time:2356ms step_avg:37.40ms step:64/1575 train_time:2395ms step_avg:37.42ms step:65/1575 train_time:2426ms step_avg:37.32ms step:66/1575 train_time:2465ms step_avg:37.35ms step:67/1575 train_time:2496ms step_avg:37.25ms step:68/1575 train_time:2535ms step_avg:37.27ms step:69/1575 train_time:2566ms step_avg:37.18ms step:70/1575 train_time:2605ms step_avg:37.21ms step:71/1575 train_time:2636ms step_avg:37.12ms step:72/1575 train_time:2674ms step_avg:37.14ms step:73/1575 train_time:2705ms step_avg:37.06ms step:74/1575 train_time:2744ms step_avg:37.08ms step:75/1575 train_time:2775ms step_avg:36.99ms step:76/1575 train_time:2814ms step_avg:37.02ms step:77/1575 train_time:2845ms step_avg:36.94ms step:78/1575 train_time:2883ms step_avg:36.96ms step:79/1575 train_time:2914ms step_avg:36.89ms step:80/1575 train_time:2953ms step_avg:36.91ms step:81/1575 train_time:2984ms step_avg:36.84ms step:82/1575 train_time:3023ms step_avg:36.87ms step:83/1575 train_time:3054ms step_avg:36.80ms step:84/1575 train_time:3093ms step_avg:36.82ms step:85/1575 train_time:3124ms step_avg:36.75ms step:86/1575 train_time:3163ms step_avg:36.78ms step:87/1575 train_time:3194ms step_avg:36.71ms step:88/1575 train_time:3233ms step_avg:36.74ms step:89/1575 train_time:3264ms step_avg:36.67ms step:90/1575 train_time:3303ms step_avg:36.70ms step:91/1575 train_time:3334ms step_avg:36.64ms step:92/1575 train_time:3372ms step_avg:36.66ms step:93/1575 train_time:3404ms step_avg:36.60ms step:94/1575 train_time:3442ms step_avg:36.62ms step:95/1575 train_time:3473ms step_avg:36.56ms step:96/1575 train_time:3513ms step_avg:36.59ms step:97/1575 train_time:3543ms step_avg:36.53ms step:98/1575 train_time:3582ms step_avg:36.55ms step:99/1575 train_time:3613ms step_avg:36.50ms step:100/1575 train_time:3652ms step_avg:36.52ms step:101/1575 train_time:3683ms step_avg:36.46ms step:102/1575 train_time:3722ms step_avg:36.49ms step:103/1575 train_time:3752ms step_avg:36.43ms step:104/1575 train_time:3792ms step_avg:36.46ms step:105/1575 train_time:3823ms step_avg:36.41ms step:106/1575 train_time:3861ms step_avg:36.43ms step:107/1575 train_time:3892ms step_avg:36.38ms step:108/1575 train_time:3932ms step_avg:36.40ms step:109/1575 train_time:3963ms step_avg:36.35ms step:110/1575 train_time:4001ms step_avg:36.37ms step:111/1575 train_time:4032ms step_avg:36.32ms step:112/1575 train_time:4070ms step_avg:36.34ms step:113/1575 train_time:4101ms step_avg:36.30ms step:114/1575 train_time:4140ms step_avg:36.32ms step:115/1575 train_time:4171ms step_avg:36.27ms step:116/1575 train_time:4210ms step_avg:36.29ms step:117/1575 train_time:4241ms step_avg:36.24ms step:118/1575 train_time:4279ms step_avg:36.26ms step:119/1575 train_time:4310ms step_avg:36.22ms step:120/1575 train_time:4349ms step_avg:36.24ms step:121/1575 train_time:4380ms step_avg:36.20ms step:122/1575 train_time:4419ms step_avg:36.22ms step:123/1575 train_time:4449ms step_avg:36.17ms step:124/1575 train_time:4488ms step_avg:36.19ms step:125/1575 train_time:4519ms step_avg:36.15ms step:126/1575 train_time:4558ms step_avg:36.17ms step:127/1575 train_time:4589ms step_avg:36.13ms step:128/1575 train_time:4628ms step_avg:36.15ms step:129/1575 train_time:4658ms step_avg:36.11ms step:130/1575 train_time:4697ms step_avg:36.13ms step:131/1575 train_time:4728ms step_avg:36.09ms step:132/1575 train_time:4767ms step_avg:36.11ms step:133/1575 train_time:4798ms step_avg:36.07ms step:134/1575 train_time:4836ms step_avg:36.09ms step:135/1575 train_time:4867ms step_avg:36.05ms step:136/1575 train_time:4906ms step_avg:36.07ms step:137/1575 train_time:4937ms step_avg:36.04ms step:138/1575 train_time:4975ms step_avg:36.05ms step:139/1575 train_time:5006ms step_avg:36.01ms step:140/1575 train_time:5045ms step_avg:36.04ms step:141/1575 train_time:5076ms step_avg:36.00ms step:142/1575 train_time:5115ms step_avg:36.02ms step:143/1575 train_time:5145ms step_avg:35.98ms step:144/1575 train_time:5184ms step_avg:36.00ms step:145/1575 train_time:5215ms step_avg:35.96ms step:146/1575 train_time:5253ms step_avg:35.98ms step:147/1575 train_time:5284ms step_avg:35.95ms step:148/1575 train_time:5323ms step_avg:35.97ms step:149/1575 train_time:5354ms step_avg:35.94ms step:150/1575 train_time:5393ms step_avg:35.95ms step:151/1575 train_time:5424ms step_avg:35.92ms step:152/1575 train_time:5463ms step_avg:35.94ms step:153/1575 train_time:5494ms step_avg:35.91ms step:154/1575 train_time:5533ms step_avg:35.93ms step:155/1575 train_time:5563ms step_avg:35.89ms step:156/1575 train_time:5602ms step_avg:35.91ms step:157/1575 train_time:5633ms step_avg:35.88ms step:158/1575 train_time:5672ms step_avg:35.90ms step:159/1575 train_time:5703ms step_avg:35.87ms step:160/1575 train_time:5742ms step_avg:35.89ms step:161/1575 train_time:5772ms step_avg:35.85ms step:162/1575 train_time:5811ms step_avg:35.87ms step:163/1575 train_time:5842ms step_avg:35.84ms step:164/1575 train_time:5881ms step_avg:35.86ms step:165/1575 train_time:5911ms step_avg:35.83ms step:166/1575 train_time:5950ms step_avg:35.84ms step:167/1575 train_time:5981ms step_avg:35.81ms step:168/1575 train_time:6019ms step_avg:35.83ms step:169/1575 train_time:6050ms step_avg:35.80ms step:170/1575 train_time:6089ms step_avg:35.82ms step:171/1575 train_time:6120ms step_avg:35.79ms step:172/1575 train_time:6159ms step_avg:35.81ms step:173/1575 train_time:6190ms step_avg:35.78ms step:174/1575 train_time:6229ms step_avg:35.80ms step:175/1575 train_time:6259ms step_avg:35.77ms step:176/1575 train_time:6298ms step_avg:35.78ms step:177/1575 train_time:6328ms step_avg:35.75ms step:178/1575 train_time:6367ms step_avg:35.77ms step:179/1575 train_time:6398ms step_avg:35.74ms step:180/1575 train_time:6437ms step_avg:35.76ms step:181/1575 train_time:6467ms step_avg:35.73ms step:182/1575 train_time:6506ms step_avg:35.75ms step:183/1575 train_time:6537ms step_avg:35.72ms step:184/1575 train_time:6575ms step_avg:35.74ms step:185/1575 train_time:6606ms step_avg:35.71ms step:186/1575 train_time:6645ms step_avg:35.73ms step:187/1575 train_time:6675ms step_avg:35.70ms step:188/1575 train_time:6714ms step_avg:35.71ms step:189/1575 train_time:6745ms step_avg:35.69ms step:190/1575 train_time:6784ms step_avg:35.71ms step:191/1575 train_time:6815ms step_avg:35.68ms step:192/1575 train_time:6854ms step_avg:35.70ms step:193/1575 train_time:6885ms step_avg:35.67ms step:194/1575 train_time:6923ms step_avg:35.69ms step:195/1575 train_time:6954ms step_avg:35.66ms step:196/1575 train_time:6993ms step_avg:35.68ms step:197/1575 train_time:7023ms step_avg:35.65ms step:198/1575 train_time:7062ms step_avg:35.67ms step:199/1575 train_time:7093ms step_avg:35.64ms step:200/1575 train_time:7132ms step_avg:35.66ms step:201/1575 train_time:7162ms step_avg:35.63ms step:202/1575 train_time:7202ms step_avg:35.65ms step:203/1575 train_time:7232ms step_avg:35.63ms step:204/1575 train_time:7270ms step_avg:35.64ms step:205/1575 train_time:7301ms step_avg:35.62ms step:206/1575 train_time:7340ms step_avg:35.63ms step:207/1575 train_time:7371ms step_avg:35.61ms step:208/1575 train_time:7410ms step_avg:35.62ms step:209/1575 train_time:7440ms step_avg:35.60ms step:210/1575 train_time:7479ms step_avg:35.61ms step:211/1575 train_time:7510ms step_avg:35.59ms step:212/1575 train_time:7548ms step_avg:35.61ms step:213/1575 train_time:7579ms step_avg:35.58ms step:214/1575 train_time:7618ms step_avg:35.60ms step:215/1575 train_time:7649ms step_avg:35.58ms step:216/1575 train_time:7688ms step_avg:35.59ms step:217/1575 train_time:7719ms step_avg:35.57ms step:218/1575 train_time:7757ms step_avg:35.58ms step:219/1575 train_time:7788ms step_avg:35.56ms step:220/1575 train_time:7827ms step_avg:35.58ms step:221/1575 train_time:7858ms step_avg:35.56ms step:222/1575 train_time:7896ms step_avg:35.57ms step:223/1575 train_time:7927ms step_avg:35.55ms step:224/1575 train_time:7966ms step_avg:35.56ms step:225/1575 train_time:7997ms step_avg:35.54ms step:226/1575 train_time:8036ms step_avg:35.56ms step:227/1575 train_time:8066ms step_avg:35.53ms step:228/1575 train_time:8105ms step_avg:35.55ms step:229/1575 train_time:8136ms step_avg:35.53ms step:230/1575 train_time:8175ms step_avg:35.54ms step:231/1575 train_time:8206ms step_avg:35.52ms step:232/1575 train_time:8244ms step_avg:35.54ms step:233/1575 train_time:8275ms step_avg:35.51ms step:234/1575 train_time:8314ms step_avg:35.53ms step:235/1575 train_time:8344ms step_avg:35.51ms step:236/1575 train_time:8383ms step_avg:35.52ms step:237/1575 train_time:8414ms step_avg:35.50ms step:238/1575 train_time:8453ms step_avg:35.52ms step:239/1575 train_time:8484ms step_avg:35.50ms step:240/1575 train_time:8523ms step_avg:35.51ms step:241/1575 train_time:8554ms step_avg:35.49ms step:242/1575 train_time:8592ms step_avg:35.51ms step:243/1575 train_time:8623ms step_avg:35.48ms step:244/1575 train_time:8662ms step_avg:35.50ms step:245/1575 train_time:8693ms step_avg:35.48ms step:246/1575 train_time:8731ms step_avg:35.49ms step:247/1575 train_time:8762ms step_avg:35.47ms step:248/1575 train_time:8801ms step_avg:35.49ms step:249/1575 train_time:8832ms step_avg:35.47ms step:250/1575 train_time:8871ms step_avg:35.48ms step:250/1575 val_loss:4.5825 train_time:8920ms step_avg:35.68ms step:251/1575 train_time:8940ms step_avg:35.62ms step:252/1575 train_time:8960ms step_avg:35.56ms step:253/1575 train_time:8978ms step_avg:35.49ms step:254/1575 train_time:9016ms step_avg:35.50ms step:255/1575 train_time:9049ms step_avg:35.49ms step:256/1575 train_time:9089ms step_avg:35.50ms step:257/1575 train_time:9120ms step_avg:35.49ms step:258/1575 train_time:9159ms step_avg:35.50ms step:259/1575 train_time:9189ms step_avg:35.48ms step:260/1575 train_time:9228ms step_avg:35.49ms step:261/1575 train_time:9259ms step_avg:35.48ms step:262/1575 train_time:9298ms step_avg:35.49ms step:263/1575 train_time:9329ms step_avg:35.47ms step:264/1575 train_time:9367ms step_avg:35.48ms step:265/1575 train_time:9398ms step_avg:35.46ms step:266/1575 train_time:9437ms step_avg:35.48ms step:267/1575 train_time:9467ms step_avg:35.46ms step:268/1575 train_time:9506ms step_avg:35.47ms step:269/1575 train_time:9536ms step_avg:35.45ms step:270/1575 train_time:9575ms step_avg:35.46ms step:271/1575 train_time:9605ms step_avg:35.44ms step:272/1575 train_time:9644ms step_avg:35.46ms step:273/1575 train_time:9675ms step_avg:35.44ms step:274/1575 train_time:9713ms step_avg:35.45ms step:275/1575 train_time:9744ms step_avg:35.43ms step:276/1575 train_time:9782ms step_avg:35.44ms step:277/1575 train_time:9813ms step_avg:35.43ms step:278/1575 train_time:9852ms step_avg:35.44ms step:279/1575 train_time:9882ms step_avg:35.42ms step:280/1575 train_time:9921ms step_avg:35.43ms step:281/1575 train_time:9951ms step_avg:35.41ms step:282/1575 train_time:9991ms step_avg:35.43ms step:283/1575 train_time:10021ms step_avg:35.41ms step:284/1575 train_time:10060ms step_avg:35.42ms step:285/1575 train_time:10090ms step_avg:35.40ms step:286/1575 train_time:10129ms step_avg:35.42ms step:287/1575 train_time:10160ms step_avg:35.40ms step:288/1575 train_time:10199ms step_avg:35.41ms step:289/1575 train_time:10230ms step_avg:35.40ms step:290/1575 train_time:10269ms step_avg:35.41ms step:291/1575 train_time:10299ms step_avg:35.39ms step:292/1575 train_time:10338ms step_avg:35.40ms step:293/1575 train_time:10369ms step_avg:35.39ms step:294/1575 train_time:10408ms step_avg:35.40ms step:295/1575 train_time:10438ms step_avg:35.38ms step:296/1575 train_time:10477ms step_avg:35.40ms step:297/1575 train_time:10508ms step_avg:35.38ms step:298/1575 train_time:10546ms step_avg:35.39ms step:299/1575 train_time:10578ms step_avg:35.38ms step:300/1575 train_time:10617ms step_avg:35.39ms step:301/1575 train_time:10647ms step_avg:35.37ms step:302/1575 train_time:10686ms step_avg:35.38ms step:303/1575 train_time:10717ms step_avg:35.37ms step:304/1575 train_time:10755ms step_avg:35.38ms step:305/1575 train_time:10786ms step_avg:35.36ms step:306/1575 train_time:10825ms step_avg:35.38ms step:307/1575 train_time:10856ms step_avg:35.36ms step:308/1575 train_time:10895ms step_avg:35.37ms step:309/1575 train_time:10926ms step_avg:35.36ms step:310/1575 train_time:10965ms step_avg:35.37ms step:311/1575 train_time:10996ms step_avg:35.36ms step:312/1575 train_time:11035ms step_avg:35.37ms step:313/1575 train_time:11066ms step_avg:35.35ms step:314/1575 train_time:11104ms step_avg:35.36ms step:315/1575 train_time:11135ms step_avg:35.35ms step:316/1575 train_time:11173ms step_avg:35.36ms step:317/1575 train_time:11204ms step_avg:35.34ms step:318/1575 train_time:11243ms step_avg:35.36ms step:319/1575 train_time:11274ms step_avg:35.34ms step:320/1575 train_time:11313ms step_avg:35.35ms step:321/1575 train_time:11343ms step_avg:35.34ms step:322/1575 train_time:11382ms step_avg:35.35ms step:323/1575 train_time:11412ms step_avg:35.33ms step:324/1575 train_time:11451ms step_avg:35.34ms step:325/1575 train_time:11482ms step_avg:35.33ms step:326/1575 train_time:11521ms step_avg:35.34ms step:327/1575 train_time:11552ms step_avg:35.33ms step:328/1575 train_time:11590ms step_avg:35.34ms step:329/1575 train_time:11621ms step_avg:35.32ms step:330/1575 train_time:11660ms step_avg:35.33ms step:331/1575 train_time:11690ms step_avg:35.32ms step:332/1575 train_time:11729ms step_avg:35.33ms step:333/1575 train_time:11760ms step_avg:35.31ms step:334/1575 train_time:11798ms step_avg:35.32ms step:335/1575 train_time:11829ms step_avg:35.31ms step:336/1575 train_time:11868ms step_avg:35.32ms step:337/1575 train_time:11899ms step_avg:35.31ms step:338/1575 train_time:11937ms step_avg:35.32ms step:339/1575 train_time:11968ms step_avg:35.30ms step:340/1575 train_time:12007ms step_avg:35.31ms step:341/1575 train_time:12038ms step_avg:35.30ms step:342/1575 train_time:12077ms step_avg:35.31ms step:343/1575 train_time:12108ms step_avg:35.30ms step:344/1575 train_time:12146ms step_avg:35.31ms step:345/1575 train_time:12177ms step_avg:35.30ms step:346/1575 train_time:12215ms step_avg:35.30ms step:347/1575 train_time:12246ms step_avg:35.29ms step:348/1575 train_time:12284ms step_avg:35.30ms step:349/1575 train_time:12315ms step_avg:35.29ms step:350/1575 train_time:12354ms step_avg:35.30ms step:351/1575 train_time:12385ms step_avg:35.28ms step:352/1575 train_time:12423ms step_avg:35.29ms step:353/1575 train_time:12454ms step_avg:35.28ms step:354/1575 train_time:12492ms step_avg:35.29ms step:355/1575 train_time:12523ms step_avg:35.28ms step:356/1575 train_time:12562ms step_avg:35.29ms step:357/1575 train_time:12592ms step_avg:35.27ms step:358/1575 train_time:12631ms step_avg:35.28ms step:359/1575 train_time:12662ms step_avg:35.27ms step:360/1575 train_time:12700ms step_avg:35.28ms step:361/1575 train_time:12731ms step_avg:35.27ms step:362/1575 train_time:12769ms step_avg:35.27ms step:363/1575 train_time:12800ms step_avg:35.26ms step:364/1575 train_time:12838ms step_avg:35.27ms step:365/1575 train_time:12869ms step_avg:35.26ms step:366/1575 train_time:12908ms step_avg:35.27ms step:367/1575 train_time:12938ms step_avg:35.25ms step:368/1575 train_time:12977ms step_avg:35.26ms step:369/1575 train_time:13008ms step_avg:35.25ms step:370/1575 train_time:13046ms step_avg:35.26ms step:371/1575 train_time:13077ms step_avg:35.25ms step:372/1575 train_time:13116ms step_avg:35.26ms step:373/1575 train_time:13146ms step_avg:35.24ms step:374/1575 train_time:13185ms step_avg:35.25ms step:375/1575 train_time:13215ms step_avg:35.24ms step:376/1575 train_time:13254ms step_avg:35.25ms step:377/1575 train_time:13285ms step_avg:35.24ms step:378/1575 train_time:13324ms step_avg:35.25ms step:379/1575 train_time:13355ms step_avg:35.24ms step:380/1575 train_time:13394ms step_avg:35.25ms step:381/1575 train_time:13424ms step_avg:35.23ms step:382/1575 train_time:13462ms step_avg:35.24ms step:383/1575 train_time:13493ms step_avg:35.23ms step:384/1575 train_time:13532ms step_avg:35.24ms step:385/1575 train_time:13563ms step_avg:35.23ms step:386/1575 train_time:13601ms step_avg:35.24ms step:387/1575 train_time:13632ms step_avg:35.22ms step:388/1575 train_time:13670ms step_avg:35.23ms step:389/1575 train_time:13701ms step_avg:35.22ms step:390/1575 train_time:13740ms step_avg:35.23ms step:391/1575 train_time:13770ms step_avg:35.22ms step:392/1575 train_time:13809ms step_avg:35.23ms step:393/1575 train_time:13840ms step_avg:35.22ms step:394/1575 train_time:13878ms step_avg:35.22ms step:395/1575 train_time:13909ms step_avg:35.21ms step:396/1575 train_time:13947ms step_avg:35.22ms step:397/1575 train_time:13978ms step_avg:35.21ms step:398/1575 train_time:14017ms step_avg:35.22ms step:399/1575 train_time:14048ms step_avg:35.21ms step:400/1575 train_time:14086ms step_avg:35.22ms step:401/1575 train_time:14117ms step_avg:35.20ms step:402/1575 train_time:14155ms step_avg:35.21ms step:403/1575 train_time:14186ms step_avg:35.20ms step:404/1575 train_time:14225ms step_avg:35.21ms step:405/1575 train_time:14256ms step_avg:35.20ms step:406/1575 train_time:14295ms step_avg:35.21ms step:407/1575 train_time:14325ms step_avg:35.20ms step:408/1575 train_time:14363ms step_avg:35.20ms step:409/1575 train_time:14394ms step_avg:35.19ms step:410/1575 train_time:14433ms step_avg:35.20ms step:411/1575 train_time:14463ms step_avg:35.19ms step:412/1575 train_time:14502ms step_avg:35.20ms step:413/1575 train_time:14533ms step_avg:35.19ms step:414/1575 train_time:14572ms step_avg:35.20ms step:415/1575 train_time:14603ms step_avg:35.19ms step:416/1575 train_time:14641ms step_avg:35.19ms step:417/1575 train_time:14672ms step_avg:35.18ms step:418/1575 train_time:14710ms step_avg:35.19ms step:419/1575 train_time:14741ms step_avg:35.18ms step:420/1575 train_time:14780ms step_avg:35.19ms step:421/1575 train_time:14810ms step_avg:35.18ms step:422/1575 train_time:14849ms step_avg:35.19ms step:423/1575 train_time:14880ms step_avg:35.18ms step:424/1575 train_time:14919ms step_avg:35.19ms step:425/1575 train_time:14949ms step_avg:35.17ms step:426/1575 train_time:14988ms step_avg:35.18ms step:427/1575 train_time:15019ms step_avg:35.17ms step:428/1575 train_time:15057ms step_avg:35.18ms step:429/1575 train_time:15088ms step_avg:35.17ms step:430/1575 train_time:15127ms step_avg:35.18ms step:431/1575 train_time:15157ms step_avg:35.17ms step:432/1575 train_time:15196ms step_avg:35.18ms step:433/1575 train_time:15227ms step_avg:35.17ms step:434/1575 train_time:15266ms step_avg:35.17ms step:435/1575 train_time:15296ms step_avg:35.16ms step:436/1575 train_time:15335ms step_avg:35.17ms step:437/1575 train_time:15366ms step_avg:35.16ms step:438/1575 train_time:15404ms step_avg:35.17ms step:439/1575 train_time:15435ms step_avg:35.16ms step:440/1575 train_time:15473ms step_avg:35.17ms step:441/1575 train_time:15504ms step_avg:35.16ms step:442/1575 train_time:15543ms step_avg:35.16ms step:443/1575 train_time:15573ms step_avg:35.15ms step:444/1575 train_time:15612ms step_avg:35.16ms step:445/1575 train_time:15643ms step_avg:35.15ms step:446/1575 train_time:15682ms step_avg:35.16ms step:447/1575 train_time:15713ms step_avg:35.15ms step:448/1575 train_time:15751ms step_avg:35.16ms step:449/1575 train_time:15782ms step_avg:35.15ms step:450/1575 train_time:15820ms step_avg:35.16ms step:451/1575 train_time:15851ms step_avg:35.15ms step:452/1575 train_time:15890ms step_avg:35.15ms step:453/1575 train_time:15921ms step_avg:35.15ms step:454/1575 train_time:15959ms step_avg:35.15ms step:455/1575 train_time:15990ms step_avg:35.14ms step:456/1575 train_time:16029ms step_avg:35.15ms step:457/1575 train_time:16060ms step_avg:35.14ms step:458/1575 train_time:16099ms step_avg:35.15ms step:459/1575 train_time:16129ms step_avg:35.14ms step:460/1575 train_time:16167ms step_avg:35.15ms step:461/1575 train_time:16198ms step_avg:35.14ms step:462/1575 train_time:16237ms step_avg:35.14ms step:463/1575 train_time:16267ms step_avg:35.13ms step:464/1575 train_time:16307ms step_avg:35.14ms step:465/1575 train_time:16338ms step_avg:35.14ms step:466/1575 train_time:16377ms step_avg:35.14ms step:467/1575 train_time:16408ms step_avg:35.13ms step:468/1575 train_time:16447ms step_avg:35.14ms step:469/1575 train_time:16478ms step_avg:35.13ms step:470/1575 train_time:16517ms step_avg:35.14ms step:471/1575 train_time:16547ms step_avg:35.13ms step:472/1575 train_time:16586ms step_avg:35.14ms step:473/1575 train_time:16617ms step_avg:35.13ms step:474/1575 train_time:16656ms step_avg:35.14ms step:475/1575 train_time:16687ms step_avg:35.13ms step:476/1575 train_time:16725ms step_avg:35.14ms step:477/1575 train_time:16756ms step_avg:35.13ms step:478/1575 train_time:16795ms step_avg:35.14ms step:479/1575 train_time:16826ms step_avg:35.13ms step:480/1575 train_time:16864ms step_avg:35.13ms step:481/1575 train_time:16895ms step_avg:35.12ms step:482/1575 train_time:16934ms step_avg:35.13ms step:483/1575 train_time:16964ms step_avg:35.12ms step:484/1575 train_time:17003ms step_avg:35.13ms step:485/1575 train_time:17034ms step_avg:35.12ms step:486/1575 train_time:17073ms step_avg:35.13ms step:487/1575 train_time:17104ms step_avg:35.12ms step:488/1575 train_time:17142ms step_avg:35.13ms step:489/1575 train_time:17173ms step_avg:35.12ms step:490/1575 train_time:17211ms step_avg:35.13ms step:491/1575 train_time:17242ms step_avg:35.12ms step:492/1575 train_time:17280ms step_avg:35.12ms step:493/1575 train_time:17311ms step_avg:35.11ms step:494/1575 train_time:17350ms step_avg:35.12ms step:495/1575 train_time:17381ms step_avg:35.11ms step:496/1575 train_time:17419ms step_avg:35.12ms step:497/1575 train_time:17450ms step_avg:35.11ms step:498/1575 train_time:17489ms step_avg:35.12ms step:499/1575 train_time:17520ms step_avg:35.11ms step:500/1575 train_time:17559ms step_avg:35.12ms step:500/1575 val_loss:4.2456 train_time:17607ms step_avg:35.21ms step:501/1575 train_time:17627ms step_avg:35.18ms step:502/1575 train_time:17647ms step_avg:35.15ms step:503/1575 train_time:17665ms step_avg:35.12ms step:504/1575 train_time:17701ms step_avg:35.12ms step:505/1575 train_time:17733ms step_avg:35.11ms step:506/1575 train_time:17772ms step_avg:35.12ms step:507/1575 train_time:17803ms step_avg:35.12ms step:508/1575 train_time:17842ms step_avg:35.12ms step:509/1575 train_time:17873ms step_avg:35.11ms step:510/1575 train_time:17912ms step_avg:35.12ms step:511/1575 train_time:17943ms step_avg:35.11ms step:512/1575 train_time:17982ms step_avg:35.12ms step:513/1575 train_time:18056ms step_avg:35.20ms step:514/1575 train_time:18112ms step_avg:35.24ms step:515/1575 train_time:18174ms step_avg:35.29ms step:516/1575 train_time:18232ms step_avg:35.33ms step:517/1575 train_time:18295ms step_avg:35.39ms step:518/1575 train_time:18354ms step_avg:35.43ms step:519/1575 train_time:18415ms step_avg:35.48ms step:520/1575 train_time:18474ms step_avg:35.53ms step:521/1575 train_time:18537ms step_avg:35.58ms step:522/1575 train_time:18596ms step_avg:35.62ms step:523/1575 train_time:18661ms step_avg:35.68ms step:524/1575 train_time:18722ms step_avg:35.73ms step:525/1575 train_time:18788ms step_avg:35.79ms step:526/1575 train_time:18846ms step_avg:35.83ms step:527/1575 train_time:18910ms step_avg:35.88ms step:528/1575 train_time:18970ms step_avg:35.93ms step:529/1575 train_time:19033ms step_avg:35.98ms step:530/1575 train_time:19093ms step_avg:36.02ms step:531/1575 train_time:19156ms step_avg:36.08ms step:532/1575 train_time:19214ms step_avg:36.12ms step:533/1575 train_time:19277ms step_avg:36.17ms step:534/1575 train_time:19336ms step_avg:36.21ms step:535/1575 train_time:19399ms step_avg:36.26ms step:536/1575 train_time:19458ms step_avg:36.30ms step:537/1575 train_time:19521ms step_avg:36.35ms step:538/1575 train_time:19580ms step_avg:36.39ms step:539/1575 train_time:19644ms step_avg:36.45ms step:540/1575 train_time:19704ms step_avg:36.49ms step:541/1575 train_time:19769ms step_avg:36.54ms step:542/1575 train_time:19828ms step_avg:36.58ms step:543/1575 train_time:19892ms step_avg:36.63ms step:544/1575 train_time:19952ms step_avg:36.68ms step:545/1575 train_time:20015ms step_avg:36.72ms step:546/1575 train_time:20075ms step_avg:36.77ms step:547/1575 train_time:20138ms step_avg:36.82ms step:548/1575 train_time:20198ms step_avg:36.86ms step:549/1575 train_time:20261ms step_avg:36.91ms step:550/1575 train_time:20321ms step_avg:36.95ms step:551/1575 train_time:20384ms step_avg:36.99ms step:552/1575 train_time:20445ms step_avg:37.04ms step:553/1575 train_time:20507ms step_avg:37.08ms step:554/1575 train_time:20566ms step_avg:37.12ms step:555/1575 train_time:20628ms step_avg:37.17ms step:556/1575 train_time:20687ms step_avg:37.21ms step:557/1575 train_time:20750ms step_avg:37.25ms step:558/1575 train_time:20810ms step_avg:37.29ms step:559/1575 train_time:20874ms step_avg:37.34ms step:560/1575 train_time:20932ms step_avg:37.38ms step:561/1575 train_time:20995ms step_avg:37.42ms step:562/1575 train_time:21055ms step_avg:37.46ms step:563/1575 train_time:21118ms step_avg:37.51ms step:564/1575 train_time:21178ms step_avg:37.55ms step:565/1575 train_time:21240ms step_avg:37.59ms step:566/1575 train_time:21300ms step_avg:37.63ms step:567/1575 train_time:21364ms step_avg:37.68ms step:568/1575 train_time:21423ms step_avg:37.72ms step:569/1575 train_time:21491ms step_avg:37.77ms step:570/1575 train_time:21547ms step_avg:37.80ms step:571/1575 train_time:21610ms step_avg:37.85ms step:572/1575 train_time:21670ms step_avg:37.88ms step:573/1575 train_time:21732ms step_avg:37.93ms step:574/1575 train_time:21791ms step_avg:37.96ms step:575/1575 train_time:21854ms step_avg:38.01ms step:576/1575 train_time:21913ms step_avg:38.04ms step:577/1575 train_time:21977ms step_avg:38.09ms step:578/1575 train_time:22037ms step_avg:38.13ms step:579/1575 train_time:22100ms step_avg:38.17ms step:580/1575 train_time:22160ms step_avg:38.21ms step:581/1575 train_time:22223ms step_avg:38.25ms step:582/1575 train_time:22282ms step_avg:38.29ms step:583/1575 train_time:22346ms step_avg:38.33ms step:584/1575 train_time:22406ms step_avg:38.37ms step:585/1575 train_time:22468ms step_avg:38.41ms step:586/1575 train_time:22527ms step_avg:38.44ms step:587/1575 train_time:22590ms step_avg:38.48ms step:588/1575 train_time:22650ms step_avg:38.52ms step:589/1575 train_time:22714ms step_avg:38.56ms step:590/1575 train_time:22772ms step_avg:38.60ms step:591/1575 train_time:22835ms step_avg:38.64ms step:592/1575 train_time:22895ms step_avg:38.67ms step:593/1575 train_time:22958ms step_avg:38.71ms step:594/1575 train_time:23017ms step_avg:38.75ms step:595/1575 train_time:23080ms step_avg:38.79ms step:596/1575 train_time:23139ms step_avg:38.82ms step:597/1575 train_time:23203ms step_avg:38.87ms step:598/1575 train_time:23263ms step_avg:38.90ms step:599/1575 train_time:23326ms step_avg:38.94ms step:600/1575 train_time:23385ms step_avg:38.98ms step:601/1575 train_time:23448ms step_avg:39.02ms step:602/1575 train_time:23508ms step_avg:39.05ms step:603/1575 train_time:23571ms step_avg:39.09ms step:604/1575 train_time:23630ms step_avg:39.12ms step:605/1575 train_time:23693ms step_avg:39.16ms step:606/1575 train_time:23752ms step_avg:39.19ms step:607/1575 train_time:23815ms step_avg:39.23ms step:608/1575 train_time:23876ms step_avg:39.27ms step:609/1575 train_time:23939ms step_avg:39.31ms step:610/1575 train_time:23997ms step_avg:39.34ms step:611/1575 train_time:24060ms step_avg:39.38ms step:612/1575 train_time:24120ms step_avg:39.41ms step:613/1575 train_time:24183ms step_avg:39.45ms step:614/1575 train_time:24243ms step_avg:39.48ms step:615/1575 train_time:24307ms step_avg:39.52ms step:616/1575 train_time:24365ms step_avg:39.55ms step:617/1575 train_time:24428ms step_avg:39.59ms step:618/1575 train_time:24488ms step_avg:39.62ms step:619/1575 train_time:24552ms step_avg:39.66ms step:620/1575 train_time:24612ms step_avg:39.70ms step:621/1575 train_time:24675ms step_avg:39.73ms step:622/1575 train_time:24734ms step_avg:39.77ms step:623/1575 train_time:24797ms step_avg:39.80ms step:624/1575 train_time:24857ms step_avg:39.84ms step:625/1575 train_time:24920ms step_avg:39.87ms step:626/1575 train_time:24979ms step_avg:39.90ms step:627/1575 train_time:25043ms step_avg:39.94ms step:628/1575 train_time:25102ms step_avg:39.97ms step:629/1575 train_time:25165ms step_avg:40.01ms step:630/1575 train_time:25224ms step_avg:40.04ms step:631/1575 train_time:25288ms step_avg:40.08ms step:632/1575 train_time:25347ms step_avg:40.11ms step:633/1575 train_time:25411ms step_avg:40.14ms step:634/1575 train_time:25471ms step_avg:40.17ms step:635/1575 train_time:25533ms step_avg:40.21ms step:636/1575 train_time:25593ms step_avg:40.24ms step:637/1575 train_time:25656ms step_avg:40.28ms step:638/1575 train_time:25715ms step_avg:40.31ms step:639/1575 train_time:25778ms step_avg:40.34ms step:640/1575 train_time:25837ms step_avg:40.37ms step:641/1575 train_time:25900ms step_avg:40.41ms step:642/1575 train_time:25959ms step_avg:40.44ms step:643/1575 train_time:26023ms step_avg:40.47ms step:644/1575 train_time:26083ms step_avg:40.50ms step:645/1575 train_time:26146ms step_avg:40.54ms step:646/1575 train_time:26205ms step_avg:40.57ms step:647/1575 train_time:26269ms step_avg:40.60ms step:648/1575 train_time:26328ms step_avg:40.63ms step:649/1575 train_time:26391ms step_avg:40.66ms step:650/1575 train_time:26451ms step_avg:40.69ms step:651/1575 train_time:26514ms step_avg:40.73ms step:652/1575 train_time:26573ms step_avg:40.76ms step:653/1575 train_time:26636ms step_avg:40.79ms step:654/1575 train_time:26696ms step_avg:40.82ms step:655/1575 train_time:26759ms step_avg:40.85ms step:656/1575 train_time:26818ms step_avg:40.88ms step:657/1575 train_time:26882ms step_avg:40.92ms step:658/1575 train_time:26940ms step_avg:40.94ms step:659/1575 train_time:27003ms step_avg:40.98ms step:660/1575 train_time:27062ms step_avg:41.00ms step:661/1575 train_time:27126ms step_avg:41.04ms step:662/1575 train_time:27185ms step_avg:41.06ms step:663/1575 train_time:27248ms step_avg:41.10ms step:664/1575 train_time:27308ms step_avg:41.13ms step:665/1575 train_time:27374ms step_avg:41.16ms step:666/1575 train_time:27431ms step_avg:41.19ms step:667/1575 train_time:27494ms step_avg:41.22ms step:668/1575 train_time:27552ms step_avg:41.25ms step:669/1575 train_time:27616ms step_avg:41.28ms step:670/1575 train_time:27675ms step_avg:41.31ms step:671/1575 train_time:27738ms step_avg:41.34ms step:672/1575 train_time:27797ms step_avg:41.36ms step:673/1575 train_time:27859ms step_avg:41.40ms step:674/1575 train_time:27919ms step_avg:41.42ms step:675/1575 train_time:27983ms step_avg:41.46ms step:676/1575 train_time:28042ms step_avg:41.48ms step:677/1575 train_time:28105ms step_avg:41.51ms step:678/1575 train_time:28165ms step_avg:41.54ms step:679/1575 train_time:28228ms step_avg:41.57ms step:680/1575 train_time:28288ms step_avg:41.60ms step:681/1575 train_time:28352ms step_avg:41.63ms step:682/1575 train_time:28411ms step_avg:41.66ms step:683/1575 train_time:28474ms step_avg:41.69ms step:684/1575 train_time:28534ms step_avg:41.72ms step:685/1575 train_time:28597ms step_avg:41.75ms step:686/1575 train_time:28655ms step_avg:41.77ms step:687/1575 train_time:28719ms step_avg:41.80ms step:688/1575 train_time:28778ms step_avg:41.83ms step:689/1575 train_time:28841ms step_avg:41.86ms step:690/1575 train_time:28902ms step_avg:41.89ms step:691/1575 train_time:28965ms step_avg:41.92ms step:692/1575 train_time:29025ms step_avg:41.94ms step:693/1575 train_time:29088ms step_avg:41.97ms step:694/1575 train_time:29146ms step_avg:42.00ms step:695/1575 train_time:29210ms step_avg:42.03ms step:696/1575 train_time:29269ms step_avg:42.05ms step:697/1575 train_time:29332ms step_avg:42.08ms step:698/1575 train_time:29391ms step_avg:42.11ms step:699/1575 train_time:29454ms step_avg:42.14ms step:700/1575 train_time:29514ms step_avg:42.16ms step:701/1575 train_time:29577ms step_avg:42.19ms step:702/1575 train_time:29636ms step_avg:42.22ms step:703/1575 train_time:29699ms step_avg:42.25ms step:704/1575 train_time:29758ms step_avg:42.27ms step:705/1575 train_time:29822ms step_avg:42.30ms step:706/1575 train_time:29882ms step_avg:42.33ms step:707/1575 train_time:29945ms step_avg:42.35ms step:708/1575 train_time:30004ms step_avg:42.38ms step:709/1575 train_time:30068ms step_avg:42.41ms step:710/1575 train_time:30128ms step_avg:42.43ms step:711/1575 train_time:30191ms step_avg:42.46ms step:712/1575 train_time:30250ms step_avg:42.49ms step:713/1575 train_time:30313ms step_avg:42.51ms step:714/1575 train_time:30374ms step_avg:42.54ms step:715/1575 train_time:30436ms step_avg:42.57ms step:716/1575 train_time:30495ms step_avg:42.59ms step:717/1575 train_time:30559ms step_avg:42.62ms step:718/1575 train_time:30617ms step_avg:42.64ms step:719/1575 train_time:30681ms step_avg:42.67ms step:720/1575 train_time:30739ms step_avg:42.69ms step:721/1575 train_time:30802ms step_avg:42.72ms step:722/1575 train_time:30862ms step_avg:42.74ms step:723/1575 train_time:30925ms step_avg:42.77ms step:724/1575 train_time:30984ms step_avg:42.80ms step:725/1575 train_time:31047ms step_avg:42.82ms step:726/1575 train_time:31106ms step_avg:42.85ms step:727/1575 train_time:31170ms step_avg:42.87ms step:728/1575 train_time:31230ms step_avg:42.90ms step:729/1575 train_time:31292ms step_avg:42.93ms step:730/1575 train_time:31352ms step_avg:42.95ms step:731/1575 train_time:31415ms step_avg:42.98ms step:732/1575 train_time:31475ms step_avg:43.00ms step:733/1575 train_time:31538ms step_avg:43.03ms step:734/1575 train_time:31597ms step_avg:43.05ms step:735/1575 train_time:31660ms step_avg:43.07ms step:736/1575 train_time:31719ms step_avg:43.10ms step:737/1575 train_time:31784ms step_avg:43.13ms step:738/1575 train_time:31843ms step_avg:43.15ms step:739/1575 train_time:31906ms step_avg:43.17ms step:740/1575 train_time:31965ms step_avg:43.20ms step:741/1575 train_time:32029ms step_avg:43.22ms step:742/1575 train_time:32092ms step_avg:43.25ms step:743/1575 train_time:32151ms step_avg:43.27ms step:744/1575 train_time:32211ms step_avg:43.29ms step:745/1575 train_time:32274ms step_avg:43.32ms step:746/1575 train_time:32333ms step_avg:43.34ms step:747/1575 train_time:32396ms step_avg:43.37ms step:748/1575 train_time:32455ms step_avg:43.39ms step:749/1575 train_time:32518ms step_avg:43.42ms step:750/1575 train_time:32578ms step_avg:43.44ms step:750/1575 val_loss:3.8849 train_time:32624ms step_avg:43.50ms step:751/1575 train_time:32644ms step_avg:43.47ms step:752/1575 train_time:32704ms step_avg:43.49ms step:753/1575 train_time:32771ms step_avg:43.52ms step:754/1575 train_time:32834ms step_avg:43.55ms step:755/1575 train_time:32897ms step_avg:43.57ms step:756/1575 train_time:32957ms step_avg:43.59ms step:757/1575 train_time:33021ms step_avg:43.62ms step:758/1575 train_time:33079ms step_avg:43.64ms step:759/1575 train_time:33141ms step_avg:43.66ms step:760/1575 train_time:33200ms step_avg:43.68ms step:761/1575 train_time:33264ms step_avg:43.71ms step:762/1575 train_time:33323ms step_avg:43.73ms step:763/1575 train_time:33384ms step_avg:43.75ms step:764/1575 train_time:33443ms step_avg:43.77ms step:765/1575 train_time:33505ms step_avg:43.80ms step:766/1575 train_time:33564ms step_avg:43.82ms step:767/1575 train_time:33628ms step_avg:43.84ms step:768/1575 train_time:33689ms step_avg:43.87ms step:769/1575 train_time:33755ms step_avg:43.89ms step:770/1575 train_time:33815ms step_avg:43.92ms step:771/1575 train_time:33878ms step_avg:43.94ms step:772/1575 train_time:33938ms step_avg:43.96ms step:773/1575 train_time:34001ms step_avg:43.99ms step:774/1575 train_time:34060ms step_avg:44.01ms step:775/1575 train_time:34123ms step_avg:44.03ms step:776/1575 train_time:34182ms step_avg:44.05ms step:777/1575 train_time:34245ms step_avg:44.07ms step:778/1575 train_time:34303ms step_avg:44.09ms step:779/1575 train_time:34366ms step_avg:44.12ms step:780/1575 train_time:34425ms step_avg:44.13ms step:781/1575 train_time:34487ms step_avg:44.16ms step:782/1575 train_time:34549ms step_avg:44.18ms step:783/1575 train_time:34612ms step_avg:44.20ms step:784/1575 train_time:34672ms step_avg:44.22ms step:785/1575 train_time:34735ms step_avg:44.25ms step:786/1575 train_time:34795ms step_avg:44.27ms step:787/1575 train_time:34858ms step_avg:44.29ms step:788/1575 train_time:34918ms step_avg:44.31ms step:789/1575 train_time:34981ms step_avg:44.34ms step:790/1575 train_time:35041ms step_avg:44.36ms step:791/1575 train_time:35104ms step_avg:44.38ms step:792/1575 train_time:35163ms step_avg:44.40ms step:793/1575 train_time:35226ms step_avg:44.42ms step:794/1575 train_time:35285ms step_avg:44.44ms step:795/1575 train_time:35348ms step_avg:44.46ms step:796/1575 train_time:35407ms step_avg:44.48ms step:797/1575 train_time:35470ms step_avg:44.50ms step:798/1575 train_time:35530ms step_avg:44.52ms step:799/1575 train_time:35592ms step_avg:44.55ms step:800/1575 train_time:35652ms step_avg:44.57ms step:801/1575 train_time:35716ms step_avg:44.59ms step:802/1575 train_time:35781ms step_avg:44.62ms step:803/1575 train_time:35840ms step_avg:44.63ms step:804/1575 train_time:35901ms step_avg:44.65ms step:805/1575 train_time:35964ms step_avg:44.68ms step:806/1575 train_time:36024ms step_avg:44.69ms step:807/1575 train_time:36087ms step_avg:44.72ms step:808/1575 train_time:36146ms step_avg:44.73ms step:809/1575 train_time:36209ms step_avg:44.76ms step:810/1575 train_time:36268ms step_avg:44.78ms step:811/1575 train_time:36332ms step_avg:44.80ms step:812/1575 train_time:36391ms step_avg:44.82ms step:813/1575 train_time:36454ms step_avg:44.84ms step:814/1575 train_time:36514ms step_avg:44.86ms step:815/1575 train_time:36577ms step_avg:44.88ms step:816/1575 train_time:36636ms step_avg:44.90ms step:817/1575 train_time:36699ms step_avg:44.92ms step:818/1575 train_time:36759ms step_avg:44.94ms step:819/1575 train_time:36822ms step_avg:44.96ms step:820/1575 train_time:36881ms step_avg:44.98ms step:821/1575 train_time:36945ms step_avg:45.00ms step:822/1575 train_time:37004ms step_avg:45.02ms step:823/1575 train_time:37068ms step_avg:45.04ms step:824/1575 train_time:37127ms step_avg:45.06ms step:825/1575 train_time:37190ms step_avg:45.08ms step:826/1575 train_time:37250ms step_avg:45.10ms step:827/1575 train_time:37313ms step_avg:45.12ms step:828/1575 train_time:37372ms step_avg:45.14ms step:829/1575 train_time:37436ms step_avg:45.16ms step:830/1575 train_time:37495ms step_avg:45.17ms step:831/1575 train_time:37558ms step_avg:45.20ms step:832/1575 train_time:37617ms step_avg:45.21ms step:833/1575 train_time:37680ms step_avg:45.23ms step:834/1575 train_time:37739ms step_avg:45.25ms step:835/1575 train_time:37802ms step_avg:45.27ms step:836/1575 train_time:37863ms step_avg:45.29ms step:837/1575 train_time:37925ms step_avg:45.31ms step:838/1575 train_time:37984ms step_avg:45.33ms step:839/1575 train_time:38047ms step_avg:45.35ms step:840/1575 train_time:38107ms step_avg:45.37ms step:841/1575 train_time:38171ms step_avg:45.39ms step:842/1575 train_time:38229ms step_avg:45.40ms step:843/1575 train_time:38292ms step_avg:45.42ms step:844/1575 train_time:38354ms step_avg:45.44ms step:845/1575 train_time:38417ms step_avg:45.46ms step:846/1575 train_time:38475ms step_avg:45.48ms step:847/1575 train_time:38538ms step_avg:45.50ms step:848/1575 train_time:38597ms step_avg:45.52ms step:849/1575 train_time:38660ms step_avg:45.54ms step:850/1575 train_time:38720ms step_avg:45.55ms step:851/1575 train_time:38783ms step_avg:45.57ms step:852/1575 train_time:38843ms step_avg:45.59ms step:853/1575 train_time:38907ms step_avg:45.61ms step:854/1575 train_time:38965ms step_avg:45.63ms step:855/1575 train_time:39030ms step_avg:45.65ms step:856/1575 train_time:39089ms step_avg:45.66ms step:857/1575 train_time:39152ms step_avg:45.68ms step:858/1575 train_time:39211ms step_avg:45.70ms step:859/1575 train_time:39274ms step_avg:45.72ms step:860/1575 train_time:39334ms step_avg:45.74ms step:861/1575 train_time:39397ms step_avg:45.76ms step:862/1575 train_time:39457ms step_avg:45.77ms step:863/1575 train_time:39519ms step_avg:45.79ms step:864/1575 train_time:39579ms step_avg:45.81ms step:865/1575 train_time:39642ms step_avg:45.83ms step:866/1575 train_time:39701ms step_avg:45.84ms step:867/1575 train_time:39764ms step_avg:45.86ms step:868/1575 train_time:39824ms step_avg:45.88ms step:869/1575 train_time:39887ms step_avg:45.90ms step:870/1575 train_time:39946ms step_avg:45.91ms step:871/1575 train_time:40009ms step_avg:45.93ms step:872/1575 train_time:40068ms step_avg:45.95ms step:873/1575 train_time:40131ms step_avg:45.97ms step:874/1575 train_time:40191ms step_avg:45.98ms step:875/1575 train_time:40254ms step_avg:46.00ms step:876/1575 train_time:40313ms step_avg:46.02ms step:877/1575 train_time:40377ms step_avg:46.04ms step:878/1575 train_time:40436ms step_avg:46.06ms step:879/1575 train_time:40504ms step_avg:46.08ms step:880/1575 train_time:40559ms step_avg:46.09ms step:881/1575 train_time:40621ms step_avg:46.11ms step:882/1575 train_time:40680ms step_avg:46.12ms step:883/1575 train_time:40744ms step_avg:46.14ms step:884/1575 train_time:40803ms step_avg:46.16ms step:885/1575 train_time:40866ms step_avg:46.18ms step:886/1575 train_time:40931ms step_avg:46.20ms step:887/1575 train_time:40991ms step_avg:46.21ms step:888/1575 train_time:41050ms step_avg:46.23ms step:889/1575 train_time:41112ms step_avg:46.25ms step:890/1575 train_time:41171ms step_avg:46.26ms step:891/1575 train_time:41234ms step_avg:46.28ms step:892/1575 train_time:41294ms step_avg:46.29ms step:893/1575 train_time:41357ms step_avg:46.31ms step:894/1575 train_time:41416ms step_avg:46.33ms step:895/1575 train_time:41480ms step_avg:46.35ms step:896/1575 train_time:41540ms step_avg:46.36ms step:897/1575 train_time:41602ms step_avg:46.38ms step:898/1575 train_time:41661ms step_avg:46.39ms step:899/1575 train_time:41724ms step_avg:46.41ms step:900/1575 train_time:41783ms step_avg:46.43ms step:901/1575 train_time:41846ms step_avg:46.44ms step:902/1575 train_time:41906ms step_avg:46.46ms step:903/1575 train_time:41969ms step_avg:46.48ms step:904/1575 train_time:42029ms step_avg:46.49ms step:905/1575 train_time:42092ms step_avg:46.51ms step:906/1575 train_time:42151ms step_avg:46.52ms step:907/1575 train_time:42215ms step_avg:46.54ms step:908/1575 train_time:42275ms step_avg:46.56ms step:909/1575 train_time:42338ms step_avg:46.58ms step:910/1575 train_time:42399ms step_avg:46.59ms step:911/1575 train_time:42461ms step_avg:46.61ms step:912/1575 train_time:42520ms step_avg:46.62ms step:913/1575 train_time:42584ms step_avg:46.64ms step:914/1575 train_time:42643ms step_avg:46.66ms step:915/1575 train_time:42706ms step_avg:46.67ms step:916/1575 train_time:42765ms step_avg:46.69ms step:917/1575 train_time:42828ms step_avg:46.70ms step:918/1575 train_time:42887ms step_avg:46.72ms step:919/1575 train_time:42950ms step_avg:46.74ms step:920/1575 train_time:43010ms step_avg:46.75ms step:921/1575 train_time:43073ms step_avg:46.77ms step:922/1575 train_time:43132ms step_avg:46.78ms step:923/1575 train_time:43195ms step_avg:46.80ms step:924/1575 train_time:43255ms step_avg:46.81ms step:925/1575 train_time:43319ms step_avg:46.83ms step:926/1575 train_time:43378ms step_avg:46.84ms step:927/1575 train_time:43442ms step_avg:46.86ms step:928/1575 train_time:43502ms step_avg:46.88ms step:929/1575 train_time:43564ms step_avg:46.89ms step:930/1575 train_time:43624ms step_avg:46.91ms step:931/1575 train_time:43687ms step_avg:46.92ms step:932/1575 train_time:43746ms step_avg:46.94ms step:933/1575 train_time:43809ms step_avg:46.96ms step:934/1575 train_time:43869ms step_avg:46.97ms step:935/1575 train_time:43932ms step_avg:46.99ms step:936/1575 train_time:43991ms step_avg:47.00ms step:937/1575 train_time:44054ms step_avg:47.02ms step:938/1575 train_time:44114ms step_avg:47.03ms step:939/1575 train_time:44177ms step_avg:47.05ms step:940/1575 train_time:44236ms step_avg:47.06ms step:941/1575 train_time:44299ms step_avg:47.08ms step:942/1575 train_time:44359ms step_avg:47.09ms step:943/1575 train_time:44422ms step_avg:47.11ms step:944/1575 train_time:44481ms step_avg:47.12ms step:945/1575 train_time:44545ms step_avg:47.14ms step:946/1575 train_time:44604ms step_avg:47.15ms step:947/1575 train_time:44667ms step_avg:47.17ms step:948/1575 train_time:44728ms step_avg:47.18ms step:949/1575 train_time:44791ms step_avg:47.20ms step:950/1575 train_time:44849ms step_avg:47.21ms step:951/1575 train_time:44912ms step_avg:47.23ms step:952/1575 train_time:44971ms step_avg:47.24ms step:953/1575 train_time:45037ms step_avg:47.26ms step:954/1575 train_time:45094ms step_avg:47.27ms step:955/1575 train_time:45158ms step_avg:47.29ms step:956/1575 train_time:45217ms step_avg:47.30ms step:957/1575 train_time:45280ms step_avg:47.32ms step:958/1575 train_time:45339ms step_avg:47.33ms step:959/1575 train_time:45403ms step_avg:47.34ms step:960/1575 train_time:45462ms step_avg:47.36ms step:961/1575 train_time:45525ms step_avg:47.37ms step:962/1575 train_time:45584ms step_avg:47.38ms step:963/1575 train_time:45647ms step_avg:47.40ms step:964/1575 train_time:45706ms step_avg:47.41ms step:965/1575 train_time:45769ms step_avg:47.43ms step:966/1575 train_time:45829ms step_avg:47.44ms step:967/1575 train_time:45893ms step_avg:47.46ms step:968/1575 train_time:45952ms step_avg:47.47ms step:969/1575 train_time:46015ms step_avg:47.49ms step:970/1575 train_time:46074ms step_avg:47.50ms step:971/1575 train_time:46138ms step_avg:47.52ms step:972/1575 train_time:46197ms step_avg:47.53ms step:973/1575 train_time:46261ms step_avg:47.54ms step:974/1575 train_time:46320ms step_avg:47.56ms step:975/1575 train_time:46383ms step_avg:47.57ms step:976/1575 train_time:46443ms step_avg:47.58ms step:977/1575 train_time:46505ms step_avg:47.60ms step:978/1575 train_time:46564ms step_avg:47.61ms step:979/1575 train_time:46627ms step_avg:47.63ms step:980/1575 train_time:46687ms step_avg:47.64ms step:981/1575 train_time:46750ms step_avg:47.66ms step:982/1575 train_time:46810ms step_avg:47.67ms step:983/1575 train_time:46873ms step_avg:47.68ms step:984/1575 train_time:46932ms step_avg:47.70ms step:985/1575 train_time:46997ms step_avg:47.71ms step:986/1575 train_time:47056ms step_avg:47.72ms step:987/1575 train_time:47120ms step_avg:47.74ms step:988/1575 train_time:47179ms step_avg:47.75ms step:989/1575 train_time:47242ms step_avg:47.77ms step:990/1575 train_time:47302ms step_avg:47.78ms step:991/1575 train_time:47364ms step_avg:47.79ms step:992/1575 train_time:47424ms step_avg:47.81ms step:993/1575 train_time:47488ms step_avg:47.82ms step:994/1575 train_time:47548ms step_avg:47.83ms step:995/1575 train_time:47610ms step_avg:47.85ms step:996/1575 train_time:47670ms step_avg:47.86ms step:997/1575 train_time:47733ms step_avg:47.88ms step:998/1575 train_time:47792ms step_avg:47.89ms step:999/1575 train_time:47855ms step_avg:47.90ms step:1000/1575 train_time:47914ms step_avg:47.91ms step:1000/1575 val_loss:3.5839 train_time:47961ms step_avg:47.96ms step:1001/1575 train_time:47982ms step_avg:47.93ms step:1002/1575 train_time:48043ms step_avg:47.95ms step:1003/1575 train_time:48108ms step_avg:47.96ms step:1004/1575 train_time:48174ms step_avg:47.98ms step:1005/1575 train_time:48236ms step_avg:48.00ms step:1006/1575 train_time:48295ms step_avg:48.01ms step:1007/1575 train_time:48358ms step_avg:48.02ms step:1008/1575 train_time:48417ms step_avg:48.03ms step:1009/1575 train_time:48480ms step_avg:48.05ms step:1010/1575 train_time:48539ms step_avg:48.06ms step:1011/1575 train_time:48602ms step_avg:48.07ms step:1012/1575 train_time:48662ms step_avg:48.08ms step:1013/1575 train_time:48726ms step_avg:48.10ms step:1014/1575 train_time:48784ms step_avg:48.11ms step:1015/1575 train_time:48847ms step_avg:48.13ms step:1016/1575 train_time:48906ms step_avg:48.14ms step:1017/1575 train_time:48969ms step_avg:48.15ms step:1018/1575 train_time:49029ms step_avg:48.16ms step:1019/1575 train_time:49093ms step_avg:48.18ms step:1020/1575 train_time:49153ms step_avg:48.19ms step:1021/1575 train_time:49217ms step_avg:48.21ms step:1022/1575 train_time:49277ms step_avg:48.22ms step:1023/1575 train_time:49341ms step_avg:48.23ms step:1024/1575 train_time:49400ms step_avg:48.24ms step:1025/1575 train_time:49475ms step_avg:48.27ms step:1026/1575 train_time:49556ms step_avg:48.30ms step:1027/1575 train_time:49645ms step_avg:48.34ms step:1028/1575 train_time:49730ms step_avg:48.38ms step:1029/1575 train_time:49819ms step_avg:48.41ms step:1030/1575 train_time:49904ms step_avg:48.45ms step:1031/1575 train_time:49994ms step_avg:48.49ms step:1032/1575 train_time:50080ms step_avg:48.53ms step:1033/1575 train_time:50170ms step_avg:48.57ms step:1034/1575 train_time:50256ms step_avg:48.60ms step:1035/1575 train_time:50346ms step_avg:48.64ms step:1036/1575 train_time:50431ms step_avg:48.68ms step:1037/1575 train_time:50520ms step_avg:48.72ms step:1038/1575 train_time:50606ms step_avg:48.75ms step:1039/1575 train_time:50694ms step_avg:48.79ms step:1040/1575 train_time:50780ms step_avg:48.83ms step:1041/1575 train_time:50869ms step_avg:48.87ms step:1042/1575 train_time:50955ms step_avg:48.90ms step:1043/1575 train_time:51046ms step_avg:48.94ms step:1044/1575 train_time:51132ms step_avg:48.98ms step:1045/1575 train_time:51223ms step_avg:49.02ms step:1046/1575 train_time:51308ms step_avg:49.05ms step:1047/1575 train_time:51397ms step_avg:49.09ms step:1048/1575 train_time:51483ms step_avg:49.13ms step:1049/1575 train_time:51572ms step_avg:49.16ms step:1050/1575 train_time:51657ms step_avg:49.20ms step:1051/1575 train_time:51746ms step_avg:49.24ms step:1052/1575 train_time:51833ms step_avg:49.27ms step:1053/1575 train_time:51923ms step_avg:49.31ms step:1054/1575 train_time:52009ms step_avg:49.34ms step:1055/1575 train_time:52098ms step_avg:49.38ms step:1056/1575 train_time:52185ms step_avg:49.42ms step:1057/1575 train_time:52273ms step_avg:49.45ms step:1058/1575 train_time:52360ms step_avg:49.49ms step:1059/1575 train_time:52449ms step_avg:49.53ms step:1060/1575 train_time:52535ms step_avg:49.56ms step:1061/1575 train_time:52624ms step_avg:49.60ms step:1062/1575 train_time:52710ms step_avg:49.63ms step:1063/1575 train_time:52799ms step_avg:49.67ms step:1064/1575 train_time:52885ms step_avg:49.70ms step:1065/1575 train_time:52973ms step_avg:49.74ms step:1066/1575 train_time:53058ms step_avg:49.77ms step:1067/1575 train_time:53148ms step_avg:49.81ms step:1068/1575 train_time:53234ms step_avg:49.84ms step:1069/1575 train_time:53325ms step_avg:49.88ms step:1070/1575 train_time:53411ms step_avg:49.92ms step:1071/1575 train_time:53500ms step_avg:49.95ms step:1072/1575 train_time:53586ms step_avg:49.99ms step:1073/1575 train_time:53675ms step_avg:50.02ms step:1074/1575 train_time:53760ms step_avg:50.06ms step:1075/1575 train_time:53850ms step_avg:50.09ms step:1076/1575 train_time:53936ms step_avg:50.13ms step:1077/1575 train_time:54025ms step_avg:50.16ms step:1078/1575 train_time:54112ms step_avg:50.20ms step:1079/1575 train_time:54201ms step_avg:50.23ms step:1080/1575 train_time:54287ms step_avg:50.27ms step:1081/1575 train_time:54378ms step_avg:50.30ms step:1082/1575 train_time:54463ms step_avg:50.34ms step:1083/1575 train_time:54552ms step_avg:50.37ms step:1084/1575 train_time:54638ms step_avg:50.40ms step:1085/1575 train_time:54727ms step_avg:50.44ms step:1086/1575 train_time:54813ms step_avg:50.47ms step:1087/1575 train_time:54902ms step_avg:50.51ms step:1088/1575 train_time:54988ms step_avg:50.54ms step:1089/1575 train_time:55078ms step_avg:50.58ms step:1090/1575 train_time:55165ms step_avg:50.61ms step:1091/1575 train_time:55255ms step_avg:50.65ms step:1092/1575 train_time:55340ms step_avg:50.68ms step:1093/1575 train_time:55430ms step_avg:50.71ms step:1094/1575 train_time:55517ms step_avg:50.75ms step:1095/1575 train_time:55605ms step_avg:50.78ms step:1096/1575 train_time:55690ms step_avg:50.81ms step:1097/1575 train_time:55778ms step_avg:50.85ms step:1098/1575 train_time:55864ms step_avg:50.88ms step:1099/1575 train_time:55953ms step_avg:50.91ms step:1100/1575 train_time:56038ms step_avg:50.94ms step:1101/1575 train_time:56128ms step_avg:50.98ms step:1102/1575 train_time:56214ms step_avg:51.01ms step:1103/1575 train_time:56304ms step_avg:51.05ms step:1104/1575 train_time:56389ms step_avg:51.08ms step:1105/1575 train_time:56480ms step_avg:51.11ms step:1106/1575 train_time:56566ms step_avg:51.14ms step:1107/1575 train_time:56656ms step_avg:51.18ms step:1108/1575 train_time:56741ms step_avg:51.21ms step:1109/1575 train_time:56833ms step_avg:51.25ms step:1110/1575 train_time:56920ms step_avg:51.28ms step:1111/1575 train_time:57006ms step_avg:51.31ms step:1112/1575 train_time:57091ms step_avg:51.34ms step:1113/1575 train_time:57182ms step_avg:51.38ms step:1114/1575 train_time:57266ms step_avg:51.41ms step:1115/1575 train_time:57355ms step_avg:51.44ms step:1116/1575 train_time:57441ms step_avg:51.47ms step:1117/1575 train_time:57530ms step_avg:51.50ms step:1118/1575 train_time:57617ms step_avg:51.54ms step:1119/1575 train_time:57707ms step_avg:51.57ms step:1120/1575 train_time:57792ms step_avg:51.60ms step:1121/1575 train_time:57881ms step_avg:51.63ms step:1122/1575 train_time:57967ms step_avg:51.66ms step:1123/1575 train_time:58056ms step_avg:51.70ms step:1124/1575 train_time:58142ms step_avg:51.73ms step:1125/1575 train_time:58231ms step_avg:51.76ms step:1126/1575 train_time:58318ms step_avg:51.79ms step:1127/1575 train_time:58406ms step_avg:51.82ms step:1128/1575 train_time:58493ms step_avg:51.86ms step:1129/1575 train_time:58582ms step_avg:51.89ms step:1130/1575 train_time:58667ms step_avg:51.92ms step:1131/1575 train_time:58758ms step_avg:51.95ms step:1132/1575 train_time:58845ms step_avg:51.98ms step:1133/1575 train_time:58934ms step_avg:52.02ms step:1134/1575 train_time:59020ms step_avg:52.05ms step:1135/1575 train_time:59109ms step_avg:52.08ms step:1136/1575 train_time:59195ms step_avg:52.11ms step:1137/1575 train_time:59284ms step_avg:52.14ms step:1138/1575 train_time:59369ms step_avg:52.17ms step:1139/1575 train_time:59458ms step_avg:52.20ms step:1140/1575 train_time:59545ms step_avg:52.23ms step:1141/1575 train_time:59633ms step_avg:52.26ms step:1142/1575 train_time:59719ms step_avg:52.29ms step:1143/1575 train_time:59809ms step_avg:52.33ms step:1144/1575 train_time:59895ms step_avg:52.36ms step:1145/1575 train_time:59985ms step_avg:52.39ms step:1146/1575 train_time:60075ms step_avg:52.42ms step:1147/1575 train_time:60162ms step_avg:52.45ms step:1148/1575 train_time:60248ms step_avg:52.48ms step:1149/1575 train_time:60338ms step_avg:52.51ms step:1150/1575 train_time:60423ms step_avg:52.54ms step:1151/1575 train_time:60513ms step_avg:52.57ms step:1152/1575 train_time:60598ms step_avg:52.60ms step:1153/1575 train_time:60688ms step_avg:52.63ms step:1154/1575 train_time:60773ms step_avg:52.66ms step:1155/1575 train_time:60863ms step_avg:52.70ms step:1156/1575 train_time:60949ms step_avg:52.72ms step:1157/1575 train_time:61039ms step_avg:52.76ms step:1158/1575 train_time:61125ms step_avg:52.78ms step:1159/1575 train_time:61214ms step_avg:52.82ms step:1160/1575 train_time:61299ms step_avg:52.84ms step:1161/1575 train_time:61389ms step_avg:52.88ms step:1162/1575 train_time:61476ms step_avg:52.91ms step:1163/1575 train_time:61565ms step_avg:52.94ms step:1164/1575 train_time:61650ms step_avg:52.96ms step:1165/1575 train_time:61740ms step_avg:53.00ms step:1166/1575 train_time:61825ms step_avg:53.02ms step:1167/1575 train_time:61915ms step_avg:53.06ms step:1168/1575 train_time:62001ms step_avg:53.08ms step:1169/1575 train_time:62091ms step_avg:53.11ms step:1170/1575 train_time:62176ms step_avg:53.14ms step:1171/1575 train_time:62265ms step_avg:53.17ms step:1172/1575 train_time:62353ms step_avg:53.20ms step:1173/1575 train_time:62442ms step_avg:53.23ms step:1174/1575 train_time:62528ms step_avg:53.26ms step:1175/1575 train_time:62618ms step_avg:53.29ms step:1176/1575 train_time:62704ms step_avg:53.32ms step:1177/1575 train_time:62793ms step_avg:53.35ms step:1178/1575 train_time:62878ms step_avg:53.38ms step:1179/1575 train_time:62968ms step_avg:53.41ms step:1180/1575 train_time:63054ms step_avg:53.44ms step:1181/1575 train_time:63145ms step_avg:53.47ms step:1182/1575 train_time:63229ms step_avg:53.49ms step:1183/1575 train_time:63319ms step_avg:53.52ms step:1184/1575 train_time:63406ms step_avg:53.55ms step:1185/1575 train_time:63495ms step_avg:53.58ms step:1186/1575 train_time:63581ms step_avg:53.61ms step:1187/1575 train_time:63671ms step_avg:53.64ms step:1188/1575 train_time:63756ms step_avg:53.67ms step:1189/1575 train_time:63845ms step_avg:53.70ms step:1190/1575 train_time:63931ms step_avg:53.72ms step:1191/1575 train_time:64022ms step_avg:53.75ms step:1192/1575 train_time:64107ms step_avg:53.78ms step:1193/1575 train_time:64197ms step_avg:53.81ms step:1194/1575 train_time:64282ms step_avg:53.84ms step:1195/1575 train_time:64371ms step_avg:53.87ms step:1196/1575 train_time:64456ms step_avg:53.89ms step:1197/1575 train_time:64545ms step_avg:53.92ms step:1198/1575 train_time:64632ms step_avg:53.95ms step:1199/1575 train_time:64721ms step_avg:53.98ms step:1200/1575 train_time:64807ms step_avg:54.01ms step:1201/1575 train_time:64897ms step_avg:54.04ms step:1202/1575 train_time:64982ms step_avg:54.06ms step:1203/1575 train_time:65072ms step_avg:54.09ms step:1204/1575 train_time:65158ms step_avg:54.12ms step:1205/1575 train_time:65248ms step_avg:54.15ms step:1206/1575 train_time:65334ms step_avg:54.17ms step:1207/1575 train_time:65424ms step_avg:54.20ms step:1208/1575 train_time:65510ms step_avg:54.23ms step:1209/1575 train_time:65601ms step_avg:54.26ms step:1210/1575 train_time:65687ms step_avg:54.29ms step:1211/1575 train_time:65776ms step_avg:54.32ms step:1212/1575 train_time:65861ms step_avg:54.34ms step:1213/1575 train_time:65950ms step_avg:54.37ms step:1214/1575 train_time:66037ms step_avg:54.40ms step:1215/1575 train_time:66126ms step_avg:54.42ms step:1216/1575 train_time:66211ms step_avg:54.45ms step:1217/1575 train_time:66302ms step_avg:54.48ms step:1218/1575 train_time:66388ms step_avg:54.51ms step:1219/1575 train_time:66478ms step_avg:54.53ms step:1220/1575 train_time:66564ms step_avg:54.56ms step:1221/1575 train_time:66654ms step_avg:54.59ms step:1222/1575 train_time:66740ms step_avg:54.62ms step:1223/1575 train_time:66828ms step_avg:54.64ms step:1224/1575 train_time:66914ms step_avg:54.67ms step:1225/1575 train_time:67004ms step_avg:54.70ms step:1226/1575 train_time:67089ms step_avg:54.72ms step:1227/1575 train_time:67179ms step_avg:54.75ms step:1228/1575 train_time:67265ms step_avg:54.78ms step:1229/1575 train_time:67354ms step_avg:54.80ms step:1230/1575 train_time:67439ms step_avg:54.83ms step:1231/1575 train_time:67529ms step_avg:54.86ms step:1232/1575 train_time:67615ms step_avg:54.88ms step:1233/1575 train_time:67705ms step_avg:54.91ms step:1234/1575 train_time:67791ms step_avg:54.94ms step:1235/1575 train_time:67881ms step_avg:54.96ms step:1236/1575 train_time:67966ms step_avg:54.99ms step:1237/1575 train_time:68055ms step_avg:55.02ms step:1238/1575 train_time:68141ms step_avg:55.04ms step:1239/1575 train_time:68232ms step_avg:55.07ms step:1240/1575 train_time:68317ms step_avg:55.09ms step:1241/1575 train_time:68406ms step_avg:55.12ms step:1242/1575 train_time:68492ms step_avg:55.15ms step:1243/1575 train_time:68582ms step_avg:55.17ms step:1244/1575 train_time:68667ms step_avg:55.20ms step:1245/1575 train_time:68758ms step_avg:55.23ms step:1246/1575 train_time:68843ms step_avg:55.25ms step:1247/1575 train_time:68932ms step_avg:55.28ms step:1248/1575 train_time:69019ms step_avg:55.30ms step:1249/1575 train_time:69108ms step_avg:55.33ms step:1250/1575 train_time:69194ms step_avg:55.36ms step:1250/1575 val_loss:3.4077 train_time:69266ms step_avg:55.41ms step:1251/1575 train_time:69287ms step_avg:55.39ms step:1252/1575 train_time:69376ms step_avg:55.41ms step:1253/1575 train_time:69470ms step_avg:55.44ms step:1254/1575 train_time:69556ms step_avg:55.47ms step:1255/1575 train_time:69645ms step_avg:55.49ms step:1256/1575 train_time:69731ms step_avg:55.52ms step:1257/1575 train_time:69821ms step_avg:55.55ms step:1258/1575 train_time:69905ms step_avg:55.57ms step:1259/1575 train_time:69994ms step_avg:55.60ms step:1260/1575 train_time:70080ms step_avg:55.62ms step:1261/1575 train_time:70168ms step_avg:55.64ms step:1262/1575 train_time:70255ms step_avg:55.67ms step:1263/1575 train_time:70349ms step_avg:55.70ms step:1264/1575 train_time:70438ms step_avg:55.73ms step:1265/1575 train_time:70528ms step_avg:55.75ms step:1266/1575 train_time:70614ms step_avg:55.78ms step:1267/1575 train_time:70704ms step_avg:55.80ms step:1268/1575 train_time:70789ms step_avg:55.83ms step:1269/1575 train_time:70878ms step_avg:55.85ms step:1270/1575 train_time:70963ms step_avg:55.88ms step:1271/1575 train_time:71052ms step_avg:55.90ms step:1272/1575 train_time:71138ms step_avg:55.93ms step:1273/1575 train_time:71227ms step_avg:55.95ms step:1274/1575 train_time:71314ms step_avg:55.98ms step:1275/1575 train_time:71405ms step_avg:56.00ms step:1276/1575 train_time:71492ms step_avg:56.03ms step:1277/1575 train_time:71583ms step_avg:56.06ms step:1278/1575 train_time:71670ms step_avg:56.08ms step:1279/1575 train_time:71759ms step_avg:56.11ms step:1280/1575 train_time:71844ms step_avg:56.13ms step:1281/1575 train_time:71933ms step_avg:56.15ms step:1282/1575 train_time:72019ms step_avg:56.18ms step:1283/1575 train_time:72107ms step_avg:56.20ms step:1284/1575 train_time:72193ms step_avg:56.22ms step:1285/1575 train_time:72283ms step_avg:56.25ms step:1286/1575 train_time:72370ms step_avg:56.28ms step:1287/1575 train_time:72462ms step_avg:56.30ms step:1288/1575 train_time:72548ms step_avg:56.33ms step:1289/1575 train_time:72637ms step_avg:56.35ms step:1290/1575 train_time:72723ms step_avg:56.37ms step:1291/1575 train_time:72813ms step_avg:56.40ms step:1292/1575 train_time:72899ms step_avg:56.42ms step:1293/1575 train_time:72989ms step_avg:56.45ms step:1294/1575 train_time:73073ms step_avg:56.47ms step:1295/1575 train_time:73162ms step_avg:56.50ms step:1296/1575 train_time:73249ms step_avg:56.52ms step:1297/1575 train_time:73339ms step_avg:56.55ms step:1298/1575 train_time:73425ms step_avg:56.57ms step:1299/1575 train_time:73515ms step_avg:56.59ms step:1300/1575 train_time:73601ms step_avg:56.62ms step:1301/1575 train_time:73690ms step_avg:56.64ms step:1302/1575 train_time:73778ms step_avg:56.66ms step:1303/1575 train_time:73867ms step_avg:56.69ms step:1304/1575 train_time:73952ms step_avg:56.71ms step:1305/1575 train_time:74041ms step_avg:56.74ms step:1306/1575 train_time:74127ms step_avg:56.76ms step:1307/1575 train_time:74216ms step_avg:56.78ms step:1308/1575 train_time:74302ms step_avg:56.81ms step:1309/1575 train_time:74391ms step_avg:56.83ms step:1310/1575 train_time:74477ms step_avg:56.85ms step:1311/1575 train_time:74567ms step_avg:56.88ms step:1312/1575 train_time:74653ms step_avg:56.90ms step:1313/1575 train_time:74744ms step_avg:56.93ms step:1314/1575 train_time:74831ms step_avg:56.95ms step:1315/1575 train_time:74919ms step_avg:56.97ms step:1316/1575 train_time:75004ms step_avg:56.99ms step:1317/1575 train_time:75094ms step_avg:57.02ms step:1318/1575 train_time:75180ms step_avg:57.04ms step:1319/1575 train_time:75269ms step_avg:57.07ms step:1320/1575 train_time:75355ms step_avg:57.09ms step:1321/1575 train_time:75445ms step_avg:57.11ms step:1322/1575 train_time:75532ms step_avg:57.13ms step:1323/1575 train_time:75622ms step_avg:57.16ms step:1324/1575 train_time:75708ms step_avg:57.18ms step:1325/1575 train_time:75798ms step_avg:57.21ms step:1326/1575 train_time:75883ms step_avg:57.23ms step:1327/1575 train_time:75972ms step_avg:57.25ms step:1328/1575 train_time:76058ms step_avg:57.27ms step:1329/1575 train_time:76147ms step_avg:57.30ms step:1330/1575 train_time:76233ms step_avg:57.32ms step:1331/1575 train_time:76324ms step_avg:57.34ms step:1332/1575 train_time:76410ms step_avg:57.36ms step:1333/1575 train_time:76499ms step_avg:57.39ms step:1334/1575 train_time:76585ms step_avg:57.41ms step:1335/1575 train_time:76675ms step_avg:57.43ms step:1336/1575 train_time:76760ms step_avg:57.46ms step:1337/1575 train_time:76849ms step_avg:57.48ms step:1338/1575 train_time:76936ms step_avg:57.50ms step:1339/1575 train_time:77025ms step_avg:57.52ms step:1340/1575 train_time:77110ms step_avg:57.54ms step:1341/1575 train_time:77200ms step_avg:57.57ms step:1342/1575 train_time:77286ms step_avg:57.59ms step:1343/1575 train_time:77376ms step_avg:57.61ms step:1344/1575 train_time:77461ms step_avg:57.63ms step:1345/1575 train_time:77552ms step_avg:57.66ms step:1346/1575 train_time:77639ms step_avg:57.68ms step:1347/1575 train_time:77727ms step_avg:57.70ms step:1348/1575 train_time:77816ms step_avg:57.73ms step:1349/1575 train_time:77906ms step_avg:57.75ms step:1350/1575 train_time:77989ms step_avg:57.77ms step:1351/1575 train_time:78078ms step_avg:57.79ms step:1352/1575 train_time:78163ms step_avg:57.81ms step:1353/1575 train_time:78254ms step_avg:57.84ms step:1354/1575 train_time:78341ms step_avg:57.86ms step:1355/1575 train_time:78430ms step_avg:57.88ms step:1356/1575 train_time:78516ms step_avg:57.90ms step:1357/1575 train_time:78606ms step_avg:57.93ms step:1358/1575 train_time:78696ms step_avg:57.95ms step:1359/1575 train_time:78784ms step_avg:57.97ms step:1360/1575 train_time:78869ms step_avg:57.99ms step:1361/1575 train_time:78958ms step_avg:58.01ms step:1362/1575 train_time:79043ms step_avg:58.03ms step:1363/1575 train_time:79133ms step_avg:58.06ms step:1364/1575 train_time:79218ms step_avg:58.08ms step:1365/1575 train_time:79308ms step_avg:58.10ms step:1366/1575 train_time:79395ms step_avg:58.12ms step:1367/1575 train_time:79484ms step_avg:58.14ms step:1368/1575 train_time:79570ms step_avg:58.17ms step:1369/1575 train_time:79659ms step_avg:58.19ms step:1370/1575 train_time:79745ms step_avg:58.21ms step:1371/1575 train_time:79836ms step_avg:58.23ms step:1372/1575 train_time:79920ms step_avg:58.25ms step:1373/1575 train_time:80009ms step_avg:58.27ms step:1374/1575 train_time:80094ms step_avg:58.29ms step:1375/1575 train_time:80184ms step_avg:58.32ms step:1376/1575 train_time:80269ms step_avg:58.34ms step:1377/1575 train_time:80360ms step_avg:58.36ms step:1378/1575 train_time:80446ms step_avg:58.38ms step:1379/1575 train_time:80536ms step_avg:58.40ms step:1380/1575 train_time:80621ms step_avg:58.42ms step:1381/1575 train_time:80711ms step_avg:58.44ms step:1382/1575 train_time:80798ms step_avg:58.46ms step:1383/1575 train_time:80887ms step_avg:58.49ms step:1384/1575 train_time:80973ms step_avg:58.51ms step:1385/1575 train_time:81063ms step_avg:58.53ms step:1386/1575 train_time:81149ms step_avg:58.55ms step:1387/1575 train_time:81239ms step_avg:58.57ms step:1388/1575 train_time:81324ms step_avg:58.59ms step:1389/1575 train_time:81415ms step_avg:58.61ms step:1390/1575 train_time:81500ms step_avg:58.63ms step:1391/1575 train_time:81589ms step_avg:58.65ms step:1392/1575 train_time:81675ms step_avg:58.67ms step:1393/1575 train_time:81765ms step_avg:58.70ms step:1394/1575 train_time:81851ms step_avg:58.72ms step:1395/1575 train_time:81940ms step_avg:58.74ms step:1396/1575 train_time:82026ms step_avg:58.76ms step:1397/1575 train_time:82115ms step_avg:58.78ms step:1398/1575 train_time:82202ms step_avg:58.80ms step:1399/1575 train_time:82290ms step_avg:58.82ms step:1400/1575 train_time:82377ms step_avg:58.84ms step:1401/1575 train_time:82466ms step_avg:58.86ms step:1402/1575 train_time:82552ms step_avg:58.88ms step:1403/1575 train_time:82643ms step_avg:58.90ms step:1404/1575 train_time:82729ms step_avg:58.92ms step:1405/1575 train_time:82818ms step_avg:58.95ms step:1406/1575 train_time:82903ms step_avg:58.96ms step:1407/1575 train_time:82992ms step_avg:58.99ms step:1408/1575 train_time:83078ms step_avg:59.00ms step:1409/1575 train_time:83167ms step_avg:59.03ms step:1410/1575 train_time:83254ms step_avg:59.05ms step:1411/1575 train_time:83344ms step_avg:59.07ms step:1412/1575 train_time:83429ms step_avg:59.09ms step:1413/1575 train_time:83518ms step_avg:59.11ms step:1414/1575 train_time:83604ms step_avg:59.13ms step:1415/1575 train_time:83694ms step_avg:59.15ms step:1416/1575 train_time:83780ms step_avg:59.17ms step:1417/1575 train_time:83869ms step_avg:59.19ms step:1418/1575 train_time:83956ms step_avg:59.21ms step:1419/1575 train_time:84046ms step_avg:59.23ms step:1420/1575 train_time:84131ms step_avg:59.25ms step:1421/1575 train_time:84222ms step_avg:59.27ms step:1422/1575 train_time:84307ms step_avg:59.29ms step:1423/1575 train_time:84397ms step_avg:59.31ms step:1424/1575 train_time:84483ms step_avg:59.33ms step:1425/1575 train_time:84572ms step_avg:59.35ms step:1426/1575 train_time:84658ms step_avg:59.37ms step:1427/1575 train_time:84747ms step_avg:59.39ms step:1428/1575 train_time:84832ms step_avg:59.41ms step:1429/1575 train_time:84923ms step_avg:59.43ms step:1430/1575 train_time:85010ms step_avg:59.45ms step:1431/1575 train_time:85099ms step_avg:59.47ms step:1432/1575 train_time:85184ms step_avg:59.49ms step:1433/1575 train_time:85274ms step_avg:59.51ms step:1434/1575 train_time:85360ms step_avg:59.53ms step:1435/1575 train_time:85450ms step_avg:59.55ms step:1436/1575 train_time:85536ms step_avg:59.57ms step:1437/1575 train_time:85626ms step_avg:59.59ms step:1438/1575 train_time:85712ms step_avg:59.61ms step:1439/1575 train_time:85802ms step_avg:59.63ms step:1440/1575 train_time:85888ms step_avg:59.64ms step:1441/1575 train_time:85978ms step_avg:59.67ms step:1442/1575 train_time:86064ms step_avg:59.68ms step:1443/1575 train_time:86153ms step_avg:59.70ms step:1444/1575 train_time:86239ms step_avg:59.72ms step:1445/1575 train_time:86328ms step_avg:59.74ms step:1446/1575 train_time:86414ms step_avg:59.76ms step:1447/1575 train_time:86503ms step_avg:59.78ms step:1448/1575 train_time:86590ms step_avg:59.80ms step:1449/1575 train_time:86680ms step_avg:59.82ms step:1450/1575 train_time:86765ms step_avg:59.84ms step:1451/1575 train_time:86855ms step_avg:59.86ms step:1452/1575 train_time:86940ms step_avg:59.88ms step:1453/1575 train_time:87029ms step_avg:59.90ms step:1454/1575 train_time:87116ms step_avg:59.91ms step:1455/1575 train_time:87206ms step_avg:59.94ms step:1456/1575 train_time:87293ms step_avg:59.95ms step:1457/1575 train_time:87384ms step_avg:59.97ms step:1458/1575 train_time:87469ms step_avg:59.99ms step:1459/1575 train_time:87558ms step_avg:60.01ms step:1460/1575 train_time:87644ms step_avg:60.03ms step:1461/1575 train_time:87733ms step_avg:60.05ms step:1462/1575 train_time:87820ms step_avg:60.07ms step:1463/1575 train_time:87909ms step_avg:60.09ms step:1464/1575 train_time:87995ms step_avg:60.11ms step:1465/1575 train_time:88084ms step_avg:60.13ms step:1466/1575 train_time:88169ms step_avg:60.14ms step:1467/1575 train_time:88259ms step_avg:60.16ms step:1468/1575 train_time:88345ms step_avg:60.18ms step:1469/1575 train_time:88434ms step_avg:60.20ms step:1470/1575 train_time:88520ms step_avg:60.22ms step:1471/1575 train_time:88609ms step_avg:60.24ms step:1472/1575 train_time:88695ms step_avg:60.25ms step:1473/1575 train_time:88785ms step_avg:60.27ms step:1474/1575 train_time:88870ms step_avg:60.29ms step:1475/1575 train_time:88960ms step_avg:60.31ms step:1476/1575 train_time:89046ms step_avg:60.33ms step:1477/1575 train_time:89136ms step_avg:60.35ms step:1478/1575 train_time:89222ms step_avg:60.37ms step:1479/1575 train_time:89311ms step_avg:60.39ms step:1480/1575 train_time:89398ms step_avg:60.40ms step:1481/1575 train_time:89486ms step_avg:60.42ms step:1482/1575 train_time:89572ms step_avg:60.44ms step:1483/1575 train_time:89661ms step_avg:60.46ms step:1484/1575 train_time:89747ms step_avg:60.48ms step:1485/1575 train_time:89836ms step_avg:60.50ms step:1486/1575 train_time:89922ms step_avg:60.51ms step:1487/1575 train_time:90012ms step_avg:60.53ms step:1488/1575 train_time:90098ms step_avg:60.55ms step:1489/1575 train_time:90188ms step_avg:60.57ms step:1490/1575 train_time:90274ms step_avg:60.59ms step:1491/1575 train_time:90363ms step_avg:60.61ms step:1492/1575 train_time:90449ms step_avg:60.62ms step:1493/1575 train_time:90540ms step_avg:60.64ms step:1494/1575 train_time:90626ms step_avg:60.66ms step:1495/1575 train_time:90715ms step_avg:60.68ms step:1496/1575 train_time:90802ms step_avg:60.70ms step:1497/1575 train_time:90891ms step_avg:60.72ms step:1498/1575 train_time:90977ms step_avg:60.73ms step:1499/1575 train_time:91067ms step_avg:60.75ms step:1500/1575 train_time:91153ms step_avg:60.77ms step:1500/1575 val_loss:3.3007 train_time:91226ms step_avg:60.82ms step:1501/1575 train_time:91247ms step_avg:60.79ms step:1502/1575 train_time:91334ms step_avg:60.81ms step:1503/1575 train_time:91427ms step_avg:60.83ms step:1504/1575 train_time:91513ms step_avg:60.85ms step:1505/1575 train_time:91603ms step_avg:60.87ms step:1506/1575 train_time:91688ms step_avg:60.88ms step:1507/1575 train_time:91776ms step_avg:60.90ms step:1508/1575 train_time:91860ms step_avg:60.92ms step:1509/1575 train_time:91950ms step_avg:60.93ms step:1510/1575 train_time:92037ms step_avg:60.95ms step:1511/1575 train_time:92125ms step_avg:60.97ms step:1512/1575 train_time:92210ms step_avg:60.99ms step:1513/1575 train_time:92302ms step_avg:61.01ms step:1514/1575 train_time:92389ms step_avg:61.02ms step:1515/1575 train_time:92481ms step_avg:61.04ms step:1516/1575 train_time:92567ms step_avg:61.06ms step:1517/1575 train_time:92656ms step_avg:61.08ms step:1518/1575 train_time:92742ms step_avg:61.09ms step:1519/1575 train_time:92829ms step_avg:61.11ms step:1520/1575 train_time:92916ms step_avg:61.13ms step:1521/1575 train_time:93004ms step_avg:61.15ms step:1522/1575 train_time:93090ms step_avg:61.16ms step:1523/1575 train_time:93180ms step_avg:61.18ms step:1524/1575 train_time:93267ms step_avg:61.20ms step:1525/1575 train_time:93359ms step_avg:61.22ms step:1526/1575 train_time:93445ms step_avg:61.24ms step:1527/1575 train_time:93535ms step_avg:61.25ms step:1528/1575 train_time:93621ms step_avg:61.27ms step:1529/1575 train_time:93710ms step_avg:61.29ms step:1530/1575 train_time:93795ms step_avg:61.30ms step:1531/1575 train_time:93884ms step_avg:61.32ms step:1532/1575 train_time:93970ms step_avg:61.34ms step:1533/1575 train_time:94059ms step_avg:61.36ms step:1534/1575 train_time:94144ms step_avg:61.37ms step:1535/1575 train_time:94235ms step_avg:61.39ms step:1536/1575 train_time:94330ms step_avg:61.41ms step:1537/1575 train_time:94418ms step_avg:61.43ms step:1538/1575 train_time:94505ms step_avg:61.45ms step:1539/1575 train_time:94595ms step_avg:61.47ms step:1540/1575 train_time:94681ms step_avg:61.48ms step:1541/1575 train_time:94770ms step_avg:61.50ms step:1542/1575 train_time:94855ms step_avg:61.51ms step:1543/1575 train_time:94944ms step_avg:61.53ms step:1544/1575 train_time:95030ms step_avg:61.55ms step:1545/1575 train_time:95119ms step_avg:61.57ms step:1546/1575 train_time:95205ms step_avg:61.58ms step:1547/1575 train_time:95296ms step_avg:61.60ms step:1548/1575 train_time:95383ms step_avg:61.62ms step:1549/1575 train_time:95474ms step_avg:61.64ms step:1550/1575 train_time:95559ms step_avg:61.65ms step:1551/1575 train_time:95649ms step_avg:61.67ms step:1552/1575 train_time:95735ms step_avg:61.69ms step:1553/1575 train_time:95825ms step_avg:61.70ms step:1554/1575 train_time:95911ms step_avg:61.72ms step:1555/1575 train_time:96001ms step_avg:61.74ms step:1556/1575 train_time:96087ms step_avg:61.75ms step:1557/1575 train_time:96176ms step_avg:61.77ms step:1558/1575 train_time:96262ms step_avg:61.79ms step:1559/1575 train_time:96352ms step_avg:61.80ms step:1560/1575 train_time:96438ms step_avg:61.82ms step:1561/1575 train_time:96528ms step_avg:61.84ms step:1562/1575 train_time:96615ms step_avg:61.85ms step:1563/1575 train_time:96705ms step_avg:61.87ms step:1564/1575 train_time:96791ms step_avg:61.89ms step:1565/1575 train_time:96881ms step_avg:61.90ms step:1566/1575 train_time:96969ms step_avg:61.92ms step:1567/1575 train_time:97060ms step_avg:61.94ms step:1568/1575 train_time:97145ms step_avg:61.95ms step:1569/1575 train_time:97241ms step_avg:61.98ms step:1570/1575 train_time:97325ms step_avg:61.99ms step:1571/1575 train_time:97413ms step_avg:62.01ms step:1572/1575 train_time:97499ms step_avg:62.02ms step:1573/1575 train_time:97588ms step_avg:62.04ms step:1574/1575 train_time:97675ms step_avg:62.05ms step:1575/1575 train_time:97763ms step_avg:62.07ms step:1575/1575 val_loss:3.2786 train_time:97831ms step_avg:62.11ms peak memory allocated: 30933 MiB reserved: 46918 MiB