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:26:52 2026 +-----------------------------------------------------------------------------------------+ | NVIDIA-SMI 570.148.08 Driver Version: 570.148.08 CUDA Version: 12.8 | |-----------------------------------------+------------------------+----------------------+ | GPU Name Persistence-M | Bus-Id Disp.A | Volatile Uncorr. ECC | | Fan Temp Perf Pwr:Usage/Cap | Memory-Usage | GPU-Util Compute M. | | | | MIG M. | |=========================================+========================+======================| | 0 NVIDIA H100 80GB HBM3 On | 00000000:61:00.0 Off | 0 | | N/A 34C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:62:00.0 Off | 0 | | N/A 38C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 40C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 35C P0 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 41C P0 130W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 40C P0 123W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 36C P0 118W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 247336 C /usr/bin/python3 1510MiB | | 1 N/A N/A 247337 C /usr/bin/python3 1510MiB | | 2 N/A N/A 247338 C /usr/bin/python3 1510MiB | | 3 N/A N/A 247339 C /usr/bin/python3 1510MiB | | 4 N/A N/A 247340 C /usr/bin/python3 1510MiB | | 5 N/A N/A 247341 C /usr/bin/python3 1510MiB | | 6 N/A N/A 247342 C /usr/bin/python3 1510MiB | | 7 N/A N/A 247343 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.8282 train_time:0ms step_avg:0.04ms step:1/1575 train_time:81ms step_avg:80.72ms step:2/1575 train_time:104ms step_avg:52.21ms step:3/1575 train_time:125ms step_avg:41.81ms step:4/1575 train_time:149ms step_avg:37.35ms step:5/1575 train_time:180ms step_avg:35.97ms step:6/1575 train_time:286ms step_avg:47.59ms step:7/1575 train_time:304ms step_avg:43.47ms step:8/1575 train_time:325ms step_avg:40.67ms step:9/1575 train_time:356ms step_avg:39.55ms step:10/1575 train_time:395ms step_avg:39.49ms step:11/1575 train_time:426ms step_avg:38.69ms step:12/1575 train_time:464ms step_avg:38.69ms step:13/1575 train_time:495ms step_avg:38.10ms step:14/1575 train_time:534ms step_avg:38.16ms step:15/1575 train_time:565ms step_avg:37.66ms step:16/1575 train_time:604ms step_avg:37.72ms step:17/1575 train_time:635ms step_avg:37.33ms step:18/1575 train_time:674ms step_avg:37.43ms step:19/1575 train_time:705ms step_avg:37.10ms step:20/1575 train_time:744ms step_avg:37.18ms step:21/1575 train_time:774ms step_avg:36.88ms step:22/1575 train_time:813ms step_avg:36.96ms step:23/1575 train_time:844ms step_avg:36.70ms step:24/1575 train_time:883ms step_avg:36.78ms step:25/1575 train_time:914ms step_avg:36.55ms step:26/1575 train_time:952ms step_avg:36.63ms step:27/1575 train_time:984ms step_avg:36.43ms step:28/1575 train_time:1023ms step_avg:36.52ms step:29/1575 train_time:1053ms step_avg:36.31ms step:30/1575 train_time:1092ms step_avg:36.40ms step:31/1575 train_time:1123ms step_avg:36.22ms step:32/1575 train_time:1162ms step_avg:36.31ms step:33/1575 train_time:1193ms step_avg:36.16ms step:34/1575 train_time:1232ms step_avg:36.25ms step:35/1575 train_time:1264ms step_avg:36.12ms step:36/1575 train_time:1303ms step_avg:36.20ms step:37/1575 train_time:1334ms step_avg:36.06ms step:38/1575 train_time:1373ms step_avg:36.14ms step:39/1575 train_time:1405ms step_avg:36.02ms step:40/1575 train_time:1444ms step_avg:36.09ms step:41/1575 train_time:1474ms step_avg:35.96ms step:42/1575 train_time:1513ms step_avg:36.03ms step:43/1575 train_time:1544ms step_avg:35.90ms step:44/1575 train_time:1582ms step_avg:35.96ms step:45/1575 train_time:1613ms step_avg:35.85ms step:46/1575 train_time:1652ms step_avg:35.91ms step:47/1575 train_time:1683ms step_avg:35.81ms step:48/1575 train_time:1722ms step_avg:35.87ms step:49/1575 train_time:1753ms step_avg:35.77ms step:50/1575 train_time:1791ms step_avg:35.82ms step:51/1575 train_time:1822ms step_avg:35.73ms step:52/1575 train_time:1861ms step_avg:35.79ms step:53/1575 train_time:1892ms step_avg:35.70ms step:54/1575 train_time:1931ms step_avg:35.75ms step:55/1575 train_time:1962ms step_avg:35.67ms step:56/1575 train_time:2000ms step_avg:35.72ms step:57/1575 train_time:2031ms step_avg:35.64ms step:58/1575 train_time:2070ms step_avg:35.69ms step:59/1575 train_time:2101ms step_avg:35.61ms step:60/1575 train_time:2140ms step_avg:35.67ms step:61/1575 train_time:2171ms step_avg:35.59ms step:62/1575 train_time:2209ms step_avg:35.64ms step:63/1575 train_time:2240ms step_avg:35.56ms step:64/1575 train_time:2279ms step_avg:35.61ms step:65/1575 train_time:2310ms step_avg:35.54ms step:66/1575 train_time:2348ms step_avg:35.58ms step:67/1575 train_time:2380ms step_avg:35.52ms step:68/1575 train_time:2419ms step_avg:35.57ms step:69/1575 train_time:2449ms step_avg:35.50ms step:70/1575 train_time:2488ms step_avg:35.54ms step:71/1575 train_time:2519ms step_avg:35.47ms step:72/1575 train_time:2558ms step_avg:35.52ms step:73/1575 train_time:2589ms step_avg:35.46ms step:74/1575 train_time:2628ms step_avg:35.51ms step:75/1575 train_time:2658ms step_avg:35.44ms step:76/1575 train_time:2696ms step_avg:35.48ms step:77/1575 train_time:2727ms step_avg:35.42ms step:78/1575 train_time:2766ms step_avg:35.46ms step:79/1575 train_time:2797ms step_avg:35.40ms step:80/1575 train_time:2836ms step_avg:35.45ms step:81/1575 train_time:2867ms step_avg:35.39ms step:82/1575 train_time:2905ms step_avg:35.43ms step:83/1575 train_time:2936ms step_avg:35.38ms step:84/1575 train_time:2975ms step_avg:35.42ms step:85/1575 train_time:3006ms step_avg:35.36ms step:86/1575 train_time:3044ms step_avg:35.40ms step:87/1575 train_time:3075ms step_avg:35.35ms step:88/1575 train_time:3114ms step_avg:35.39ms step:89/1575 train_time:3145ms step_avg:35.34ms step:90/1575 train_time:3185ms step_avg:35.39ms step:91/1575 train_time:3215ms step_avg:35.33ms step:92/1575 train_time:3253ms step_avg:35.36ms step:93/1575 train_time:3284ms step_avg:35.32ms step:94/1575 train_time:3323ms step_avg:35.36ms step:95/1575 train_time:3354ms step_avg:35.31ms step:96/1575 train_time:3393ms step_avg:35.34ms step:97/1575 train_time:3424ms step_avg:35.29ms step:98/1575 train_time:3462ms step_avg:35.33ms step:99/1575 train_time:3493ms step_avg:35.28ms step:100/1575 train_time:3532ms step_avg:35.32ms step:101/1575 train_time:3562ms step_avg:35.27ms step:102/1575 train_time:3601ms step_avg:35.31ms step:103/1575 train_time:3632ms step_avg:35.27ms step:104/1575 train_time:3671ms step_avg:35.30ms step:105/1575 train_time:3702ms step_avg:35.26ms step:106/1575 train_time:3741ms step_avg:35.29ms step:107/1575 train_time:3772ms step_avg:35.25ms step:108/1575 train_time:3810ms step_avg:35.28ms step:109/1575 train_time:3842ms step_avg:35.24ms step:110/1575 train_time:3880ms step_avg:35.28ms step:111/1575 train_time:3911ms step_avg:35.24ms step:112/1575 train_time:3950ms step_avg:35.27ms step:113/1575 train_time:3981ms step_avg:35.23ms step:114/1575 train_time:4020ms step_avg:35.27ms step:115/1575 train_time:4051ms step_avg:35.22ms step:116/1575 train_time:4089ms step_avg:35.25ms step:117/1575 train_time:4120ms step_avg:35.21ms step:118/1575 train_time:4159ms step_avg:35.24ms step:119/1575 train_time:4189ms step_avg:35.20ms step:120/1575 train_time:4228ms step_avg:35.23ms step:121/1575 train_time:4259ms step_avg:35.20ms step:122/1575 train_time:4297ms step_avg:35.22ms step:123/1575 train_time:4328ms step_avg:35.19ms step:124/1575 train_time:4367ms step_avg:35.22ms step:125/1575 train_time:4398ms step_avg:35.18ms step:126/1575 train_time:4437ms step_avg:35.21ms step:127/1575 train_time:4468ms step_avg:35.18ms step:128/1575 train_time:4506ms step_avg:35.20ms step:129/1575 train_time:4537ms step_avg:35.17ms step:130/1575 train_time:4576ms step_avg:35.20ms step:131/1575 train_time:4607ms step_avg:35.17ms step:132/1575 train_time:4646ms step_avg:35.19ms step:133/1575 train_time:4676ms step_avg:35.16ms step:134/1575 train_time:4715ms step_avg:35.19ms step:135/1575 train_time:4746ms step_avg:35.15ms step:136/1575 train_time:4784ms step_avg:35.18ms step:137/1575 train_time:4815ms step_avg:35.15ms step:138/1575 train_time:4854ms step_avg:35.17ms step:139/1575 train_time:4884ms step_avg:35.14ms step:140/1575 train_time:4923ms step_avg:35.17ms step:141/1575 train_time:4954ms step_avg:35.14ms step:142/1575 train_time:4993ms step_avg:35.16ms step:143/1575 train_time:5025ms step_avg:35.14ms step:144/1575 train_time:5063ms step_avg:35.16ms step:145/1575 train_time:5094ms step_avg:35.13ms step:146/1575 train_time:5133ms step_avg:35.16ms step:147/1575 train_time:5164ms step_avg:35.13ms step:148/1575 train_time:5203ms step_avg:35.15ms step:149/1575 train_time:5233ms step_avg:35.12ms step:150/1575 train_time:5272ms step_avg:35.15ms step:151/1575 train_time:5303ms step_avg:35.12ms step:152/1575 train_time:5342ms step_avg:35.14ms step:153/1575 train_time:5373ms step_avg:35.12ms step:154/1575 train_time:5411ms step_avg:35.14ms step:155/1575 train_time:5442ms step_avg:35.11ms step:156/1575 train_time:5481ms step_avg:35.14ms step:157/1575 train_time:5512ms step_avg:35.11ms step:158/1575 train_time:5551ms step_avg:35.13ms step:159/1575 train_time:5582ms step_avg:35.11ms step:160/1575 train_time:5621ms step_avg:35.13ms step:161/1575 train_time:5651ms step_avg:35.10ms step:162/1575 train_time:5690ms step_avg:35.12ms step:163/1575 train_time:5721ms step_avg:35.10ms step:164/1575 train_time:5759ms step_avg:35.12ms step:165/1575 train_time:5791ms step_avg:35.09ms step:166/1575 train_time:5829ms step_avg:35.12ms step:167/1575 train_time:5860ms step_avg:35.09ms step:168/1575 train_time:5899ms step_avg:35.11ms step:169/1575 train_time:5929ms step_avg:35.09ms step:170/1575 train_time:5968ms step_avg:35.10ms step:171/1575 train_time:5999ms step_avg:35.08ms step:172/1575 train_time:6037ms step_avg:35.10ms step:173/1575 train_time:6068ms step_avg:35.07ms step:174/1575 train_time:6106ms step_avg:35.09ms step:175/1575 train_time:6137ms step_avg:35.07ms step:176/1575 train_time:6176ms step_avg:35.09ms step:177/1575 train_time:6207ms step_avg:35.07ms step:178/1575 train_time:6245ms step_avg:35.08ms step:179/1575 train_time:6276ms step_avg:35.06ms step:180/1575 train_time:6314ms step_avg:35.08ms step:181/1575 train_time:6345ms step_avg:35.05ms step:182/1575 train_time:6383ms step_avg:35.07ms step:183/1575 train_time:6414ms step_avg:35.05ms step:184/1575 train_time:6453ms step_avg:35.07ms step:185/1575 train_time:6484ms step_avg:35.05ms step:186/1575 train_time:6522ms step_avg:35.07ms step:187/1575 train_time:6553ms step_avg:35.04ms step:188/1575 train_time:6592ms step_avg:35.06ms step:189/1575 train_time:6623ms step_avg:35.04ms step:190/1575 train_time:6661ms step_avg:35.06ms step:191/1575 train_time:6692ms step_avg:35.04ms step:192/1575 train_time:6730ms step_avg:35.05ms step:193/1575 train_time:6761ms step_avg:35.03ms step:194/1575 train_time:6800ms step_avg:35.05ms step:195/1575 train_time:6831ms step_avg:35.03ms step:196/1575 train_time:6869ms step_avg:35.05ms step:197/1575 train_time:6900ms step_avg:35.03ms step:198/1575 train_time:6939ms step_avg:35.05ms step:199/1575 train_time:6970ms step_avg:35.02ms step:200/1575 train_time:7008ms step_avg:35.04ms step:201/1575 train_time:7039ms step_avg:35.02ms step:202/1575 train_time:7077ms step_avg:35.04ms step:203/1575 train_time:7108ms step_avg:35.02ms step:204/1575 train_time:7147ms step_avg:35.03ms step:205/1575 train_time:7178ms step_avg:35.01ms step:206/1575 train_time:7216ms step_avg:35.03ms step:207/1575 train_time:7247ms step_avg:35.01ms step:208/1575 train_time:7286ms step_avg:35.03ms step:209/1575 train_time:7316ms step_avg:35.01ms step:210/1575 train_time:7356ms step_avg:35.03ms step:211/1575 train_time:7386ms step_avg:35.01ms step:212/1575 train_time:7425ms step_avg:35.02ms step:213/1575 train_time:7455ms step_avg:35.00ms step:214/1575 train_time:7494ms step_avg:35.02ms step:215/1575 train_time:7525ms step_avg:35.00ms step:216/1575 train_time:7563ms step_avg:35.02ms step:217/1575 train_time:7594ms step_avg:35.00ms step:218/1575 train_time:7633ms step_avg:35.01ms step:219/1575 train_time:7664ms step_avg:35.00ms step:220/1575 train_time:7703ms step_avg:35.01ms step:221/1575 train_time:7734ms step_avg:35.00ms step:222/1575 train_time:7772ms step_avg:35.01ms step:223/1575 train_time:7803ms step_avg:34.99ms step:224/1575 train_time:7841ms step_avg:35.01ms step:225/1575 train_time:7872ms step_avg:34.99ms step:226/1575 train_time:7911ms step_avg:35.00ms step:227/1575 train_time:7942ms step_avg:34.99ms step:228/1575 train_time:7981ms step_avg:35.00ms step:229/1575 train_time:8011ms step_avg:34.98ms step:230/1575 train_time:8050ms step_avg:35.00ms step:231/1575 train_time:8081ms step_avg:34.98ms step:232/1575 train_time:8119ms step_avg:35.00ms step:233/1575 train_time:8150ms step_avg:34.98ms step:234/1575 train_time:8189ms step_avg:34.99ms step:235/1575 train_time:8220ms step_avg:34.98ms step:236/1575 train_time:8259ms step_avg:35.00ms step:237/1575 train_time:8290ms step_avg:34.98ms step:238/1575 train_time:8328ms step_avg:34.99ms step:239/1575 train_time:8359ms step_avg:34.98ms step:240/1575 train_time:8398ms step_avg:34.99ms step:241/1575 train_time:8429ms step_avg:34.97ms step:242/1575 train_time:8467ms step_avg:34.99ms step:243/1575 train_time:8498ms step_avg:34.97ms step:244/1575 train_time:8537ms step_avg:34.99ms step:245/1575 train_time:8568ms step_avg:34.97ms step:246/1575 train_time:8606ms step_avg:34.98ms step:247/1575 train_time:8637ms step_avg:34.97ms step:248/1575 train_time:8676ms step_avg:34.98ms step:249/1575 train_time:8706ms step_avg:34.97ms step:250/1575 train_time:8745ms step_avg:34.98ms step:250/1575 val_loss:4.5749 train_time:8793ms step_avg:35.17ms step:251/1575 train_time:8813ms step_avg:35.11ms step:252/1575 train_time:8833ms step_avg:35.05ms step:253/1575 train_time:8851ms step_avg:34.99ms step:254/1575 train_time:8887ms step_avg:34.99ms step:255/1575 train_time:8919ms step_avg:34.98ms step:256/1575 train_time:8959ms step_avg:34.99ms step:257/1575 train_time:8991ms step_avg:34.99ms step:258/1575 train_time:9031ms step_avg:35.00ms step:259/1575 train_time:9062ms step_avg:34.99ms step:260/1575 train_time:9101ms step_avg:35.00ms step:261/1575 train_time:9132ms step_avg:34.99ms step:262/1575 train_time:9170ms step_avg:35.00ms step:263/1575 train_time:9201ms step_avg:34.98ms step:264/1575 train_time:9239ms step_avg:35.00ms step:265/1575 train_time:9270ms step_avg:34.98ms step:266/1575 train_time:9309ms step_avg:34.99ms step:267/1575 train_time:9339ms step_avg:34.98ms step:268/1575 train_time:9378ms step_avg:34.99ms step:269/1575 train_time:9409ms step_avg:34.98ms step:270/1575 train_time:9447ms step_avg:34.99ms step:271/1575 train_time:9478ms step_avg:34.97ms step:272/1575 train_time:9516ms step_avg:34.99ms step:273/1575 train_time:9547ms step_avg:34.97ms step:274/1575 train_time:9585ms step_avg:34.98ms step:275/1575 train_time:9616ms step_avg:34.97ms step:276/1575 train_time:9655ms step_avg:34.98ms step:277/1575 train_time:9686ms step_avg:34.97ms step:278/1575 train_time:9724ms step_avg:34.98ms step:279/1575 train_time:9755ms step_avg:34.96ms step:280/1575 train_time:9793ms step_avg:34.98ms step:281/1575 train_time:9824ms step_avg:34.96ms step:282/1575 train_time:9863ms step_avg:34.98ms step:283/1575 train_time:9894ms step_avg:34.96ms step:284/1575 train_time:9932ms step_avg:34.97ms step:285/1575 train_time:9963ms step_avg:34.96ms step:286/1575 train_time:10002ms step_avg:34.97ms step:287/1575 train_time:10032ms step_avg:34.96ms step:288/1575 train_time:10071ms step_avg:34.97ms step:289/1575 train_time:10102ms step_avg:34.95ms step:290/1575 train_time:10140ms step_avg:34.97ms step:291/1575 train_time:10171ms step_avg:34.95ms step:292/1575 train_time:10210ms step_avg:34.96ms step:293/1575 train_time:10240ms step_avg:34.95ms step:294/1575 train_time:10279ms step_avg:34.96ms step:295/1575 train_time:10310ms step_avg:34.95ms step:296/1575 train_time:10348ms step_avg:34.96ms step:297/1575 train_time:10379ms step_avg:34.95ms step:298/1575 train_time:10418ms step_avg:34.96ms step:299/1575 train_time:10448ms step_avg:34.94ms step:300/1575 train_time:10487ms step_avg:34.96ms step:301/1575 train_time:10517ms step_avg:34.94ms step:302/1575 train_time:10556ms step_avg:34.95ms step:303/1575 train_time:10586ms step_avg:34.94ms step:304/1575 train_time:10625ms step_avg:34.95ms step:305/1575 train_time:10656ms step_avg:34.94ms step:306/1575 train_time:10694ms step_avg:34.95ms step:307/1575 train_time:10725ms step_avg:34.93ms step:308/1575 train_time:10763ms step_avg:34.95ms step:309/1575 train_time:10794ms step_avg:34.93ms step:310/1575 train_time:10832ms step_avg:34.94ms step:311/1575 train_time:10864ms step_avg:34.93ms step:312/1575 train_time:10902ms step_avg:34.94ms step:313/1575 train_time:10933ms step_avg:34.93ms step:314/1575 train_time:10971ms step_avg:34.94ms step:315/1575 train_time:11002ms step_avg:34.93ms step:316/1575 train_time:11041ms step_avg:34.94ms step:317/1575 train_time:11071ms step_avg:34.93ms step:318/1575 train_time:11110ms step_avg:34.94ms step:319/1575 train_time:11141ms step_avg:34.92ms step:320/1575 train_time:11180ms step_avg:34.94ms step:321/1575 train_time:11211ms step_avg:34.93ms step:322/1575 train_time:11249ms step_avg:34.94ms step:323/1575 train_time:11280ms step_avg:34.92ms step:324/1575 train_time:11320ms step_avg:34.94ms step:325/1575 train_time:11350ms step_avg:34.92ms step:326/1575 train_time:11389ms step_avg:34.93ms step:327/1575 train_time:11420ms step_avg:34.92ms step:328/1575 train_time:11458ms step_avg:34.93ms step:329/1575 train_time:11489ms step_avg:34.92ms step:330/1575 train_time:11527ms step_avg:34.93ms step:331/1575 train_time:11558ms step_avg:34.92ms step:332/1575 train_time:11597ms step_avg:34.93ms step:333/1575 train_time:11628ms step_avg:34.92ms step:334/1575 train_time:11667ms step_avg:34.93ms step:335/1575 train_time:11698ms step_avg:34.92ms step:336/1575 train_time:11736ms step_avg:34.93ms step:337/1575 train_time:11766ms step_avg:34.91ms step:338/1575 train_time:11805ms step_avg:34.93ms step:339/1575 train_time:11835ms step_avg:34.91ms step:340/1575 train_time:11874ms step_avg:34.92ms step:341/1575 train_time:11905ms step_avg:34.91ms step:342/1575 train_time:11943ms step_avg:34.92ms step:343/1575 train_time:11974ms step_avg:34.91ms step:344/1575 train_time:12013ms step_avg:34.92ms step:345/1575 train_time:12044ms step_avg:34.91ms step:346/1575 train_time:12083ms step_avg:34.92ms step:347/1575 train_time:12113ms step_avg:34.91ms step:348/1575 train_time:12152ms step_avg:34.92ms step:349/1575 train_time:12182ms step_avg:34.91ms step:350/1575 train_time:12221ms step_avg:34.92ms step:351/1575 train_time:12252ms step_avg:34.91ms step:352/1575 train_time:12291ms step_avg:34.92ms step:353/1575 train_time:12321ms step_avg:34.90ms step:354/1575 train_time:12359ms step_avg:34.91ms step:355/1575 train_time:12390ms step_avg:34.90ms step:356/1575 train_time:12428ms step_avg:34.91ms step:357/1575 train_time:12460ms step_avg:34.90ms step:358/1575 train_time:12499ms step_avg:34.91ms step:359/1575 train_time:12528ms step_avg:34.90ms step:360/1575 train_time:12567ms step_avg:34.91ms step:361/1575 train_time:12597ms step_avg:34.90ms step:362/1575 train_time:12636ms step_avg:34.91ms step:363/1575 train_time:12667ms step_avg:34.89ms step:364/1575 train_time:12705ms step_avg:34.91ms step:365/1575 train_time:12736ms step_avg:34.89ms step:366/1575 train_time:12774ms step_avg:34.90ms step:367/1575 train_time:12805ms step_avg:34.89ms step:368/1575 train_time:12843ms step_avg:34.90ms step:369/1575 train_time:12875ms step_avg:34.89ms step:370/1575 train_time:12914ms step_avg:34.90ms step:371/1575 train_time:12944ms step_avg:34.89ms step:372/1575 train_time:12983ms step_avg:34.90ms step:373/1575 train_time:13014ms step_avg:34.89ms step:374/1575 train_time:13053ms step_avg:34.90ms step:375/1575 train_time:13083ms step_avg:34.89ms step:376/1575 train_time:13122ms step_avg:34.90ms step:377/1575 train_time:13153ms step_avg:34.89ms step:378/1575 train_time:13192ms step_avg:34.90ms step:379/1575 train_time:13222ms step_avg:34.89ms step:380/1575 train_time:13261ms step_avg:34.90ms step:381/1575 train_time:13292ms step_avg:34.89ms step:382/1575 train_time:13330ms step_avg:34.90ms step:383/1575 train_time:13360ms step_avg:34.88ms step:384/1575 train_time:13399ms step_avg:34.89ms step:385/1575 train_time:13430ms step_avg:34.88ms step:386/1575 train_time:13469ms step_avg:34.89ms step:387/1575 train_time:13499ms step_avg:34.88ms step:388/1575 train_time:13538ms step_avg:34.89ms step:389/1575 train_time:13569ms step_avg:34.88ms step:390/1575 train_time:13607ms step_avg:34.89ms step:391/1575 train_time:13638ms step_avg:34.88ms step:392/1575 train_time:13676ms step_avg:34.89ms step:393/1575 train_time:13707ms step_avg:34.88ms step:394/1575 train_time:13746ms step_avg:34.89ms step:395/1575 train_time:13776ms step_avg:34.88ms step:396/1575 train_time:13815ms step_avg:34.89ms step:397/1575 train_time:13846ms step_avg:34.88ms step:398/1575 train_time:13885ms step_avg:34.89ms step:399/1575 train_time:13915ms step_avg:34.87ms step:400/1575 train_time:13954ms step_avg:34.88ms step:401/1575 train_time:13984ms step_avg:34.87ms step:402/1575 train_time:14024ms step_avg:34.88ms step:403/1575 train_time:14054ms step_avg:34.87ms step:404/1575 train_time:14092ms step_avg:34.88ms step:405/1575 train_time:14123ms step_avg:34.87ms step:406/1575 train_time:14162ms step_avg:34.88ms step:407/1575 train_time:14193ms step_avg:34.87ms step:408/1575 train_time:14231ms step_avg:34.88ms step:409/1575 train_time:14262ms step_avg:34.87ms step:410/1575 train_time:14301ms step_avg:34.88ms step:411/1575 train_time:14331ms step_avg:34.87ms step:412/1575 train_time:14370ms step_avg:34.88ms step:413/1575 train_time:14401ms step_avg:34.87ms step:414/1575 train_time:14439ms step_avg:34.88ms step:415/1575 train_time:14470ms step_avg:34.87ms step:416/1575 train_time:14509ms step_avg:34.88ms step:417/1575 train_time:14539ms step_avg:34.87ms step:418/1575 train_time:14578ms step_avg:34.88ms step:419/1575 train_time:14609ms step_avg:34.87ms step:420/1575 train_time:14648ms step_avg:34.88ms step:421/1575 train_time:14678ms step_avg:34.87ms step:422/1575 train_time:14724ms step_avg:34.89ms step:423/1575 train_time:14748ms step_avg:34.86ms step:424/1575 train_time:14786ms step_avg:34.87ms step:425/1575 train_time:14817ms step_avg:34.86ms step:426/1575 train_time:14856ms step_avg:34.87ms step:427/1575 train_time:14887ms step_avg:34.86ms step:428/1575 train_time:14925ms step_avg:34.87ms step:429/1575 train_time:14956ms step_avg:34.86ms step:430/1575 train_time:14994ms step_avg:34.87ms step:431/1575 train_time:15025ms step_avg:34.86ms step:432/1575 train_time:15063ms step_avg:34.87ms step:433/1575 train_time:15094ms step_avg:34.86ms step:434/1575 train_time:15133ms step_avg:34.87ms step:435/1575 train_time:15163ms step_avg:34.86ms step:436/1575 train_time:15202ms step_avg:34.87ms step:437/1575 train_time:15233ms step_avg:34.86ms step:438/1575 train_time:15271ms step_avg:34.87ms step:439/1575 train_time:15302ms step_avg:34.86ms step:440/1575 train_time:15341ms step_avg:34.87ms step:441/1575 train_time:15372ms step_avg:34.86ms step:442/1575 train_time:15410ms step_avg:34.86ms step:443/1575 train_time:15441ms step_avg:34.86ms step:444/1575 train_time:15480ms step_avg:34.86ms step:445/1575 train_time:15510ms step_avg:34.85ms step:446/1575 train_time:15549ms step_avg:34.86ms step:447/1575 train_time:15580ms step_avg:34.85ms step:448/1575 train_time:15618ms step_avg:34.86ms step:449/1575 train_time:15649ms step_avg:34.85ms step:450/1575 train_time:15687ms step_avg:34.86ms step:451/1575 train_time:15718ms step_avg:34.85ms step:452/1575 train_time:15756ms step_avg:34.86ms step:453/1575 train_time:15787ms step_avg:34.85ms step:454/1575 train_time:15826ms step_avg:34.86ms step:455/1575 train_time:15857ms step_avg:34.85ms step:456/1575 train_time:15895ms step_avg:34.86ms step:457/1575 train_time:15926ms step_avg:34.85ms step:458/1575 train_time:15965ms step_avg:34.86ms step:459/1575 train_time:15995ms step_avg:34.85ms step:460/1575 train_time:16034ms step_avg:34.86ms step:461/1575 train_time:16065ms step_avg:34.85ms step:462/1575 train_time:16104ms step_avg:34.86ms step:463/1575 train_time:16135ms step_avg:34.85ms step:464/1575 train_time:16173ms step_avg:34.86ms step:465/1575 train_time:16204ms step_avg:34.85ms step:466/1575 train_time:16243ms step_avg:34.86ms step:467/1575 train_time:16273ms step_avg:34.85ms step:468/1575 train_time:16312ms step_avg:34.86ms step:469/1575 train_time:16343ms step_avg:34.85ms step:470/1575 train_time:16381ms step_avg:34.85ms step:471/1575 train_time:16412ms step_avg:34.85ms step:472/1575 train_time:16451ms step_avg:34.85ms step:473/1575 train_time:16482ms step_avg:34.85ms step:474/1575 train_time:16521ms step_avg:34.85ms step:475/1575 train_time:16551ms step_avg:34.84ms step:476/1575 train_time:16590ms step_avg:34.85ms step:477/1575 train_time:16620ms step_avg:34.84ms step:478/1575 train_time:16659ms step_avg:34.85ms step:479/1575 train_time:16690ms step_avg:34.84ms step:480/1575 train_time:16729ms step_avg:34.85ms step:481/1575 train_time:16760ms step_avg:34.84ms step:482/1575 train_time:16799ms step_avg:34.85ms step:483/1575 train_time:16829ms step_avg:34.84ms step:484/1575 train_time:16868ms step_avg:34.85ms step:485/1575 train_time:16898ms step_avg:34.84ms step:486/1575 train_time:16937ms step_avg:34.85ms step:487/1575 train_time:16968ms step_avg:34.84ms step:488/1575 train_time:17006ms step_avg:34.85ms step:489/1575 train_time:17037ms step_avg:34.84ms step:490/1575 train_time:17075ms step_avg:34.85ms step:491/1575 train_time:17106ms step_avg:34.84ms step:492/1575 train_time:17145ms step_avg:34.85ms step:493/1575 train_time:17175ms step_avg:34.84ms step:494/1575 train_time:17214ms step_avg:34.85ms step:495/1575 train_time:17245ms step_avg:34.84ms step:496/1575 train_time:17283ms step_avg:34.85ms step:497/1575 train_time:17314ms step_avg:34.84ms step:498/1575 train_time:17353ms step_avg:34.84ms step:499/1575 train_time:17383ms step_avg:34.84ms step:500/1575 train_time:17422ms step_avg:34.84ms step:500/1575 val_loss:4.2357 train_time:17470ms step_avg:34.94ms step:501/1575 train_time:17491ms step_avg:34.91ms step:502/1575 train_time:17511ms step_avg:34.88ms step:503/1575 train_time:17529ms step_avg:34.85ms step:504/1575 train_time:17565ms step_avg:34.85ms step:505/1575 train_time:17598ms step_avg:34.85ms step:506/1575 train_time:17638ms step_avg:34.86ms step:507/1575 train_time:17669ms step_avg:34.85ms step:508/1575 train_time:17708ms step_avg:34.86ms step:509/1575 train_time:17739ms step_avg:34.85ms step:510/1575 train_time:17777ms step_avg:34.86ms step:511/1575 train_time:17808ms step_avg:34.85ms step:512/1575 train_time:17848ms step_avg:34.86ms step:513/1575 train_time:17924ms step_avg:34.94ms step:514/1575 train_time:17978ms step_avg:34.98ms step:515/1575 train_time:18039ms step_avg:35.03ms step:516/1575 train_time:18098ms step_avg:35.07ms step:517/1575 train_time:18160ms step_avg:35.13ms step:518/1575 train_time:18219ms step_avg:35.17ms step:519/1575 train_time:18281ms step_avg:35.22ms step:520/1575 train_time:18340ms step_avg:35.27ms step:521/1575 train_time:18404ms step_avg:35.32ms step:522/1575 train_time:18463ms step_avg:35.37ms step:523/1575 train_time:18528ms step_avg:35.43ms step:524/1575 train_time:18588ms step_avg:35.47ms step:525/1575 train_time:18652ms step_avg:35.53ms step:526/1575 train_time:18712ms step_avg:35.57ms step:527/1575 train_time:18776ms step_avg:35.63ms step:528/1575 train_time:18837ms step_avg:35.68ms step:529/1575 train_time:18900ms step_avg:35.73ms step:530/1575 train_time:18960ms step_avg:35.77ms step:531/1575 train_time:19023ms step_avg:35.82ms step:532/1575 train_time:19082ms step_avg:35.87ms step:533/1575 train_time:19145ms step_avg:35.92ms step:534/1575 train_time:19203ms step_avg:35.96ms step:535/1575 train_time:19267ms step_avg:36.01ms step:536/1575 train_time:19327ms step_avg:36.06ms step:537/1575 train_time:19391ms step_avg:36.11ms step:538/1575 train_time:19450ms step_avg:36.15ms step:539/1575 train_time:19513ms step_avg:36.20ms step:540/1575 train_time:19573ms step_avg:36.25ms step:541/1575 train_time:19637ms step_avg:36.30ms step:542/1575 train_time:19696ms step_avg:36.34ms step:543/1575 train_time:19761ms step_avg:36.39ms step:544/1575 train_time:19821ms step_avg:36.44ms step:545/1575 train_time:19884ms step_avg:36.49ms step:546/1575 train_time:19944ms step_avg:36.53ms step:547/1575 train_time:20008ms step_avg:36.58ms step:548/1575 train_time:20067ms step_avg:36.62ms step:549/1575 train_time:20131ms step_avg:36.67ms step:550/1575 train_time:20190ms step_avg:36.71ms step:551/1575 train_time:20253ms step_avg:36.76ms step:552/1575 train_time:20313ms step_avg:36.80ms step:553/1575 train_time:20377ms step_avg:36.85ms step:554/1575 train_time:20437ms step_avg:36.89ms step:555/1575 train_time:20499ms step_avg:36.94ms step:556/1575 train_time:20559ms step_avg:36.98ms step:557/1575 train_time:20622ms step_avg:37.02ms step:558/1575 train_time:20681ms step_avg:37.06ms step:559/1575 train_time:20745ms step_avg:37.11ms step:560/1575 train_time:20805ms step_avg:37.15ms step:561/1575 train_time:20868ms step_avg:37.20ms step:562/1575 train_time:20928ms step_avg:37.24ms step:563/1575 train_time:20990ms step_avg:37.28ms step:564/1575 train_time:21050ms step_avg:37.32ms step:565/1575 train_time:21114ms step_avg:37.37ms step:566/1575 train_time:21172ms step_avg:37.41ms step:567/1575 train_time:21235ms step_avg:37.45ms step:568/1575 train_time:21297ms step_avg:37.49ms step:569/1575 train_time:21364ms step_avg:37.55ms step:570/1575 train_time:21421ms step_avg:37.58ms step:571/1575 train_time:21483ms step_avg:37.62ms step:572/1575 train_time:21542ms step_avg:37.66ms step:573/1575 train_time:21604ms step_avg:37.70ms step:574/1575 train_time:21663ms step_avg:37.74ms step:575/1575 train_time:21726ms step_avg:37.78ms step:576/1575 train_time:21786ms step_avg:37.82ms step:577/1575 train_time:21851ms step_avg:37.87ms step:578/1575 train_time:21911ms step_avg:37.91ms step:579/1575 train_time:21972ms step_avg:37.95ms step:580/1575 train_time:22032ms step_avg:37.99ms step:581/1575 train_time:22095ms step_avg:38.03ms step:582/1575 train_time:22155ms step_avg:38.07ms step:583/1575 train_time:22219ms step_avg:38.11ms step:584/1575 train_time:22278ms step_avg:38.15ms step:585/1575 train_time:22341ms step_avg:38.19ms step:586/1575 train_time:22400ms step_avg:38.23ms step:587/1575 train_time:22463ms step_avg:38.27ms step:588/1575 train_time:22523ms step_avg:38.30ms step:589/1575 train_time:22586ms step_avg:38.35ms step:590/1575 train_time:22645ms step_avg:38.38ms step:591/1575 train_time:22708ms step_avg:38.42ms step:592/1575 train_time:22768ms step_avg:38.46ms step:593/1575 train_time:22831ms step_avg:38.50ms step:594/1575 train_time:22890ms step_avg:38.54ms step:595/1575 train_time:22954ms step_avg:38.58ms step:596/1575 train_time:23013ms step_avg:38.61ms step:597/1575 train_time:23077ms step_avg:38.65ms step:598/1575 train_time:23136ms step_avg:38.69ms step:599/1575 train_time:23199ms step_avg:38.73ms step:600/1575 train_time:23259ms step_avg:38.76ms step:601/1575 train_time:23322ms step_avg:38.81ms step:602/1575 train_time:23381ms step_avg:38.84ms step:603/1575 train_time:23444ms step_avg:38.88ms step:604/1575 train_time:23503ms step_avg:38.91ms step:605/1575 train_time:23566ms step_avg:38.95ms step:606/1575 train_time:23625ms step_avg:38.99ms step:607/1575 train_time:23689ms step_avg:39.03ms step:608/1575 train_time:23748ms step_avg:39.06ms step:609/1575 train_time:23812ms step_avg:39.10ms step:610/1575 train_time:23871ms step_avg:39.13ms step:611/1575 train_time:23934ms step_avg:39.17ms step:612/1575 train_time:23994ms step_avg:39.21ms step:613/1575 train_time:24057ms step_avg:39.25ms step:614/1575 train_time:24118ms step_avg:39.28ms step:615/1575 train_time:24180ms step_avg:39.32ms step:616/1575 train_time:24240ms step_avg:39.35ms step:617/1575 train_time:24303ms step_avg:39.39ms step:618/1575 train_time:24362ms step_avg:39.42ms step:619/1575 train_time:24425ms step_avg:39.46ms step:620/1575 train_time:24484ms step_avg:39.49ms step:621/1575 train_time:24547ms step_avg:39.53ms step:622/1575 train_time:24606ms step_avg:39.56ms step:623/1575 train_time:24669ms step_avg:39.60ms step:624/1575 train_time:24729ms step_avg:39.63ms step:625/1575 train_time:24792ms step_avg:39.67ms step:626/1575 train_time:24851ms step_avg:39.70ms step:627/1575 train_time:24914ms step_avg:39.74ms step:628/1575 train_time:24974ms step_avg:39.77ms step:629/1575 train_time:25037ms step_avg:39.80ms step:630/1575 train_time:25097ms step_avg:39.84ms step:631/1575 train_time:25160ms step_avg:39.87ms step:632/1575 train_time:25220ms step_avg:39.90ms step:633/1575 train_time:25283ms step_avg:39.94ms step:634/1575 train_time:25344ms step_avg:39.97ms step:635/1575 train_time:25405ms step_avg:40.01ms step:636/1575 train_time:25465ms step_avg:40.04ms step:637/1575 train_time:25529ms step_avg:40.08ms step:638/1575 train_time:25588ms step_avg:40.11ms step:639/1575 train_time:25651ms step_avg:40.14ms step:640/1575 train_time:25711ms step_avg:40.17ms step:641/1575 train_time:25774ms step_avg:40.21ms step:642/1575 train_time:25833ms step_avg:40.24ms step:643/1575 train_time:25897ms step_avg:40.27ms step:644/1575 train_time:25956ms step_avg:40.30ms step:645/1575 train_time:26020ms step_avg:40.34ms step:646/1575 train_time:26079ms step_avg:40.37ms step:647/1575 train_time:26143ms step_avg:40.41ms step:648/1575 train_time:26202ms step_avg:40.44ms step:649/1575 train_time:26265ms step_avg:40.47ms step:650/1575 train_time:26324ms step_avg:40.50ms step:651/1575 train_time:26388ms step_avg:40.53ms step:652/1575 train_time:26447ms step_avg:40.56ms step:653/1575 train_time:26511ms step_avg:40.60ms step:654/1575 train_time:26570ms step_avg:40.63ms step:655/1575 train_time:26633ms step_avg:40.66ms step:656/1575 train_time:26693ms step_avg:40.69ms step:657/1575 train_time:26757ms step_avg:40.73ms step:658/1575 train_time:26815ms step_avg:40.75ms step:659/1575 train_time:26878ms step_avg:40.79ms step:660/1575 train_time:26938ms step_avg:40.81ms step:661/1575 train_time:27003ms step_avg:40.85ms step:662/1575 train_time:27063ms step_avg:40.88ms step:663/1575 train_time:27125ms step_avg:40.91ms step:664/1575 train_time:27184ms step_avg:40.94ms step:665/1575 train_time:27246ms step_avg:40.97ms step:666/1575 train_time:27306ms step_avg:41.00ms step:667/1575 train_time:27369ms step_avg:41.03ms step:668/1575 train_time:27428ms step_avg:41.06ms step:669/1575 train_time:27491ms step_avg:41.09ms step:670/1575 train_time:27550ms step_avg:41.12ms step:671/1575 train_time:27614ms step_avg:41.15ms step:672/1575 train_time:27672ms step_avg:41.18ms step:673/1575 train_time:27736ms step_avg:41.21ms step:674/1575 train_time:27795ms step_avg:41.24ms step:675/1575 train_time:27858ms step_avg:41.27ms step:676/1575 train_time:27917ms step_avg:41.30ms step:677/1575 train_time:27981ms step_avg:41.33ms step:678/1575 train_time:28043ms step_avg:41.36ms step:679/1575 train_time:28104ms step_avg:41.39ms step:680/1575 train_time:28164ms step_avg:41.42ms step:681/1575 train_time:28227ms step_avg:41.45ms step:682/1575 train_time:28286ms step_avg:41.48ms step:683/1575 train_time:28349ms step_avg:41.51ms step:684/1575 train_time:28408ms step_avg:41.53ms step:685/1575 train_time:28472ms step_avg:41.56ms step:686/1575 train_time:28531ms step_avg:41.59ms step:687/1575 train_time:28595ms step_avg:41.62ms step:688/1575 train_time:28654ms step_avg:41.65ms step:689/1575 train_time:28717ms step_avg:41.68ms step:690/1575 train_time:28776ms step_avg:41.70ms step:691/1575 train_time:28840ms step_avg:41.74ms step:692/1575 train_time:28900ms step_avg:41.76ms step:693/1575 train_time:28963ms step_avg:41.79ms step:694/1575 train_time:29023ms step_avg:41.82ms step:695/1575 train_time:29086ms step_avg:41.85ms step:696/1575 train_time:29146ms step_avg:41.88ms step:697/1575 train_time:29209ms step_avg:41.91ms step:698/1575 train_time:29268ms step_avg:41.93ms step:699/1575 train_time:29332ms step_avg:41.96ms step:700/1575 train_time:29391ms step_avg:41.99ms step:701/1575 train_time:29454ms step_avg:42.02ms step:702/1575 train_time:29514ms step_avg:42.04ms step:703/1575 train_time:29576ms step_avg:42.07ms step:704/1575 train_time:29637ms step_avg:42.10ms step:705/1575 train_time:29700ms step_avg:42.13ms step:706/1575 train_time:29759ms step_avg:42.15ms step:707/1575 train_time:29823ms step_avg:42.18ms step:708/1575 train_time:29883ms step_avg:42.21ms step:709/1575 train_time:29945ms step_avg:42.24ms step:710/1575 train_time:30005ms step_avg:42.26ms step:711/1575 train_time:30069ms step_avg:42.29ms step:712/1575 train_time:30128ms step_avg:42.31ms step:713/1575 train_time:30192ms step_avg:42.34ms step:714/1575 train_time:30251ms step_avg:42.37ms step:715/1575 train_time:30315ms step_avg:42.40ms step:716/1575 train_time:30374ms step_avg:42.42ms step:717/1575 train_time:30438ms step_avg:42.45ms step:718/1575 train_time:30497ms step_avg:42.47ms step:719/1575 train_time:30560ms step_avg:42.50ms step:720/1575 train_time:30619ms step_avg:42.53ms step:721/1575 train_time:30682ms step_avg:42.55ms step:722/1575 train_time:30741ms step_avg:42.58ms step:723/1575 train_time:30805ms step_avg:42.61ms step:724/1575 train_time:30865ms step_avg:42.63ms step:725/1575 train_time:30927ms step_avg:42.66ms step:726/1575 train_time:30987ms step_avg:42.68ms step:727/1575 train_time:31051ms step_avg:42.71ms step:728/1575 train_time:31110ms step_avg:42.73ms step:729/1575 train_time:31174ms step_avg:42.76ms step:730/1575 train_time:31234ms step_avg:42.79ms step:731/1575 train_time:31297ms step_avg:42.81ms step:732/1575 train_time:31356ms step_avg:42.84ms step:733/1575 train_time:31419ms step_avg:42.86ms step:734/1575 train_time:31479ms step_avg:42.89ms step:735/1575 train_time:31542ms step_avg:42.91ms step:736/1575 train_time:31601ms step_avg:42.94ms step:737/1575 train_time:31665ms step_avg:42.96ms step:738/1575 train_time:31723ms step_avg:42.99ms step:739/1575 train_time:31787ms step_avg:43.01ms step:740/1575 train_time:31846ms step_avg:43.04ms step:741/1575 train_time:31909ms step_avg:43.06ms step:742/1575 train_time:31968ms step_avg:43.08ms step:743/1575 train_time:32032ms step_avg:43.11ms step:744/1575 train_time:32092ms step_avg:43.13ms step:745/1575 train_time:32155ms step_avg:43.16ms step:746/1575 train_time:32213ms step_avg:43.18ms step:747/1575 train_time:32277ms step_avg:43.21ms step:748/1575 train_time:32336ms step_avg:43.23ms step:749/1575 train_time:32400ms step_avg:43.26ms step:750/1575 train_time:32459ms step_avg:43.28ms step:750/1575 val_loss:3.8798 train_time:32505ms step_avg:43.34ms step:751/1575 train_time:32526ms step_avg:43.31ms step:752/1575 train_time:32584ms step_avg:43.33ms step:753/1575 train_time:32650ms step_avg:43.36ms step:754/1575 train_time:32710ms step_avg:43.38ms step:755/1575 train_time:32773ms step_avg:43.41ms step:756/1575 train_time:32832ms step_avg:43.43ms step:757/1575 train_time:32894ms step_avg:43.45ms step:758/1575 train_time:32954ms step_avg:43.48ms step:759/1575 train_time:33017ms step_avg:43.50ms step:760/1575 train_time:33076ms step_avg:43.52ms step:761/1575 train_time:33139ms step_avg:43.55ms step:762/1575 train_time:33198ms step_avg:43.57ms step:763/1575 train_time:33261ms step_avg:43.59ms step:764/1575 train_time:33320ms step_avg:43.61ms step:765/1575 train_time:33383ms step_avg:43.64ms step:766/1575 train_time:33443ms step_avg:43.66ms step:767/1575 train_time:33506ms step_avg:43.69ms step:768/1575 train_time:33567ms step_avg:43.71ms step:769/1575 train_time:33632ms step_avg:43.74ms step:770/1575 train_time:33693ms step_avg:43.76ms step:771/1575 train_time:33756ms step_avg:43.78ms step:772/1575 train_time:33816ms step_avg:43.80ms step:773/1575 train_time:33880ms step_avg:43.83ms step:774/1575 train_time:33939ms step_avg:43.85ms step:775/1575 train_time:34002ms step_avg:43.87ms step:776/1575 train_time:34061ms step_avg:43.89ms step:777/1575 train_time:34123ms step_avg:43.92ms step:778/1575 train_time:34183ms step_avg:43.94ms step:779/1575 train_time:34246ms step_avg:43.96ms step:780/1575 train_time:34304ms step_avg:43.98ms step:781/1575 train_time:34367ms step_avg:44.00ms step:782/1575 train_time:34426ms step_avg:44.02ms step:783/1575 train_time:34489ms step_avg:44.05ms step:784/1575 train_time:34549ms step_avg:44.07ms step:785/1575 train_time:34613ms step_avg:44.09ms step:786/1575 train_time:34673ms step_avg:44.11ms step:787/1575 train_time:34736ms step_avg:44.14ms step:788/1575 train_time:34795ms step_avg:44.16ms step:789/1575 train_time:34858ms step_avg:44.18ms step:790/1575 train_time:34917ms step_avg:44.20ms step:791/1575 train_time:34981ms step_avg:44.22ms step:792/1575 train_time:35041ms step_avg:44.24ms step:793/1575 train_time:35104ms step_avg:44.27ms step:794/1575 train_time:35164ms step_avg:44.29ms step:795/1575 train_time:35226ms step_avg:44.31ms step:796/1575 train_time:35285ms step_avg:44.33ms step:797/1575 train_time:35348ms step_avg:44.35ms step:798/1575 train_time:35407ms step_avg:44.37ms step:799/1575 train_time:35471ms step_avg:44.39ms step:800/1575 train_time:35529ms step_avg:44.41ms step:801/1575 train_time:35593ms step_avg:44.44ms step:802/1575 train_time:35652ms step_avg:44.45ms step:803/1575 train_time:35716ms step_avg:44.48ms step:804/1575 train_time:35776ms step_avg:44.50ms step:805/1575 train_time:35839ms step_avg:44.52ms step:806/1575 train_time:35899ms step_avg:44.54ms step:807/1575 train_time:35963ms step_avg:44.56ms step:808/1575 train_time:36022ms step_avg:44.58ms step:809/1575 train_time:36085ms step_avg:44.60ms step:810/1575 train_time:36144ms step_avg:44.62ms step:811/1575 train_time:36208ms step_avg:44.65ms step:812/1575 train_time:36267ms step_avg:44.66ms step:813/1575 train_time:36329ms step_avg:44.69ms step:814/1575 train_time:36390ms step_avg:44.70ms step:815/1575 train_time:36451ms step_avg:44.73ms step:816/1575 train_time:36511ms step_avg:44.74ms step:817/1575 train_time:36573ms step_avg:44.77ms step:818/1575 train_time:36633ms step_avg:44.78ms step:819/1575 train_time:36697ms step_avg:44.81ms step:820/1575 train_time:36756ms step_avg:44.82ms step:821/1575 train_time:36820ms step_avg:44.85ms step:822/1575 train_time:36879ms step_avg:44.87ms step:823/1575 train_time:36943ms step_avg:44.89ms step:824/1575 train_time:37002ms step_avg:44.91ms step:825/1575 train_time:37065ms step_avg:44.93ms step:826/1575 train_time:37124ms step_avg:44.94ms step:827/1575 train_time:37188ms step_avg:44.97ms step:828/1575 train_time:37247ms step_avg:44.98ms step:829/1575 train_time:37310ms step_avg:45.01ms step:830/1575 train_time:37370ms step_avg:45.02ms step:831/1575 train_time:37433ms step_avg:45.05ms step:832/1575 train_time:37493ms step_avg:45.06ms step:833/1575 train_time:37556ms step_avg:45.08ms step:834/1575 train_time:37615ms step_avg:45.10ms step:835/1575 train_time:37679ms step_avg:45.12ms step:836/1575 train_time:37738ms step_avg:45.14ms step:837/1575 train_time:37801ms step_avg:45.16ms step:838/1575 train_time:37860ms step_avg:45.18ms step:839/1575 train_time:37926ms step_avg:45.20ms step:840/1575 train_time:37985ms step_avg:45.22ms step:841/1575 train_time:38048ms step_avg:45.24ms step:842/1575 train_time:38107ms step_avg:45.26ms step:843/1575 train_time:38171ms step_avg:45.28ms step:844/1575 train_time:38229ms step_avg:45.29ms step:845/1575 train_time:38292ms step_avg:45.32ms step:846/1575 train_time:38351ms step_avg:45.33ms step:847/1575 train_time:38414ms step_avg:45.35ms step:848/1575 train_time:38473ms step_avg:45.37ms step:849/1575 train_time:38536ms step_avg:45.39ms step:850/1575 train_time:38596ms step_avg:45.41ms step:851/1575 train_time:38660ms step_avg:45.43ms step:852/1575 train_time:38719ms step_avg:45.44ms step:853/1575 train_time:38782ms step_avg:45.47ms step:854/1575 train_time:38841ms step_avg:45.48ms step:855/1575 train_time:38904ms step_avg:45.50ms step:856/1575 train_time:38964ms step_avg:45.52ms step:857/1575 train_time:39027ms step_avg:45.54ms step:858/1575 train_time:39088ms step_avg:45.56ms step:859/1575 train_time:39150ms step_avg:45.58ms step:860/1575 train_time:39209ms step_avg:45.59ms step:861/1575 train_time:39271ms step_avg:45.61ms step:862/1575 train_time:39332ms step_avg:45.63ms step:863/1575 train_time:39395ms step_avg:45.65ms step:864/1575 train_time:39454ms step_avg:45.66ms step:865/1575 train_time:39524ms step_avg:45.69ms step:866/1575 train_time:39576ms step_avg:45.70ms step:867/1575 train_time:39642ms step_avg:45.72ms step:868/1575 train_time:39700ms step_avg:45.74ms step:869/1575 train_time:39763ms step_avg:45.76ms step:870/1575 train_time:39823ms step_avg:45.77ms step:871/1575 train_time:39887ms step_avg:45.79ms step:872/1575 train_time:39946ms step_avg:45.81ms step:873/1575 train_time:40014ms step_avg:45.84ms step:874/1575 train_time:40069ms step_avg:45.85ms step:875/1575 train_time:40132ms step_avg:45.87ms step:876/1575 train_time:40191ms step_avg:45.88ms step:877/1575 train_time:40254ms step_avg:45.90ms step:878/1575 train_time:40313ms step_avg:45.91ms step:879/1575 train_time:40376ms step_avg:45.93ms step:880/1575 train_time:40435ms step_avg:45.95ms step:881/1575 train_time:40498ms step_avg:45.97ms step:882/1575 train_time:40558ms step_avg:45.98ms step:883/1575 train_time:40621ms step_avg:46.00ms step:884/1575 train_time:40680ms step_avg:46.02ms step:885/1575 train_time:40743ms step_avg:46.04ms step:886/1575 train_time:40808ms step_avg:46.06ms step:887/1575 train_time:40869ms step_avg:46.08ms step:888/1575 train_time:40928ms step_avg:46.09ms step:889/1575 train_time:40992ms step_avg:46.11ms step:890/1575 train_time:41051ms step_avg:46.12ms step:891/1575 train_time:41113ms step_avg:46.14ms step:892/1575 train_time:41173ms step_avg:46.16ms step:893/1575 train_time:41235ms step_avg:46.18ms step:894/1575 train_time:41294ms step_avg:46.19ms step:895/1575 train_time:41357ms step_avg:46.21ms step:896/1575 train_time:41416ms step_avg:46.22ms step:897/1575 train_time:41481ms step_avg:46.24ms step:898/1575 train_time:41540ms step_avg:46.26ms step:899/1575 train_time:41603ms step_avg:46.28ms step:900/1575 train_time:41661ms step_avg:46.29ms step:901/1575 train_time:41724ms step_avg:46.31ms step:902/1575 train_time:41784ms step_avg:46.32ms step:903/1575 train_time:41847ms step_avg:46.34ms step:904/1575 train_time:41907ms step_avg:46.36ms step:905/1575 train_time:41970ms step_avg:46.38ms step:906/1575 train_time:42030ms step_avg:46.39ms step:907/1575 train_time:42092ms step_avg:46.41ms step:908/1575 train_time:42152ms step_avg:46.42ms step:909/1575 train_time:42216ms step_avg:46.44ms step:910/1575 train_time:42274ms step_avg:46.46ms step:911/1575 train_time:42336ms step_avg:46.47ms step:912/1575 train_time:42395ms step_avg:46.49ms step:913/1575 train_time:42459ms step_avg:46.50ms step:914/1575 train_time:42518ms step_avg:46.52ms step:915/1575 train_time:42582ms step_avg:46.54ms step:916/1575 train_time:42641ms step_avg:46.55ms step:917/1575 train_time:42704ms step_avg:46.57ms step:918/1575 train_time:42764ms step_avg:46.58ms step:919/1575 train_time:42828ms step_avg:46.60ms step:920/1575 train_time:42888ms step_avg:46.62ms step:921/1575 train_time:42951ms step_avg:46.64ms step:922/1575 train_time:43011ms step_avg:46.65ms step:923/1575 train_time:43074ms step_avg:46.67ms step:924/1575 train_time:43133ms step_avg:46.68ms step:925/1575 train_time:43196ms step_avg:46.70ms step:926/1575 train_time:43255ms step_avg:46.71ms step:927/1575 train_time:43318ms step_avg:46.73ms step:928/1575 train_time:43378ms step_avg:46.74ms step:929/1575 train_time:43440ms step_avg:46.76ms step:930/1575 train_time:43499ms step_avg:46.77ms step:931/1575 train_time:43563ms step_avg:46.79ms step:932/1575 train_time:43622ms step_avg:46.81ms step:933/1575 train_time:43685ms step_avg:46.82ms step:934/1575 train_time:43745ms step_avg:46.84ms step:935/1575 train_time:43809ms step_avg:46.85ms step:936/1575 train_time:43869ms step_avg:46.87ms step:937/1575 train_time:43932ms step_avg:46.89ms step:938/1575 train_time:43991ms step_avg:46.90ms step:939/1575 train_time:44054ms step_avg:46.92ms step:940/1575 train_time:44113ms step_avg:46.93ms step:941/1575 train_time:44177ms step_avg:46.95ms step:942/1575 train_time:44235ms step_avg:46.96ms step:943/1575 train_time:44299ms step_avg:46.98ms step:944/1575 train_time:44358ms step_avg:46.99ms step:945/1575 train_time:44421ms step_avg:47.01ms step:946/1575 train_time:44480ms step_avg:47.02ms step:947/1575 train_time:44543ms step_avg:47.04ms step:948/1575 train_time:44604ms step_avg:47.05ms step:949/1575 train_time:44666ms step_avg:47.07ms step:950/1575 train_time:44725ms step_avg:47.08ms step:951/1575 train_time:44789ms step_avg:47.10ms step:952/1575 train_time:44848ms step_avg:47.11ms step:953/1575 train_time:44911ms step_avg:47.13ms step:954/1575 train_time:44971ms step_avg:47.14ms step:955/1575 train_time:45033ms step_avg:47.15ms step:956/1575 train_time:45092ms step_avg:47.17ms step:957/1575 train_time:45155ms step_avg:47.18ms step:958/1575 train_time:45214ms step_avg:47.20ms step:959/1575 train_time:45277ms step_avg:47.21ms step:960/1575 train_time:45336ms step_avg:47.23ms step:961/1575 train_time:45403ms step_avg:47.25ms step:962/1575 train_time:45461ms step_avg:47.26ms step:963/1575 train_time:45523ms step_avg:47.27ms step:964/1575 train_time:45583ms step_avg:47.29ms step:965/1575 train_time:45646ms step_avg:47.30ms step:966/1575 train_time:45705ms step_avg:47.31ms step:967/1575 train_time:45770ms step_avg:47.33ms step:968/1575 train_time:45829ms step_avg:47.34ms step:969/1575 train_time:45892ms step_avg:47.36ms step:970/1575 train_time:45953ms step_avg:47.37ms step:971/1575 train_time:46015ms step_avg:47.39ms step:972/1575 train_time:46074ms step_avg:47.40ms step:973/1575 train_time:46138ms step_avg:47.42ms step:974/1575 train_time:46196ms step_avg:47.43ms step:975/1575 train_time:46261ms step_avg:47.45ms step:976/1575 train_time:46319ms step_avg:47.46ms step:977/1575 train_time:46382ms step_avg:47.47ms step:978/1575 train_time:46441ms step_avg:47.49ms step:979/1575 train_time:46506ms step_avg:47.50ms step:980/1575 train_time:46565ms step_avg:47.52ms step:981/1575 train_time:46630ms step_avg:47.53ms step:982/1575 train_time:46689ms step_avg:47.54ms step:983/1575 train_time:46751ms step_avg:47.56ms step:984/1575 train_time:46810ms step_avg:47.57ms step:985/1575 train_time:46874ms step_avg:47.59ms step:986/1575 train_time:46933ms step_avg:47.60ms step:987/1575 train_time:46996ms step_avg:47.62ms step:988/1575 train_time:47057ms step_avg:47.63ms step:989/1575 train_time:47120ms step_avg:47.64ms step:990/1575 train_time:47180ms step_avg:47.66ms step:991/1575 train_time:47241ms step_avg:47.67ms step:992/1575 train_time:47301ms step_avg:47.68ms step:993/1575 train_time:47364ms step_avg:47.70ms step:994/1575 train_time:47423ms step_avg:47.71ms step:995/1575 train_time:47487ms step_avg:47.73ms step:996/1575 train_time:47546ms step_avg:47.74ms step:997/1575 train_time:47610ms step_avg:47.75ms step:998/1575 train_time:47669ms step_avg:47.76ms step:999/1575 train_time:47732ms step_avg:47.78ms step:1000/1575 train_time:47791ms step_avg:47.79ms step:1000/1575 val_loss:3.5839 train_time:47837ms step_avg:47.84ms step:1001/1575 train_time:47858ms step_avg:47.81ms step:1002/1575 train_time:47917ms step_avg:47.82ms step:1003/1575 train_time:47983ms step_avg:47.84ms step:1004/1575 train_time:48044ms step_avg:47.85ms step:1005/1575 train_time:48107ms step_avg:47.87ms step:1006/1575 train_time:48167ms step_avg:47.88ms step:1007/1575 train_time:48229ms step_avg:47.89ms step:1008/1575 train_time:48289ms step_avg:47.91ms step:1009/1575 train_time:48351ms step_avg:47.92ms step:1010/1575 train_time:48409ms step_avg:47.93ms step:1011/1575 train_time:48473ms step_avg:47.95ms step:1012/1575 train_time:48532ms step_avg:47.96ms step:1013/1575 train_time:48594ms step_avg:47.97ms step:1014/1575 train_time:48654ms step_avg:47.98ms step:1015/1575 train_time:48717ms step_avg:48.00ms step:1016/1575 train_time:48776ms step_avg:48.01ms step:1017/1575 train_time:48841ms step_avg:48.02ms step:1018/1575 train_time:48901ms step_avg:48.04ms step:1019/1575 train_time:48966ms step_avg:48.05ms step:1020/1575 train_time:49026ms step_avg:48.06ms step:1021/1575 train_time:49089ms step_avg:48.08ms step:1022/1575 train_time:49151ms step_avg:48.09ms step:1023/1575 train_time:49212ms step_avg:48.11ms step:1024/1575 train_time:49270ms step_avg:48.12ms step:1025/1575 train_time:49344ms step_avg:48.14ms step:1026/1575 train_time:49427ms step_avg:48.17ms step:1027/1575 train_time:49515ms step_avg:48.21ms step:1028/1575 train_time:49599ms step_avg:48.25ms step:1029/1575 train_time:49687ms step_avg:48.29ms step:1030/1575 train_time:49773ms step_avg:48.32ms step:1031/1575 train_time:49864ms step_avg:48.36ms step:1032/1575 train_time:49951ms step_avg:48.40ms step:1033/1575 train_time:50041ms step_avg:48.44ms step:1034/1575 train_time:50127ms step_avg:48.48ms step:1035/1575 train_time:50215ms step_avg:48.52ms step:1036/1575 train_time:50301ms step_avg:48.55ms step:1037/1575 train_time:50390ms step_avg:48.59ms step:1038/1575 train_time:50475ms step_avg:48.63ms step:1039/1575 train_time:50564ms step_avg:48.67ms step:1040/1575 train_time:50649ms step_avg:48.70ms step:1041/1575 train_time:50738ms step_avg:48.74ms step:1042/1575 train_time:50826ms step_avg:48.78ms step:1043/1575 train_time:50916ms step_avg:48.82ms step:1044/1575 train_time:51001ms step_avg:48.85ms step:1045/1575 train_time:51093ms step_avg:48.89ms step:1046/1575 train_time:51179ms step_avg:48.93ms step:1047/1575 train_time:51267ms step_avg:48.97ms step:1048/1575 train_time:51353ms step_avg:49.00ms step:1049/1575 train_time:51442ms step_avg:49.04ms step:1050/1575 train_time:51527ms step_avg:49.07ms step:1051/1575 train_time:51615ms step_avg:49.11ms step:1052/1575 train_time:51701ms step_avg:49.15ms step:1053/1575 train_time:51792ms step_avg:49.19ms step:1054/1575 train_time:51877ms step_avg:49.22ms step:1055/1575 train_time:51967ms step_avg:49.26ms step:1056/1575 train_time:52053ms step_avg:49.29ms step:1057/1575 train_time:52143ms step_avg:49.33ms step:1058/1575 train_time:52229ms step_avg:49.37ms step:1059/1575 train_time:52318ms step_avg:49.40ms step:1060/1575 train_time:52403ms step_avg:49.44ms step:1061/1575 train_time:52492ms step_avg:49.47ms step:1062/1575 train_time:52577ms step_avg:49.51ms step:1063/1575 train_time:52667ms step_avg:49.55ms step:1064/1575 train_time:52753ms step_avg:49.58ms step:1065/1575 train_time:52843ms step_avg:49.62ms step:1066/1575 train_time:52929ms step_avg:49.65ms step:1067/1575 train_time:53018ms step_avg:49.69ms step:1068/1575 train_time:53105ms step_avg:49.72ms step:1069/1575 train_time:53195ms step_avg:49.76ms step:1070/1575 train_time:53280ms step_avg:49.79ms step:1071/1575 train_time:53369ms step_avg:49.83ms step:1072/1575 train_time:53455ms step_avg:49.87ms step:1073/1575 train_time:53544ms step_avg:49.90ms step:1074/1575 train_time:53629ms step_avg:49.93ms step:1075/1575 train_time:53717ms step_avg:49.97ms step:1076/1575 train_time:53803ms step_avg:50.00ms step:1077/1575 train_time:53893ms step_avg:50.04ms step:1078/1575 train_time:53979ms step_avg:50.07ms step:1079/1575 train_time:54069ms step_avg:50.11ms step:1080/1575 train_time:54154ms step_avg:50.14ms step:1081/1575 train_time:54243ms step_avg:50.18ms step:1082/1575 train_time:54329ms step_avg:50.21ms step:1083/1575 train_time:54417ms step_avg:50.25ms step:1084/1575 train_time:54503ms step_avg:50.28ms step:1085/1575 train_time:54593ms step_avg:50.32ms step:1086/1575 train_time:54679ms step_avg:50.35ms step:1087/1575 train_time:54768ms step_avg:50.38ms step:1088/1575 train_time:54853ms step_avg:50.42ms step:1089/1575 train_time:54943ms step_avg:50.45ms step:1090/1575 train_time:55029ms step_avg:50.49ms step:1091/1575 train_time:55118ms step_avg:50.52ms step:1092/1575 train_time:55204ms step_avg:50.55ms step:1093/1575 train_time:55293ms step_avg:50.59ms step:1094/1575 train_time:55378ms step_avg:50.62ms step:1095/1575 train_time:55468ms step_avg:50.66ms step:1096/1575 train_time:55552ms step_avg:50.69ms step:1097/1575 train_time:55642ms step_avg:50.72ms step:1098/1575 train_time:55728ms step_avg:50.75ms step:1099/1575 train_time:55817ms step_avg:50.79ms step:1100/1575 train_time:55903ms step_avg:50.82ms step:1101/1575 train_time:55994ms step_avg:50.86ms step:1102/1575 train_time:56079ms step_avg:50.89ms step:1103/1575 train_time:56168ms step_avg:50.92ms step:1104/1575 train_time:56254ms step_avg:50.95ms step:1105/1575 train_time:56345ms step_avg:50.99ms step:1106/1575 train_time:56430ms step_avg:51.02ms step:1107/1575 train_time:56518ms step_avg:51.06ms step:1108/1575 train_time:56604ms step_avg:51.09ms step:1109/1575 train_time:56694ms step_avg:51.12ms step:1110/1575 train_time:56780ms step_avg:51.15ms step:1111/1575 train_time:56868ms step_avg:51.19ms step:1112/1575 train_time:56954ms step_avg:51.22ms step:1113/1575 train_time:57043ms step_avg:51.25ms step:1114/1575 train_time:57129ms step_avg:51.28ms step:1115/1575 train_time:57219ms step_avg:51.32ms step:1116/1575 train_time:57304ms step_avg:51.35ms step:1117/1575 train_time:57394ms step_avg:51.38ms step:1118/1575 train_time:57479ms step_avg:51.41ms step:1119/1575 train_time:57570ms step_avg:51.45ms step:1120/1575 train_time:57659ms step_avg:51.48ms step:1121/1575 train_time:57747ms step_avg:51.51ms step:1122/1575 train_time:57835ms step_avg:51.55ms step:1123/1575 train_time:57921ms step_avg:51.58ms step:1124/1575 train_time:58006ms step_avg:51.61ms step:1125/1575 train_time:58097ms step_avg:51.64ms step:1126/1575 train_time:58181ms step_avg:51.67ms step:1127/1575 train_time:58270ms step_avg:51.70ms step:1128/1575 train_time:58356ms step_avg:51.73ms step:1129/1575 train_time:58445ms step_avg:51.77ms step:1130/1575 train_time:58530ms step_avg:51.80ms step:1131/1575 train_time:58619ms step_avg:51.83ms step:1132/1575 train_time:58705ms step_avg:51.86ms step:1133/1575 train_time:58794ms step_avg:51.89ms step:1134/1575 train_time:58880ms step_avg:51.92ms step:1135/1575 train_time:58969ms step_avg:51.96ms step:1136/1575 train_time:59055ms step_avg:51.99ms step:1137/1575 train_time:59144ms step_avg:52.02ms step:1138/1575 train_time:59229ms step_avg:52.05ms step:1139/1575 train_time:59318ms step_avg:52.08ms step:1140/1575 train_time:59404ms step_avg:52.11ms step:1141/1575 train_time:59493ms step_avg:52.14ms step:1142/1575 train_time:59578ms step_avg:52.17ms step:1143/1575 train_time:59668ms step_avg:52.20ms step:1144/1575 train_time:59753ms step_avg:52.23ms step:1145/1575 train_time:59843ms step_avg:52.26ms step:1146/1575 train_time:59935ms step_avg:52.30ms step:1147/1575 train_time:60020ms step_avg:52.33ms step:1148/1575 train_time:60107ms step_avg:52.36ms step:1149/1575 train_time:60197ms step_avg:52.39ms step:1150/1575 train_time:60282ms step_avg:52.42ms step:1151/1575 train_time:60372ms step_avg:52.45ms step:1152/1575 train_time:60457ms step_avg:52.48ms step:1153/1575 train_time:60547ms step_avg:52.51ms step:1154/1575 train_time:60632ms step_avg:52.54ms step:1155/1575 train_time:60721ms step_avg:52.57ms step:1156/1575 train_time:60807ms step_avg:52.60ms step:1157/1575 train_time:60897ms step_avg:52.63ms step:1158/1575 train_time:60983ms step_avg:52.66ms step:1159/1575 train_time:61076ms step_avg:52.70ms step:1160/1575 train_time:61161ms step_avg:52.73ms step:1161/1575 train_time:61249ms step_avg:52.76ms step:1162/1575 train_time:61334ms step_avg:52.78ms step:1163/1575 train_time:61423ms step_avg:52.81ms step:1164/1575 train_time:61509ms step_avg:52.84ms step:1165/1575 train_time:61598ms step_avg:52.87ms step:1166/1575 train_time:61684ms step_avg:52.90ms step:1167/1575 train_time:61775ms step_avg:52.93ms step:1168/1575 train_time:61862ms step_avg:52.96ms step:1169/1575 train_time:61952ms step_avg:53.00ms step:1170/1575 train_time:62038ms step_avg:53.02ms step:1171/1575 train_time:62128ms step_avg:53.06ms step:1172/1575 train_time:62213ms step_avg:53.08ms step:1173/1575 train_time:62301ms step_avg:53.11ms step:1174/1575 train_time:62387ms step_avg:53.14ms step:1175/1575 train_time:62476ms step_avg:53.17ms step:1176/1575 train_time:62567ms step_avg:53.20ms step:1177/1575 train_time:62652ms step_avg:53.23ms step:1178/1575 train_time:62738ms step_avg:53.26ms step:1179/1575 train_time:62828ms step_avg:53.29ms step:1180/1575 train_time:62913ms step_avg:53.32ms step:1181/1575 train_time:63002ms step_avg:53.35ms step:1182/1575 train_time:63088ms step_avg:53.37ms step:1183/1575 train_time:63177ms step_avg:53.40ms step:1184/1575 train_time:63262ms step_avg:53.43ms step:1185/1575 train_time:63352ms step_avg:53.46ms step:1186/1575 train_time:63437ms step_avg:53.49ms step:1187/1575 train_time:63527ms step_avg:53.52ms step:1188/1575 train_time:63613ms step_avg:53.55ms step:1189/1575 train_time:63703ms step_avg:53.58ms step:1190/1575 train_time:63788ms step_avg:53.60ms step:1191/1575 train_time:63878ms step_avg:53.63ms step:1192/1575 train_time:63964ms step_avg:53.66ms step:1193/1575 train_time:64053ms step_avg:53.69ms step:1194/1575 train_time:64139ms step_avg:53.72ms step:1195/1575 train_time:64228ms step_avg:53.75ms step:1196/1575 train_time:64313ms step_avg:53.77ms step:1197/1575 train_time:64403ms step_avg:53.80ms step:1198/1575 train_time:64488ms step_avg:53.83ms step:1199/1575 train_time:64578ms step_avg:53.86ms step:1200/1575 train_time:64663ms step_avg:53.89ms step:1201/1575 train_time:64753ms step_avg:53.92ms step:1202/1575 train_time:64840ms step_avg:53.94ms step:1203/1575 train_time:64929ms step_avg:53.97ms step:1204/1575 train_time:65014ms step_avg:54.00ms step:1205/1575 train_time:65104ms step_avg:54.03ms step:1206/1575 train_time:65189ms step_avg:54.05ms step:1207/1575 train_time:65278ms step_avg:54.08ms step:1208/1575 train_time:65366ms step_avg:54.11ms step:1209/1575 train_time:65455ms step_avg:54.14ms step:1210/1575 train_time:65540ms step_avg:54.17ms step:1211/1575 train_time:65629ms step_avg:54.19ms step:1212/1575 train_time:65714ms step_avg:54.22ms step:1213/1575 train_time:65803ms step_avg:54.25ms step:1214/1575 train_time:65889ms step_avg:54.27ms step:1215/1575 train_time:65978ms step_avg:54.30ms step:1216/1575 train_time:66064ms step_avg:54.33ms step:1217/1575 train_time:66154ms step_avg:54.36ms step:1218/1575 train_time:66241ms step_avg:54.39ms step:1219/1575 train_time:66330ms step_avg:54.41ms step:1220/1575 train_time:66414ms step_avg:54.44ms step:1221/1575 train_time:66503ms step_avg:54.47ms step:1222/1575 train_time:66589ms step_avg:54.49ms step:1223/1575 train_time:66678ms step_avg:54.52ms step:1224/1575 train_time:66763ms step_avg:54.54ms step:1225/1575 train_time:66852ms step_avg:54.57ms step:1226/1575 train_time:66938ms step_avg:54.60ms step:1227/1575 train_time:67028ms step_avg:54.63ms step:1228/1575 train_time:67119ms step_avg:54.66ms step:1229/1575 train_time:67203ms step_avg:54.68ms step:1230/1575 train_time:67288ms step_avg:54.71ms step:1231/1575 train_time:67379ms step_avg:54.73ms step:1232/1575 train_time:67464ms step_avg:54.76ms step:1233/1575 train_time:67554ms step_avg:54.79ms step:1234/1575 train_time:67639ms step_avg:54.81ms step:1235/1575 train_time:67729ms step_avg:54.84ms step:1236/1575 train_time:67814ms step_avg:54.87ms step:1237/1575 train_time:67904ms step_avg:54.89ms step:1238/1575 train_time:67990ms step_avg:54.92ms step:1239/1575 train_time:68078ms step_avg:54.95ms step:1240/1575 train_time:68165ms step_avg:54.97ms step:1241/1575 train_time:68255ms step_avg:55.00ms step:1242/1575 train_time:68342ms step_avg:55.03ms step:1243/1575 train_time:68431ms step_avg:55.05ms step:1244/1575 train_time:68517ms step_avg:55.08ms step:1245/1575 train_time:68606ms step_avg:55.11ms step:1246/1575 train_time:68691ms step_avg:55.13ms step:1247/1575 train_time:68782ms step_avg:55.16ms step:1248/1575 train_time:68868ms step_avg:55.18ms step:1249/1575 train_time:68957ms step_avg:55.21ms step:1250/1575 train_time:69043ms step_avg:55.23ms step:1250/1575 val_loss:3.4077 train_time:69115ms step_avg:55.29ms step:1251/1575 train_time:69136ms step_avg:55.26ms step:1252/1575 train_time:69226ms step_avg:55.29ms step:1253/1575 train_time:69318ms step_avg:55.32ms step:1254/1575 train_time:69405ms step_avg:55.35ms step:1255/1575 train_time:69494ms step_avg:55.37ms step:1256/1575 train_time:69579ms step_avg:55.40ms step:1257/1575 train_time:69667ms step_avg:55.42ms step:1258/1575 train_time:69752ms step_avg:55.45ms step:1259/1575 train_time:69840ms step_avg:55.47ms step:1260/1575 train_time:69925ms step_avg:55.50ms step:1261/1575 train_time:70013ms step_avg:55.52ms step:1262/1575 train_time:70099ms step_avg:55.55ms step:1263/1575 train_time:70191ms step_avg:55.58ms step:1264/1575 train_time:70278ms step_avg:55.60ms step:1265/1575 train_time:70369ms step_avg:55.63ms step:1266/1575 train_time:70457ms step_avg:55.65ms step:1267/1575 train_time:70544ms step_avg:55.68ms step:1268/1575 train_time:70630ms step_avg:55.70ms step:1269/1575 train_time:70717ms step_avg:55.73ms step:1270/1575 train_time:70802ms step_avg:55.75ms step:1271/1575 train_time:70890ms step_avg:55.78ms step:1272/1575 train_time:70975ms step_avg:55.80ms step:1273/1575 train_time:71066ms step_avg:55.83ms step:1274/1575 train_time:71152ms step_avg:55.85ms step:1275/1575 train_time:71242ms step_avg:55.88ms step:1276/1575 train_time:71329ms step_avg:55.90ms step:1277/1575 train_time:71418ms step_avg:55.93ms step:1278/1575 train_time:71505ms step_avg:55.95ms step:1279/1575 train_time:71594ms step_avg:55.98ms step:1280/1575 train_time:71679ms step_avg:56.00ms step:1281/1575 train_time:71769ms step_avg:56.03ms step:1282/1575 train_time:71854ms step_avg:56.05ms step:1283/1575 train_time:71943ms step_avg:56.07ms step:1284/1575 train_time:72029ms step_avg:56.10ms step:1285/1575 train_time:72118ms step_avg:56.12ms step:1286/1575 train_time:72206ms step_avg:56.15ms step:1287/1575 train_time:72295ms step_avg:56.17ms step:1288/1575 train_time:72381ms step_avg:56.20ms step:1289/1575 train_time:72471ms step_avg:56.22ms step:1290/1575 train_time:72556ms step_avg:56.25ms step:1291/1575 train_time:72646ms step_avg:56.27ms step:1292/1575 train_time:72731ms step_avg:56.29ms step:1293/1575 train_time:72819ms step_avg:56.32ms step:1294/1575 train_time:72905ms step_avg:56.34ms step:1295/1575 train_time:72994ms step_avg:56.37ms step:1296/1575 train_time:73080ms step_avg:56.39ms step:1297/1575 train_time:73170ms step_avg:56.41ms step:1298/1575 train_time:73255ms step_avg:56.44ms step:1299/1575 train_time:73345ms step_avg:56.46ms step:1300/1575 train_time:73431ms step_avg:56.49ms step:1301/1575 train_time:73520ms step_avg:56.51ms step:1302/1575 train_time:73606ms step_avg:56.53ms step:1303/1575 train_time:73695ms step_avg:56.56ms step:1304/1575 train_time:73781ms step_avg:56.58ms step:1305/1575 train_time:73870ms step_avg:56.61ms step:1306/1575 train_time:73955ms step_avg:56.63ms step:1307/1575 train_time:74045ms step_avg:56.65ms step:1308/1575 train_time:74130ms step_avg:56.67ms step:1309/1575 train_time:74220ms step_avg:56.70ms step:1310/1575 train_time:74307ms step_avg:56.72ms step:1311/1575 train_time:74397ms step_avg:56.75ms step:1312/1575 train_time:74483ms step_avg:56.77ms step:1313/1575 train_time:74573ms step_avg:56.80ms step:1314/1575 train_time:74659ms step_avg:56.82ms step:1315/1575 train_time:74748ms step_avg:56.84ms step:1316/1575 train_time:74833ms step_avg:56.86ms step:1317/1575 train_time:74922ms step_avg:56.89ms step:1318/1575 train_time:75008ms step_avg:56.91ms step:1319/1575 train_time:75097ms step_avg:56.93ms step:1320/1575 train_time:75183ms step_avg:56.96ms step:1321/1575 train_time:75275ms step_avg:56.98ms step:1322/1575 train_time:75360ms step_avg:57.00ms step:1323/1575 train_time:75450ms step_avg:57.03ms step:1324/1575 train_time:75535ms step_avg:57.05ms step:1325/1575 train_time:75626ms step_avg:57.08ms step:1326/1575 train_time:75711ms step_avg:57.10ms step:1327/1575 train_time:75799ms step_avg:57.12ms step:1328/1575 train_time:75885ms step_avg:57.14ms step:1329/1575 train_time:75975ms step_avg:57.17ms step:1330/1575 train_time:76060ms step_avg:57.19ms step:1331/1575 train_time:76149ms step_avg:57.21ms step:1332/1575 train_time:76234ms step_avg:57.23ms step:1333/1575 train_time:76324ms step_avg:57.26ms step:1334/1575 train_time:76411ms step_avg:57.28ms step:1335/1575 train_time:76499ms step_avg:57.30ms step:1336/1575 train_time:76585ms step_avg:57.32ms step:1337/1575 train_time:76675ms step_avg:57.35ms step:1338/1575 train_time:76759ms step_avg:57.37ms step:1339/1575 train_time:76849ms step_avg:57.39ms step:1340/1575 train_time:76935ms step_avg:57.41ms step:1341/1575 train_time:77024ms step_avg:57.44ms step:1342/1575 train_time:77110ms step_avg:57.46ms step:1343/1575 train_time:77199ms step_avg:57.48ms step:1344/1575 train_time:77285ms step_avg:57.50ms step:1345/1575 train_time:77375ms step_avg:57.53ms step:1346/1575 train_time:77461ms step_avg:57.55ms step:1347/1575 train_time:77551ms step_avg:57.57ms step:1348/1575 train_time:77636ms step_avg:57.59ms step:1349/1575 train_time:77725ms step_avg:57.62ms step:1350/1575 train_time:77810ms step_avg:57.64ms step:1351/1575 train_time:77900ms step_avg:57.66ms step:1352/1575 train_time:77985ms step_avg:57.68ms step:1353/1575 train_time:78075ms step_avg:57.71ms step:1354/1575 train_time:78161ms step_avg:57.73ms step:1355/1575 train_time:78250ms step_avg:57.75ms step:1356/1575 train_time:78337ms step_avg:57.77ms step:1357/1575 train_time:78426ms step_avg:57.79ms step:1358/1575 train_time:78516ms step_avg:57.82ms step:1359/1575 train_time:78603ms step_avg:57.84ms step:1360/1575 train_time:78689ms step_avg:57.86ms step:1361/1575 train_time:78778ms step_avg:57.88ms step:1362/1575 train_time:78863ms step_avg:57.90ms step:1363/1575 train_time:78953ms step_avg:57.93ms step:1364/1575 train_time:79040ms step_avg:57.95ms step:1365/1575 train_time:79128ms step_avg:57.97ms step:1366/1575 train_time:79213ms step_avg:57.99ms step:1367/1575 train_time:79303ms step_avg:58.01ms step:1368/1575 train_time:79390ms step_avg:58.03ms step:1369/1575 train_time:79479ms step_avg:58.06ms step:1370/1575 train_time:79564ms step_avg:58.08ms step:1371/1575 train_time:79654ms step_avg:58.10ms step:1372/1575 train_time:79740ms step_avg:58.12ms step:1373/1575 train_time:79829ms step_avg:58.14ms step:1374/1575 train_time:79915ms step_avg:58.16ms step:1375/1575 train_time:80003ms step_avg:58.18ms step:1376/1575 train_time:80089ms step_avg:58.20ms step:1377/1575 train_time:80178ms step_avg:58.23ms step:1378/1575 train_time:80264ms step_avg:58.25ms step:1379/1575 train_time:80355ms step_avg:58.27ms step:1380/1575 train_time:80440ms step_avg:58.29ms step:1381/1575 train_time:80530ms step_avg:58.31ms step:1382/1575 train_time:80616ms step_avg:58.33ms step:1383/1575 train_time:80707ms step_avg:58.36ms step:1384/1575 train_time:80791ms step_avg:58.38ms step:1385/1575 train_time:80880ms step_avg:58.40ms step:1386/1575 train_time:80966ms step_avg:58.42ms step:1387/1575 train_time:81056ms step_avg:58.44ms step:1388/1575 train_time:81141ms step_avg:58.46ms step:1389/1575 train_time:81231ms step_avg:58.48ms step:1390/1575 train_time:81316ms step_avg:58.50ms step:1391/1575 train_time:81406ms step_avg:58.52ms step:1392/1575 train_time:81491ms step_avg:58.54ms step:1393/1575 train_time:81580ms step_avg:58.56ms step:1394/1575 train_time:81666ms step_avg:58.58ms step:1395/1575 train_time:81756ms step_avg:58.61ms step:1396/1575 train_time:81842ms step_avg:58.63ms step:1397/1575 train_time:81932ms step_avg:58.65ms step:1398/1575 train_time:82017ms step_avg:58.67ms step:1399/1575 train_time:82106ms step_avg:58.69ms step:1400/1575 train_time:82191ms step_avg:58.71ms step:1401/1575 train_time:82281ms step_avg:58.73ms step:1402/1575 train_time:82367ms step_avg:58.75ms step:1403/1575 train_time:82456ms step_avg:58.77ms step:1404/1575 train_time:82542ms step_avg:58.79ms step:1405/1575 train_time:82632ms step_avg:58.81ms step:1406/1575 train_time:82717ms step_avg:58.83ms step:1407/1575 train_time:82807ms step_avg:58.85ms step:1408/1575 train_time:82894ms step_avg:58.87ms step:1409/1575 train_time:82983ms step_avg:58.89ms step:1410/1575 train_time:83069ms step_avg:58.91ms step:1411/1575 train_time:83158ms step_avg:58.94ms step:1412/1575 train_time:83243ms step_avg:58.95ms step:1413/1575 train_time:83333ms step_avg:58.98ms step:1414/1575 train_time:83420ms step_avg:59.00ms step:1415/1575 train_time:83508ms step_avg:59.02ms step:1416/1575 train_time:83594ms step_avg:59.04ms step:1417/1575 train_time:83684ms step_avg:59.06ms step:1418/1575 train_time:83770ms step_avg:59.08ms step:1419/1575 train_time:83859ms step_avg:59.10ms step:1420/1575 train_time:83945ms step_avg:59.12ms step:1421/1575 train_time:84036ms step_avg:59.14ms step:1422/1575 train_time:84120ms step_avg:59.16ms step:1423/1575 train_time:84211ms step_avg:59.18ms step:1424/1575 train_time:84297ms step_avg:59.20ms step:1425/1575 train_time:84387ms step_avg:59.22ms step:1426/1575 train_time:84472ms step_avg:59.24ms step:1427/1575 train_time:84561ms step_avg:59.26ms step:1428/1575 train_time:84646ms step_avg:59.28ms step:1429/1575 train_time:84736ms step_avg:59.30ms step:1430/1575 train_time:84822ms step_avg:59.32ms step:1431/1575 train_time:84912ms step_avg:59.34ms step:1432/1575 train_time:84998ms step_avg:59.36ms step:1433/1575 train_time:85087ms step_avg:59.38ms step:1434/1575 train_time:85172ms step_avg:59.39ms step:1435/1575 train_time:85262ms step_avg:59.42ms step:1436/1575 train_time:85349ms step_avg:59.43ms step:1437/1575 train_time:85440ms step_avg:59.46ms step:1438/1575 train_time:85524ms step_avg:59.47ms step:1439/1575 train_time:85615ms step_avg:59.50ms step:1440/1575 train_time:85699ms step_avg:59.51ms step:1441/1575 train_time:85790ms step_avg:59.54ms step:1442/1575 train_time:85876ms step_avg:59.55ms step:1443/1575 train_time:85964ms step_avg:59.57ms step:1444/1575 train_time:86050ms step_avg:59.59ms step:1445/1575 train_time:86139ms step_avg:59.61ms step:1446/1575 train_time:86225ms step_avg:59.63ms step:1447/1575 train_time:86315ms step_avg:59.65ms step:1448/1575 train_time:86401ms step_avg:59.67ms step:1449/1575 train_time:86490ms step_avg:59.69ms step:1450/1575 train_time:86576ms step_avg:59.71ms step:1451/1575 train_time:86665ms step_avg:59.73ms step:1452/1575 train_time:86750ms step_avg:59.75ms step:1453/1575 train_time:86839ms step_avg:59.77ms step:1454/1575 train_time:86926ms step_avg:59.78ms step:1455/1575 train_time:87015ms step_avg:59.80ms step:1456/1575 train_time:87100ms step_avg:59.82ms step:1457/1575 train_time:87189ms step_avg:59.84ms step:1458/1575 train_time:87275ms step_avg:59.86ms step:1459/1575 train_time:87365ms step_avg:59.88ms step:1460/1575 train_time:87451ms step_avg:59.90ms step:1461/1575 train_time:87540ms step_avg:59.92ms step:1462/1575 train_time:87627ms step_avg:59.94ms step:1463/1575 train_time:87716ms step_avg:59.96ms step:1464/1575 train_time:87802ms step_avg:59.97ms step:1465/1575 train_time:87892ms step_avg:59.99ms step:1466/1575 train_time:87978ms step_avg:60.01ms step:1467/1575 train_time:88067ms step_avg:60.03ms step:1468/1575 train_time:88152ms step_avg:60.05ms step:1469/1575 train_time:88242ms step_avg:60.07ms step:1470/1575 train_time:88328ms step_avg:60.09ms step:1471/1575 train_time:88417ms step_avg:60.11ms step:1472/1575 train_time:88505ms step_avg:60.13ms step:1473/1575 train_time:88593ms step_avg:60.14ms step:1474/1575 train_time:88679ms step_avg:60.16ms step:1475/1575 train_time:88768ms step_avg:60.18ms step:1476/1575 train_time:88853ms step_avg:60.20ms step:1477/1575 train_time:88943ms step_avg:60.22ms step:1478/1575 train_time:89029ms step_avg:60.24ms step:1479/1575 train_time:89118ms step_avg:60.26ms step:1480/1575 train_time:89203ms step_avg:60.27ms step:1481/1575 train_time:89293ms step_avg:60.29ms step:1482/1575 train_time:89379ms step_avg:60.31ms step:1483/1575 train_time:89468ms step_avg:60.33ms step:1484/1575 train_time:89554ms step_avg:60.35ms step:1485/1575 train_time:89644ms step_avg:60.37ms step:1486/1575 train_time:89729ms step_avg:60.38ms step:1487/1575 train_time:89819ms step_avg:60.40ms step:1488/1575 train_time:89905ms step_avg:60.42ms step:1489/1575 train_time:89994ms step_avg:60.44ms step:1490/1575 train_time:90080ms step_avg:60.46ms step:1491/1575 train_time:90170ms step_avg:60.48ms step:1492/1575 train_time:90255ms step_avg:60.49ms step:1493/1575 train_time:90346ms step_avg:60.51ms step:1494/1575 train_time:90432ms step_avg:60.53ms step:1495/1575 train_time:90519ms step_avg:60.55ms step:1496/1575 train_time:90605ms step_avg:60.56ms step:1497/1575 train_time:90695ms step_avg:60.58ms step:1498/1575 train_time:90781ms step_avg:60.60ms step:1499/1575 train_time:90871ms step_avg:60.62ms step:1500/1575 train_time:90957ms step_avg:60.64ms step:1500/1575 val_loss:3.3004 train_time:91028ms step_avg:60.69ms step:1501/1575 train_time:91050ms step_avg:60.66ms step:1502/1575 train_time:91137ms step_avg:60.68ms step:1503/1575 train_time:91231ms step_avg:60.70ms step:1504/1575 train_time:91318ms step_avg:60.72ms step:1505/1575 train_time:91408ms step_avg:60.74ms step:1506/1575 train_time:91493ms step_avg:60.75ms step:1507/1575 train_time:91581ms step_avg:60.77ms step:1508/1575 train_time:91666ms step_avg:60.79ms step:1509/1575 train_time:91754ms step_avg:60.80ms step:1510/1575 train_time:91840ms step_avg:60.82ms step:1511/1575 train_time:91928ms step_avg:60.84ms step:1512/1575 train_time:92014ms step_avg:60.86ms step:1513/1575 train_time:92105ms step_avg:60.88ms step:1514/1575 train_time:92194ms step_avg:60.89ms step:1515/1575 train_time:92284ms step_avg:60.91ms step:1516/1575 train_time:92371ms step_avg:60.93ms step:1517/1575 train_time:92460ms step_avg:60.95ms step:1518/1575 train_time:92546ms step_avg:60.97ms step:1519/1575 train_time:92634ms step_avg:60.98ms step:1520/1575 train_time:92718ms step_avg:61.00ms step:1521/1575 train_time:92807ms step_avg:61.02ms step:1522/1575 train_time:92893ms step_avg:61.03ms step:1523/1575 train_time:92982ms step_avg:61.05ms step:1524/1575 train_time:93070ms step_avg:61.07ms step:1525/1575 train_time:93162ms step_avg:61.09ms step:1526/1575 train_time:93248ms step_avg:61.11ms step:1527/1575 train_time:93342ms step_avg:61.13ms step:1528/1575 train_time:93424ms step_avg:61.14ms step:1529/1575 train_time:93514ms step_avg:61.16ms step:1530/1575 train_time:93599ms step_avg:61.18ms step:1531/1575 train_time:93688ms step_avg:61.19ms step:1532/1575 train_time:93773ms step_avg:61.21ms step:1533/1575 train_time:93861ms step_avg:61.23ms step:1534/1575 train_time:93946ms step_avg:61.24ms step:1535/1575 train_time:94039ms step_avg:61.26ms step:1536/1575 train_time:94133ms step_avg:61.28ms step:1537/1575 train_time:94220ms step_avg:61.30ms step:1538/1575 train_time:94307ms step_avg:61.32ms step:1539/1575 train_time:94397ms step_avg:61.34ms step:1540/1575 train_time:94483ms step_avg:61.35ms step:1541/1575 train_time:94573ms step_avg:61.37ms step:1542/1575 train_time:94658ms step_avg:61.39ms step:1543/1575 train_time:94748ms step_avg:61.41ms step:1544/1575 train_time:94834ms step_avg:61.42ms step:1545/1575 train_time:94923ms step_avg:61.44ms step:1546/1575 train_time:95009ms step_avg:61.45ms step:1547/1575 train_time:95099ms step_avg:61.47ms step:1548/1575 train_time:95187ms step_avg:61.49ms step:1549/1575 train_time:95278ms step_avg:61.51ms step:1550/1575 train_time:95364ms step_avg:61.53ms step:1551/1575 train_time:95454ms step_avg:61.54ms step:1552/1575 train_time:95539ms step_avg:61.56ms step:1553/1575 train_time:95629ms step_avg:61.58ms step:1554/1575 train_time:95715ms step_avg:61.59ms step:1555/1575 train_time:95804ms step_avg:61.61ms step:1556/1575 train_time:95890ms step_avg:61.63ms step:1557/1575 train_time:95979ms step_avg:61.64ms step:1558/1575 train_time:96065ms step_avg:61.66ms step:1559/1575 train_time:96155ms step_avg:61.68ms step:1560/1575 train_time:96241ms step_avg:61.69ms step:1561/1575 train_time:96334ms step_avg:61.71ms step:1562/1575 train_time:96419ms step_avg:61.73ms step:1563/1575 train_time:96509ms step_avg:61.75ms step:1564/1575 train_time:96595ms step_avg:61.76ms step:1565/1575 train_time:96684ms step_avg:61.78ms step:1566/1575 train_time:96771ms step_avg:61.79ms step:1567/1575 train_time:96861ms step_avg:61.81ms step:1568/1575 train_time:96947ms step_avg:61.83ms step:1569/1575 train_time:97041ms step_avg:61.85ms step:1570/1575 train_time:97125ms step_avg:61.86ms step:1571/1575 train_time:97215ms step_avg:61.88ms step:1572/1575 train_time:97299ms step_avg:61.90ms step:1573/1575 train_time:97389ms step_avg:61.91ms step:1574/1575 train_time:97475ms step_avg:61.93ms step:1575/1575 train_time:97566ms step_avg:61.95ms step:1575/1575 val_loss:3.2782 train_time:97632ms step_avg:61.99ms peak memory allocated: 31016 MiB reserved: 46998 MiB