import os import sys # Read the current file and the kernels file code ASAP, for logging with open(sys.argv[0], 'r') as f: code = f.read() with open(os.path.join(os.path.dirname(sys.argv[0]), 'triton_kernels.py'), 'r') as f: code += f"\n\n{'-'*40}\n# triton_kernels.py\n{'-'*40}\n\n" code += f.read() import copy import glob import math import threading import time import uuid from dataclasses import dataclass from itertools import accumulate from pathlib import Path import gc os.environ["PYTORCH_ALLOC_CONF"] = "expandable_segments:True" import torch import triton torch.empty( 1, device=f"cuda:{os.environ['LOCAL_RANK']}", requires_grad=True ).backward() # prevents a bug on some systems import torch._dynamo as dynamo import torch.distributed as dist import torch.nn.functional as F # torch._inductor.config.coordinate_descent_tuning = True # we have banned this flag for new records because it causes compilation to take 30min from kernels import get_kernel from torch import Tensor, nn from triton_kernels import XXT, ba_plus_cAA, FusedLinearReLUSquareFunction, FusedSoftcappedCrossEntropy dynamo.config.recompile_limit = 64 # ----------------------------------------------------------------------------- # Custom operators: FP8 matmul by @YouJiacheng # Transposed layout by @ChrisJMcCormick allows for faster gradient accumulation. @torch.library.custom_op("nanogpt::mm_t", mutates_args=()) def mm_t_op(x: Tensor, w: Tensor, x_s: float, w_s: float, grad_s: float) -> tuple[Tensor, Tensor, Tensor]: """Computes y = x @ w with F8 weights stored as (in_features, out_features).""" @torch.compile def impl(x: Tensor, w: Tensor): assert x.is_contiguous() and w.is_contiguous() assert x.shape[1] == w.shape[0] # x: (batch, in), w: (in, out) x_f8 = x.div(x_s).to(torch.float8_e4m3fn) w_f8 = w.div(w_s).to(torch.float8_e4m3fn) # _scaled_mm requires column-major B. w_f8 is row-major (in, out). # .T.contiguous().T creates a column-major view without changing logical shape. w_f8_col_major = w_f8.T.contiguous().T out = torch._scaled_mm( x_f8, w_f8_col_major, out_dtype=torch.bfloat16, scale_a=x.new_tensor(x_s, dtype=torch.float32), scale_b=x.new_tensor(w_s, dtype=torch.float32), use_fast_accum=True, ) return out, x_f8, w_f8 return impl(x, w) @mm_t_op.register_fake def _(x: Tensor, w: Tensor, *_): assert x.ndim == w.ndim == 2 assert x.shape[1] == w.shape[0] assert x.device == w.device assert x.is_contiguous() and w.is_contiguous() return x @ w, x.to(torch.float8_e4m3fn), w.to(torch.float8_e4m3fn) @torch.library.custom_op("nanogpt::mm_t_backward", mutates_args=()) def mm_t_backward_op(g: Tensor, x_f8: Tensor, w_f8: Tensor, x_s: float, w_s: float, grad_s: float) -> tuple[Tensor, Tensor]: @torch.compile def impl(grad: Tensor, x_f8: Tensor, w_f8: Tensor): assert grad.is_contiguous() x_scale = grad.new_tensor(x_s, dtype=torch.float32) w_scale = grad.new_tensor(w_s, dtype=torch.float32) grad_scale = grad.new_tensor(grad_s, dtype=torch.float32) grad_f8 = grad.div(grad_s).to(torch.float8_e5m2) # grad_x = grad @ w.T grad_x = torch._scaled_mm( grad_f8, w_f8.T, out_dtype=torch.bfloat16, scale_a=grad_scale, scale_b=w_scale, use_fast_accum=False, ) # grad_w = x.T @ grad # Result is (in, out), naturally matching weight storage. No final .T needed. grad_w = torch._scaled_mm( x_f8.T.contiguous(), grad_f8.T.contiguous().T, out_dtype=torch.float32, scale_a=x_scale, scale_b=grad_scale, use_fast_accum=False, ) return grad_x, grad_w grad_x, grad_w = impl(g, x_f8, w_f8) return grad_x, grad_w @mm_t_backward_op.register_fake def _(g: Tensor, x_f8: Tensor, w_f8: Tensor, *_): return x_f8.to(torch.bfloat16), w_f8.to(torch.float32) def backward_t(ctx, grad_out: Tensor, *_): x_f8, w_f8 = ctx.saved_tensors x_s, w_s, grad_s = ctx.scales grad_x, grad_w = torch.ops.nanogpt.mm_t_backward( grad_out, x_f8, w_f8, x_s, w_s, grad_s ) return grad_x, grad_w, None, None, None def setup_context_t(ctx: torch.autograd.function.FunctionCtx, inputs, output): *_, x_s, w_s, grad_s = inputs _, x_f8, w_f8 = output ctx.save_for_backward(x_f8, w_f8) ctx.scales = x_s, w_s, grad_s ctx.set_materialize_grads(False) mm_t_op.register_autograd(backward_t, setup_context=setup_context_t) # ----------------------------------------------------------------------------- # Polar Express # Computed for num_iters=5, safety_factor=2e-2, cushion=2 polar_express_coeffs = [ (8.156554524902461, -22.48329292557795, 15.878769915207462), (4.042929935166739, -2.808917465908714, 0.5000178451051316), (3.8916678022926607, -2.772484153217685, 0.5060648178503393), (3.285753657755655, -2.3681294933425376, 0.46449024233003106), (2.3465413258596377, -1.7097828382687081, 0.42323551169305323) ] @torch.compile(dynamic=False, fullgraph=True) # Must use dynamic=False or else it's much slower def polar_express(G: torch.Tensor, split_baddbmm: bool = False): """ Polar Express Sign Method: https://arxiv.org/pdf/2505.16932 by Noah Amsel, David Persson, Christopher Musco, Robert M. Gower. """ X = G.bfloat16() if G.size(-2) > G.size(-1): X = X.mT # Ensure spectral norm is at most 1 X = X / (X.norm(dim=(-2, -1), keepdim=True) * (1 + 2e-2) + 1e-6) # Allocate buffers X = X.contiguous() A = torch.empty((*X.shape[:-1], X.size(-2)), device=X.device, dtype=X.dtype) B = torch.empty_like(A) C = torch.empty_like(X) # Select batched vs unbatched if split_baddbmm: BX_matmul = torch.bmm if X.ndim > 2 else torch.mm else: aX_plus_BX = torch.baddbmm if X.ndim > 2 else torch.addmm # Perform the iterations for a, b, c in polar_express_coeffs: XXT(X, out=A) # A = X @ X.mT ba_plus_cAA(A, alpha=c, beta=b, out=B) # B = b * A + c * A @ A # Referencing X twice causes pytorch to make a defensive copy, # resulting in a cudaMemcpyAsync in baddbmm. # For large matrices (i.e., the mlp weights), it's faster to split # the operation into two kernels to avoid this. if split_baddbmm: BX_matmul(B, X, out=C) # C = B @ X C.add_(X, alpha=a) # C = C + a*X (in-place, X only read) else: aX_plus_BX(X, B, X, beta=a, out=C) # C = a * X + B @ X X, C = C, X # Swap references to avoid unnecessary copies if G.size(-2) > G.size(-1): X = X.mT return X # ----------------------------------------------------------------------------- # Combined NorMuon + Adam Optimizer @dataclass class ParamConfig: """Per-parameter configuration for NorMuonAndAdam optimizer.""" label: str optim: str # "adam" or "normuon" comms: str # "none", "replicated", or "sharded" adam_betas: tuple[float, float] | None lr_mul: float wd_mul: float lr: float initial_lr: float weight_decay: float # Adam-specific eps: float | None = None # NorMuon-specific reshape: tuple | None = None chunk_size: int | None = None momentum: float | None = None beta2: float | None = None per_matrix_lr_mul: list[float] | None = None class NorMuonAndAdam: """ Combined optimizer that handles both NorMuon (for projection matrices) and Adam (for embeddings/scalars/gate weights). Muon - MomentUm Orthogonalized by Newton-schulz https://kellerjordan.github.io/posts/muon/ Muon internally runs standard SGD-momentum, and then performs an orthogonalization post- processing step, in which each 2D parameter's update is replaced with the nearest orthogonal matrix. To efficiently orthogonalize each update, Muon uses a Newton-Schulz iteration (replaced here with Polar Express), which has the advantage that it can be stably run in bfloat16 on the GPU. Muon is applied only to the projection matrices in the attention and MLP layers, and is not recommended for embeddings, scalars, or individual weight vectors (e.g., bias terms or gate weights). Differences from standard Muon: - Newton-Shulz is replaced with Polar Express for the orthogonalization step - NorMuon adds a low-rank variance estimator similar to Adafactor. https://arxiv.org/pdf/2510.05491 - Cautious weight decay, a gated version of decoupled weight decay - Mantissa tracking for precision Adam (for embeddings/scalars/gates): - Standard Adam with bias correction - Cautious weight decay Configuration: Unlike torch.optim.Optimizer, this class uses per-parameter configs from a `param_table` dict and does not include parameter "groups". All parameters require a .label attribute, and a corresponding entry in the param_table to specify their hyperparameters (lr_mul, wd_mul, adam_betas, etc.). Communication and ordering: Gradient communication is explicitly scheduled rather than hook-driven. Reductions are launched in `scatter_order`, while update math and final gathers are executed in `work_order`. These orders are independent and must each contain every parameter label exactly once. Two communication modes are supported per parameter: - 'replicated': Gradients are all-reduced and each rank computes the full update. - 'sharded': Gradients are reduce-scattered, each rank updates its shard, and results are all-gathered. Adam parameters may be freely sharded. NorMuon operates on full matrices; sharding is supported by grouping matrices into parameter banks. NorMuon parameters must have a `.reshape` attribute that reshapes the bank so that the leading dimension is divisible by world_size. # Contributors include @YouJiacheng, @KonstantinWilleke, @alexrgilbert, @adricarda, # @tuttyfrutyee, @vdlad, @ryanyang0, @vagrawal, @varunneal, @chrisjmccormick """ def __init__(self, named_params, param_table: dict, scatter_order: list, work_order: list, adam_defaults: dict, normuon_defaults: dict): self.world_size = dist.get_world_size() if dist.is_initialized() else 1 # Store defaults for each optimizer type self.adam_defaults = adam_defaults self.normuon_defaults = normuon_defaults self.param_table = param_table self.scatter_order = scatter_order self.work_order = work_order # Collect params by label and build config self.param_cfgs: dict[nn.Parameter, ParamConfig] = {} self.param_states: dict[nn.Parameter, dict] = {} self._param_by_label: dict[str, nn.Parameter] = {} for name, param in named_params: label = getattr(param, "label", None) assert label is not None and label in param_table # all params must have valid label assert label not in self._param_by_label # exactly one param per label self._param_by_label[label] = param self._build_param_cfg(param, label) # Assert scatter_order and work_order match present labels exactly present = set(self._param_by_label.keys()) assert set(scatter_order) == present and set(work_order) == present # Handle world_size=1: overwrite comms to "none" if self.world_size == 1: for p_cfg in self.param_cfgs.values(): p_cfg.comms = "none" # Initialize state for all params self._init_state() # 0-D CPU tensors to avoid recompilation self._step_size_t = torch.tensor(0.0, dtype=torch.float32, device="cpu") self._eff_wd_t = torch.tensor(0.0, dtype=torch.float32, device="cpu") self._eff_lr_t = torch.tensor(0.0, dtype=torch.float32, device="cpu") # Track async operations self._reduce_futures: dict[nn.Parameter, tuple] = {} # Embed/lm_head tying state self.split_embed = False self._lm_head_param = self._param_by_label.get("lm_head") self._embed_param = self._param_by_label.get("embed") def _build_param_cfg(self, param: nn.Parameter, label: str): """Build config for a single parameter from param_table.""" table_entry = self.param_table[label] optim = table_entry["optim"] comms = table_entry["comms"] adam_betas = table_entry.get("adam_betas") lr_mul = table_entry.get("lr_mul", 1.0) wd_mul = table_entry.get("wd_mul", 1.0) if optim == "adam": chunk_size = param.shape[0] // self.world_size if comms == "sharded" else None p_cfg = ParamConfig( label=label, optim=optim, comms=comms, adam_betas=tuple(adam_betas) if adam_betas else None, lr_mul=lr_mul, wd_mul=wd_mul, lr=self.adam_defaults["lr"], initial_lr=self.adam_defaults["lr"], weight_decay=self.adam_defaults["weight_decay"], eps=self.adam_defaults["eps"], chunk_size=chunk_size, ) elif optim == "normuon": reshape = getattr(param, "reshape", None) if reshape is None: raise ValueError(f"NorMuon param {label} must have .reshape attribute") if reshape[0] % self.world_size != 0: raise ValueError(f"reshape[0]={reshape[0]} must be divisible by world_size") chunk_size = reshape[0] // self.world_size chunk_shape = (chunk_size, *reshape[1:]) # Shape-based LR multiplier for NorMuon shape_mult = max(1.0, chunk_shape[-2] / chunk_shape[-1]) ** 0.5 if len(chunk_shape) >= 2 else 1.0 lr_mul = shape_mult * lr_mul # Per-matrix LR multipliers for MLP c_proj (2x LR on odd indices) per_matrix_lr_mul = None if label == "mlp": rank = dist.get_rank() if dist.is_initialized() else 0 start_idx = rank * chunk_size per_matrix_lr_mul = [] for i in range(chunk_size): global_idx = start_idx + i is_c_proj = (global_idx % 2 == 1) per_matrix_lr_mul.append(2.0 if is_c_proj else 1.0) p_cfg = ParamConfig( label=label, optim=optim, comms=comms, adam_betas=tuple(adam_betas) if adam_betas else None, lr_mul=lr_mul, wd_mul=wd_mul, lr=self.normuon_defaults["lr"], initial_lr=self.normuon_defaults["lr"], weight_decay=self.normuon_defaults["weight_decay"], reshape=reshape, chunk_size=chunk_size, momentum=self.normuon_defaults["momentum"], beta2=self.normuon_defaults["beta2"], per_matrix_lr_mul=per_matrix_lr_mul, ) else: raise ValueError(f"Unknown optim type: {optim}") self.param_cfgs[param] = p_cfg def _init_state(self): """Initialize optimizer state for all parameters.""" for param, p_cfg in self.param_cfgs.items(): if p_cfg.optim == "adam": # Sharded params use chunk state, replicated use full state if p_cfg.comms == "sharded": chunk = param[:p_cfg.chunk_size] else: chunk = param exp_avg = torch.zeros_like(chunk, dtype=torch.float32, device=param.device) self.param_states[param] = dict(step=0, exp_avg=exp_avg, exp_avg_sq=torch.zeros_like(exp_avg)) elif p_cfg.optim == "normuon": chunk_shape = (p_cfg.chunk_size, *p_cfg.reshape[1:]) # Momentum buffer (FP32 for precision) momentum_buffer = torch.zeros( chunk_shape, dtype=torch.float32, device=param.device ) # Second momentum buffer - reduced along one dimension if chunk_shape[-2] >= chunk_shape[-1]: second_mom_shape = (*chunk_shape[:-1], 1) else: second_mom_shape = (*chunk_shape[:-2], 1, chunk_shape[-1]) second_momentum_buffer = torch.zeros( second_mom_shape, dtype=torch.float32, device=param.device ) # Mantissa buffer for precision tracking mantissa = torch.zeros( chunk_shape, dtype=torch.uint16, device=param.device ) self.param_states[param] = dict( momentum_buffer=momentum_buffer, second_momentum_buffer=second_momentum_buffer, mantissa=mantissa, ) # ----------------------------------- # Reduce/Gather operations def _launch_reduce(self, param: nn.Parameter, grad: Tensor): """Launch async reduce for a parameter based on its comms policy.""" p_cfg = self.param_cfgs[param] if p_cfg.comms == "none": if p_cfg.optim == "normuon": # NorMuon needs reshaped gradient even without communication grad = grad.view(p_cfg.reshape) self._reduce_futures[param] = (None, grad) elif p_cfg.comms == "replicated": future = dist.all_reduce(grad, op=dist.ReduceOp.AVG, async_op=True).get_future() self._reduce_futures[param] = (future, grad) elif p_cfg.comms == "sharded": if p_cfg.optim == "normuon": # NorMuon: reshape before reduce_scatter grad_reshaped = grad.view(p_cfg.reshape) grad_chunk = torch.empty( (p_cfg.chunk_size, *grad_reshaped.shape[1:]), dtype=grad.dtype, device=grad.device ) future = dist.reduce_scatter_tensor( grad_chunk, grad_reshaped.contiguous(), op=dist.ReduceOp.AVG, async_op=True ).get_future() self._reduce_futures[param] = (future, grad_chunk) else: # Adam: simple reduce_scatter grad_chunk = torch.empty_like(grad[:p_cfg.chunk_size]) future = dist.reduce_scatter_tensor( grad_chunk, grad, op=dist.ReduceOp.AVG, async_op=True ).get_future() self._reduce_futures[param] = (future, grad_chunk) def _launch_gather(self, param: nn.Parameter, p_slice: Tensor) -> "torch.futures.Future": """Launch async all_gather for a sharded parameter.""" p_cfg = self.param_cfgs[param] if p_cfg.optim == "normuon": full_param = param.data.view(p_cfg.reshape) assert full_param.is_contiguous() return dist.all_gather_into_tensor( full_param, p_slice.contiguous(), async_op=True ).get_future() else: return dist.all_gather_into_tensor( param, p_slice.contiguous(), async_op=True ).get_future() # ----------------------------------- # State management def reset(self): """Reset NorMuon momentum buffers and split_embed state (called on training reset).""" self.split_embed = False for param, p_cfg in self.param_cfgs.items(): if p_cfg.optim == "normuon": p_state = self.param_states[param] p_state["momentum_buffer"].zero_() p_state["mantissa"].zero_() p_state["second_momentum_buffer"].zero_() def copy_lm_state_to_embed(self): """ Copy the optimizer state from the lm_head to the embed at the untie point. This requires an all-gather + reshard because of different sharding: - lm_head (768, 50304) is sharded to (96, 50304) per rank (along model_dim) - embed (50304, 768) is sharded to (6288, 768) per rank (along vocab_size) We all-gather the lm_head momentum, transpose it, then each rank takes their embed shard to get the correct momentum state. """ lm_head = self._lm_head_param embed = self._embed_param lm_state = self.param_states[lm_head] embed_state = self.param_states[embed] lm_cfg = self.param_cfgs[lm_head] embed_cfg = self.param_cfgs[embed] embed_state['step'] = lm_state['step'] # Preserve step count for bias correction # Copy optimizer state with all-gather + transpose + reshard if self.world_size > 1: rank = dist.get_rank() lm_chunk_size = lm_cfg.chunk_size # 96 embed_chunk_size = embed_cfg.chunk_size # 6288 # All-gather lm_head momentum to get full (768, 50304) tensor for key in ["exp_avg", "exp_avg_sq"]: lm_chunk = lm_state[key] # (96, 50304) full_lm = torch.empty(lm_head.shape[0], lm_head.shape[1], dtype=lm_chunk.dtype, device=lm_chunk.device) dist.all_gather_into_tensor(full_lm, lm_chunk.contiguous()) embed_state[key].copy_(full_lm.T[rank * embed_chunk_size:(rank + 1) * embed_chunk_size]) else: # Single GPU: simple transpose for key in ["exp_avg", "exp_avg_sq"]: embed_state[key].copy_(lm_state[key].T) # Mark as split self.split_embed = True def state_dict(self): """Return the optimizer state as a dict.""" return { "param_states": {id(p): s for p, s in self.param_states.items()}, "param_cfgs": {id(p): s for p, s in self.param_cfgs.items()}, } def load_state_dict(self, state_dict): """Load optimizer state from a dict.""" # Build id->param mapping id_to_param = {id(p): p for p in self.param_cfgs.keys()} # Load state, preserving dtypes for param_id, saved_p_state in state_dict["param_states"].items(): if param_id in id_to_param: param = id_to_param[param_id] p_state = self.param_states[param] for k, v in saved_p_state.items(): if isinstance(v, torch.Tensor) and k in p_state: target_dtype = p_state[k].dtype p_state[k] = v.to(dtype=target_dtype, device=p_state[k].device) else: p_state[k] = v # ----------------------------------- # Unified optimizer step with explicit ordering @torch.no_grad() def step(self, do_adam: bool = True): """ Combined optimizer step with explicit ordering. Args: do_adam: If True, update Adam params. NorMuon params always updated. Flow: 1. Scatter phase: Launch reduces in scatter_order 2. Work phase: Process updates in work_order - Wait for reduce, compute update, launch gather 3. Finalize phase: Wait for gathers While the embeddings are tied: - Comms and update math are only done on lm_head. - We add embed.grad.T into lm_head.grad before comms. - After lm_head gather, we copy lm_head.data.T --> embed.data """ rank = dist.get_rank() if dist.is_initialized() else 0 lm_param, embed_param = self._lm_head_param, self._embed_param # ===== Phase 1: Launch reduces in scatter_order ===== for label in self.scatter_order: param = self._param_by_label[label] p_cfg = self.param_cfgs[param] if p_cfg.optim == "adam" and not do_adam: continue if param.grad is None: continue # lm_head when tied: aggregate embed.grad.T (transposed shapes) if label == "lm_head" and do_adam and not self.split_embed: if embed_param is not None and embed_param.grad is not None: param.grad.add_(embed_param.grad.T) # Skip embed when tied (copied from lm_head after gather) if label == "embed" and not self.split_embed: continue self._launch_reduce(param, param.grad) # ===== Phase 2: Process updates in work_order ===== gather_futures = [] lm_head_gather_future = None for label in self.work_order: param = self._param_by_label[label] if param not in self._reduce_futures: continue p_cfg = self.param_cfgs[param] if p_cfg.optim == "adam" and not do_adam: continue # Wait for reduce future, grad_chunk = self._reduce_futures[param] if future is not None: future.wait() # Apply update based on optim type if p_cfg.optim == "adam": p_slice = self._adam_update(param, grad_chunk, p_cfg, rank) else: p_slice = self._normuon_update(param, grad_chunk, p_cfg, rank) # Launch gather for sharded params if p_cfg.comms == "sharded" and self.world_size > 1: gather_fut = self._launch_gather(param, p_slice) if label == "lm_head": lm_head_gather_future = gather_fut else: gather_futures.append(gather_fut) # ===== Phase 3: Wait for gathers, sync embed if tied ===== # Wait for lm_head gather first so we can copy to embed while other gathers complete if lm_head_gather_future is not None: lm_head_gather_future.wait() # When tied: copy lm_head.T to embed if do_adam and not self.split_embed and embed_param is not None and lm_param is not None: embed_param.data.copy_(lm_param.data.T) # Wait for remaining gathers for fut in gather_futures: fut.wait() self._reduce_futures.clear() # Clear grads for updated params for param, p_cfg in self.param_cfgs.items(): if p_cfg.optim == "adam" and not do_adam: continue # Don't clear Adam grads on even steps param.grad = None # ----------------------------------- # Adam update def _adam_update(self, param: nn.Parameter, grad_chunk: Tensor, p_cfg: ParamConfig, rank: int) -> Tensor: """Apply Adam update to a parameter. Returns the updated p_slice.""" beta1, beta2 = p_cfg.adam_betas lr = p_cfg.lr * p_cfg.lr_mul # Get parameter slice if p_cfg.comms == "sharded": p_slice = param[rank * p_cfg.chunk_size:(rank + 1) * p_cfg.chunk_size] else: p_slice = param p_state = self.param_states[param] p_state["step"] += 1 t = p_state["step"] bias1, bias2 = 1 - beta1 ** t, 1 - beta2 ** t self._step_size_t.fill_(lr * (bias2 ** 0.5 / bias1)) self._eff_wd_t.fill_(lr * lr * p_cfg.weight_decay * p_cfg.wd_mul) NorMuonAndAdam._adam_update_step( p_slice, grad_chunk, p_state["exp_avg"], p_state["exp_avg_sq"], beta1, beta2, p_cfg.eps, self._step_size_t, self._eff_wd_t ) return p_slice @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _adam_update_step(p_slice, g_slice, exp_avg, exp_avg_sq, beta1, beta2, eps, step_size_t, eff_wd_t): """Compiled Adam update step.""" exp_avg.mul_(beta1).add_(g_slice, alpha=1 - beta1) exp_avg_sq.mul_(beta2).addcmul_(g_slice, g_slice, value=1 - beta2) update = exp_avg.div(exp_avg_sq.sqrt().add_(eps)).mul_(step_size_t) # Cautious weight decay mask = (update * p_slice) > 0 update.addcmul_(p_slice, mask, value=eff_wd_t) p_slice.add_(other=update, alpha=-1.0) # ----------------------------------- # NorMuon update def _normuon_update(self, param: nn.Parameter, grad_chunk: Tensor, p_cfg: ParamConfig, rank: int) -> Tensor: """Apply NorMuon update to a parameter. Returns the updated p_slice.""" chunk_shape = grad_chunk.shape p_state = self.param_states[param] grad_chunk = grad_chunk.float() # FP32 for momentum # Momentum update momentum_buffer = p_state["momentum_buffer"] momentum_buffer.lerp_(grad_chunk, 1 - p_cfg.momentum) updated_grads = grad_chunk.lerp_(momentum_buffer, p_cfg.momentum) self._eff_lr_t.fill_(p_cfg.lr_mul * p_cfg.lr) self._eff_wd_t.fill_(p_cfg.wd_mul * p_cfg.weight_decay * p_cfg.lr) # Polar Express orthogonalization is_large_matrix = chunk_shape[-2] > 1024 v_chunk = polar_express(updated_grads, split_baddbmm=is_large_matrix) # Variance reduction red_dim = -1 if chunk_shape[-2] >= chunk_shape[-1] else -2 v_chunk = NorMuonAndAdam._apply_normuon_variance_reduction( v_chunk, p_state["second_momentum_buffer"], p_cfg.beta2, red_dim ) # Update parameter, in place, with cautious weight decay param_view = param.data.view(p_cfg.reshape) p_slice = param_view[rank * p_cfg.chunk_size:(rank + 1) * p_cfg.chunk_size] # MLP has per-matrix LR multipliers (c_proj gets 2x LR) if p_cfg.per_matrix_lr_mul is not None: for mat_idx in range(p_cfg.chunk_size): self._eff_lr_t.fill_(p_cfg.lr_mul * p_cfg.per_matrix_lr_mul[mat_idx] * p_cfg.lr) self._eff_wd_t.fill_(p_cfg.wd_mul * p_cfg.weight_decay * p_cfg.lr) NorMuonAndAdam._cautious_wd_and_update_inplace( p_slice[mat_idx].view(torch.uint16), p_state["mantissa"][mat_idx], v_chunk[mat_idx], self._eff_wd_t, self._eff_lr_t ) else: NorMuonAndAdam._cautious_wd_and_update_inplace( p_slice.view(torch.uint16), p_state["mantissa"], v_chunk, self._eff_wd_t, self._eff_lr_t ) return p_slice @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _cautious_wd_and_update_inplace(p, mantissa, grad, wd_tensor, lr_tensor): """ Cautious weight decay + parameter update. wd_tensor and lr_tensor are 0-D CPU tensors. Mantissa is tracked to enable higher precision updates on bfloat16 parameters. bfloat16 format: 1 sign bit + 8 exponent bits + 7 mantissa bits = 16 bits total float32 format: 1 sign bit + 8 exponent bits + 23 mantissa bits = 32 bits total """ assert p.dtype == mantissa.dtype == torch.uint16 grad = grad.float() wd_factor = wd_tensor.to(torch.float32) lr_factor = lr_tensor.to(torch.float32) p_precise_raw = (p.to(torch.uint32) << 16) | mantissa.to(torch.uint32) p_precise = p_precise_raw.view(torch.float32) mask = (grad * p_precise) >= 0 p_precise.copy_(p_precise - (p_precise * mask * wd_factor * lr_factor) - (grad * lr_factor)) p.copy_((p_precise_raw >> 16).to(torch.uint16)) mantissa.copy_(p_precise_raw.to(torch.uint16)) @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _apply_normuon_variance_reduction(v_chunk, second_momentum_buffer, beta2, red_dim): """NorMuon variance reduction. Algebraically fuses the normalization steps to minimize memory ops.""" v_mean = v_chunk.float().square().mean(dim=red_dim, keepdim=True) red_dim_size = v_chunk.size(red_dim) v_norm_sq = v_mean.sum(dim=(-2, -1), keepdim=True).mul_(red_dim_size) v_norm = v_norm_sq.sqrt_() second_momentum_buffer.lerp_(v_mean.to(dtype=second_momentum_buffer.dtype), 1 - beta2) step_size = second_momentum_buffer.clamp_min(1e-10).rsqrt_() scaled_sq_sum = (v_mean * red_dim_size) * step_size.float().square() v_norm_new = scaled_sq_sum.sum(dim=(-2, -1), keepdim=True).sqrt_() final_scale = step_size * (v_norm / v_norm_new.clamp_min_(1e-10)) return v_chunk.mul_(final_scale.type_as(v_chunk)) # ----------------------------------------------------------------------------- # PyTorch nn.Module definitions for the model def norm(x: Tensor): return F.rms_norm(x, (x.size(-1),)) class CastedLinearT(nn.Module): """ Linear layer with transposed weight storage (in_features, out_features) which addresses the slow kernel that was used for gradient accumulation. @chrisjmccormick """ def __init__(self, in_features: int, out_features: int, use_fp8=False, x_s=1.0, w_s=1.0, grad_s=1.0): super().__init__() self.in_features = in_features self.out_features = out_features self.use_fp8 = use_fp8 self.x_s = x_s self.w_s = w_s self.grad_s = grad_s self.weight = nn.Parameter(torch.empty(in_features, out_features, dtype=torch.bfloat16)) self.reset_parameters() def reset_parameters(self) -> None: with torch.no_grad(): nn.init.zeros_(self.weight) # @Grad62304977 and others def forward(self, x: Tensor): if self.use_fp8 and self.training: _x = x.flatten(0, -2) out = torch.ops.nanogpt.mm_t(_x, self.weight, x_s=self.x_s, w_s=self.w_s, grad_s=self.grad_s)[0] return out.reshape(*x.shape[:-1], -1) else: return x @ self.weight.type_as(x) # ----------------------------------------------------------------------------- # PyTorch nn.Module definitions for the model class Yarn(nn.Module): def __init__(self, head_dim, max_seq_len): super().__init__() self.head_dim = head_dim self.max_seq_len = max_seq_len self.reset() def rotary(self, x_BTHD): assert self.factor1.size(0) >= x_BTHD.size(-3) factor1, factor2 = ( self.factor1[None, : x_BTHD.size(-3), None, :], self.factor2[None, : x_BTHD.size(-3), None, :], ) x_flip = x_BTHD.view(*x_BTHD.shape[:-1], x_BTHD.shape[-1] // 2, 2).flip(-1).view(x_BTHD.shape) return factor1 * x_BTHD + factor2 * x_flip def reset(self): angular_freq = (1 / 1024) ** torch.linspace(0, 1, steps=self.head_dim//4, dtype=torch.float32, device=device) angular_freq = angular_freq.repeat_interleave(2) # half-truncate RoPE by @YouJiacheng (w/ base freq tuning) angular_freq = torch.cat([angular_freq, angular_freq.new_zeros(self.head_dim//2)]) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=device) theta = torch.outer(t, angular_freq) self.factor1 = nn.Buffer( theta.cos().to(torch.bfloat16), persistent=False ) self.factor2 = nn.Buffer( theta.sin().to(torch.bfloat16), persistent=False ) self.factor2[..., 1::2] *= -1 self.angular_freq = angular_freq # start with 0.1, inspired by 0.12 from @leloykun and learnable scalars used by @brendanh0gan https://x.com/hi_tysam/status/1879693583898591283 self.attn_scale = 0.1 def apply(self, old_window: int, new_window: int, alpha: int=1, beta: int=32): rotations = args.block_size * old_window * self.angular_freq / (2 * torch.pi) scaling_factor = old_window / new_window interpolation_weight = torch.clamp((rotations - alpha) / (beta - alpha), 0, 1) self.angular_freq *= scaling_factor + interpolation_weight * (1 - scaling_factor) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=self.angular_freq.device) theta = torch.outer(t, self.angular_freq) self.factor1.copy_(theta.cos()) self.factor2.copy_(theta.sin()) self.factor2[..., 1::2] *= -1 self.attn_scale *= 0.2 * math.log(new_window / old_window) + 1 class YarnPairedHead(nn.Module): def __init__(self, head_dim, max_seq_len): super().__init__() self.head_dim = head_dim self.max_seq_len = max_seq_len self.reset() def rotary(self, x_BTHD): assert self.factor1.size(0) >= x_BTHD.size(-3) factor1, factor2 = ( self.factor1[None, : x_BTHD.size(-3), None, :], self.factor2[None, : x_BTHD.size(-3), None, :], ) x_flip = x_BTHD.view(*x_BTHD.shape[:-1], x_BTHD.shape[-1] // 2, 2).flip(-1).view(x_BTHD.shape) return factor1 * x_BTHD + factor2 * x_flip def reset(self): angular_freq = (1 / 1024) ** torch.linspace(0, 1, steps=self.head_dim//4, dtype=torch.float32, device=device) angular_freq = angular_freq.repeat_interleave(2) angular_freq = torch.cat([angular_freq, angular_freq.new_zeros(self.head_dim//2)]) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=device) t_even = 2 * t t_odd = 2 * t + 1 theta1 = torch.outer(t_even, angular_freq) theta2 = torch.outer(t_odd, angular_freq) self.factor1 = nn.Buffer( torch.cat((theta1.cos(),theta2.cos()), dim=-1).to(torch.bfloat16), persistent=False ) self.factor2 = nn.Buffer( torch.cat((theta1.sin(),theta2.sin()), dim=-1).to(torch.bfloat16), persistent=False ) self.factor2[..., 1::2] *= -1 self.angular_freq = angular_freq # start with 0.1, inspired by 0.12 from @leloykun and learnable scalars used by @brendanh0gan https://x.com/hi_tysam/status/1879693583898591283 self.attn_scale = 0.1 def apply(self, old_window: int, new_window: int, alpha: int=1, beta: int=32): rotations = args.block_size * old_window * self.angular_freq / (2 * torch.pi) scaling_factor = old_window / new_window interpolation_weight = torch.clamp((rotations - alpha) / (beta - alpha), 0, 1) self.angular_freq *= scaling_factor + interpolation_weight * (1 - scaling_factor) t = torch.arange(2*self.max_seq_len, dtype=torch.float32, device=self.angular_freq.device) t_even = 2 * t t_odd = 2 * t + 1 theta1 = torch.outer(t_even, self.angular_freq) theta2 = torch.outer(t_odd, self.angular_freq) self.factor1.copy_(torch.cat((theta1.cos(),theta2.cos()), dim=-1)) self.factor2.copy_( torch.cat((theta1.sin(),theta2.sin()), dim=-1)) self.factor2[..., 1::2] *= -1 self.attn_scale *= 0.2 * math.log(new_window / old_window) + 1 @dataclass class AttnArgs: ve: torch.Tensor sa_lambdas: torch.Tensor seqlens: torch.Tensor bm_size: int yarn: Yarn key_offset: bool attn_gate_w: torch.Tensor ve_gate_w: torch.Tensor flash_attn_interface = get_kernel('varunneal/flash-attention-3').flash_attn_interface class CausalSelfAttention(nn.Module): def __init__(self, dim: int, head_dim: int, num_heads: int): super().__init__() self.num_heads = num_heads self.head_dim = head_dim self.dim = dim self.hdim = num_heads * head_dim assert self.hdim == self.dim, "num_heads * head_dim must equal model_dim" # Weights are stored in parameter banks and passed via forward() def forward(self, x: Tensor, attn_args: AttnArgs, qkvo_w: Tensor): B, T = x.size(0), x.size(1) # batch size, sequence length assert B == 1, "varlen sequences requires B == 1" assert T % 16 == 0 # unpack attention args yarn = attn_args.yarn ve, sa_lambdas, key_offset = attn_args.ve, attn_args.sa_lambdas, attn_args.key_offset seqlens, bm_size = attn_args.seqlens, attn_args.bm_size # sparse gated attention to enable context based no-op by @classiclarryd # only include gates on layers with value embeds used on forward pass attn_gate_w, ve_gate_w = attn_args.attn_gate_w, attn_args.ve_gate_w q, k, v = F.linear(x, sa_lambdas[0] * qkvo_w[:self.dim * 3].type_as(x)).view(B, T, 3 * self.num_heads, self.head_dim).chunk(3, dim=-2) q, k = norm(q), norm(k) # QK norm @Grad62304977 q, k = yarn.rotary(q), yarn.rotary(k) if key_offset: # shift keys forward for the stationary head dims. Enables 1-layer induction. k[:, 1:, :, self.head_dim // 2:] = k[:, :-1, :, self.head_dim // 2:] if ve is not None: ve_gate_out = 2 * torch.sigmoid(F.linear(x[..., :12], ve_gate_w)).view(B, T, self.num_heads, 1) v = v + ve_gate_out * ve.view_as(v) # @ KoszarskyB & @Grad62304977 max_len = args.train_max_seq_len if self.training else (args.val_batch_size // (grad_accum_steps * world_size)) # use flash_attn over flex_attn @varunneal. flash_attn_varlen suggested by @YouJiacheng y = flash_attn_interface.flash_attn_varlen_func(q[0], k[0], v[0], cu_seqlens_q=seqlens, cu_seqlens_k=seqlens, max_seqlen_q=max_len, max_seqlen_k=max_len, causal=True, softmax_scale=yarn.attn_scale, window_size=(bm_size, 0)) y = y.view(B, T, self.num_heads, self.head_dim) y = y * torch.sigmoid(F.linear(x[..., :12], attn_gate_w)).view(B, T, self.num_heads, 1) y = y.contiguous().view(B, T, self.num_heads * self.head_dim) # re-assemble all head outputs side by side y = F.linear(y, sa_lambdas[1] * qkvo_w[self.dim * 3:].type_as(y)) # sa_lambdas[1] pre-multiplied to O @shenberg return y class PairedHeadCausalSelfAttention(nn.Module): """ Pairs up attention heads such that queries from head 1 can attend to keys in head 2, and vice-versa. Implemented by interleaving the k, q, and v for pairs of heads to form twice as long sequences EG [k1_h1, k2_h1, k3_h1], [k1_h2, k2_h2, k3_h2] -> [k1_h1, k1_h2, k2_h1, k2_h2, k3_h1, k3_h2], repeat for q and v """ def __init__(self, dim: int, head_dim: int, num_heads: int): super().__init__() self.num_heads = num_heads self.head_dim = head_dim self.dim = dim self.hdim = num_heads * head_dim assert self.hdim == self.dim, "num_heads * head_dim must equal model_dim" # Weights are stored in parameter banks and passed via forward() def forward(self, x: Tensor, attn_args: AttnArgs, qkvo_w: Tensor): B, T = x.size(0), x.size(1) # batch size, sequence length assert B == 1, "varlen sequences requires B == 1" assert T % 16 == 0 # unpack attention args yarn = attn_args.yarn ve, sa_lambdas = attn_args.ve, attn_args.sa_lambdas seqlens, bm_size = attn_args.seqlens, attn_args.bm_size attn_gate_w, ve_gate_w = attn_args.attn_gate_w, attn_args.ve_gate_w q, k, v = F.linear(x, sa_lambdas[0] * qkvo_w[:self.dim * 3].type_as(x)).view(B, T, 3 * self.num_heads, self.head_dim).chunk(3, dim=-2) q, k = norm(q), norm(k) # delay q,k reshape until rotary makes data contiguous, to enable view (non-copy) q = q.view(B, T, self.num_heads // 2, self.head_dim * 2) k = k.view(B, T, self.num_heads // 2, self.head_dim * 2) v = v.reshape(B, T*2, self.num_heads//2, self.head_dim) q, k = yarn.rotary(q), yarn.rotary(k) q = q.view(B, T*2, self.num_heads//2, self.head_dim) k = k.view(B, T*2, self.num_heads//2, self.head_dim) if ve is not None: ve_gate_out = 2 * torch.sigmoid(F.linear(x[..., :12], ve_gate_w)).view(B, T*2, self.num_heads//2, 1) v = v + ve_gate_out * ve.view_as(v) max_len = args.train_max_seq_len if self.training else (args.val_batch_size // (grad_accum_steps * world_size)) # paired head correction seqlens = 2 * seqlens max_len = 2 * max_len y = flash_attn_interface.flash_attn_varlen_func(q[0], k[0], v[0], cu_seqlens_q=seqlens, cu_seqlens_k=seqlens, max_seqlen_q=max_len, max_seqlen_k=max_len, causal=True, softmax_scale=yarn.attn_scale, window_size=(bm_size, 0)) y = y.view(B, T, self.num_heads, self.head_dim) y = y * torch.sigmoid(F.linear(x[..., :12], attn_gate_w)).view(B, T, self.num_heads, 1) y = y.contiguous().view(B, T, self.num_heads * self.head_dim) y = F.linear(y, sa_lambdas[1] * qkvo_w[self.dim * 3:].type_as(y)) return y class MLP(nn.Module): def __init__(self): super().__init__() # Weights are stored in parameter banks and passed via forward() def forward(self, x: Tensor, c_fc: Tensor, c_proj: Tensor): # relu(x)^2: # https://arxiv.org/abs/2109.08668v2; ~1-2% better than GELU; suggested by @SKYLINEZ007 and @Grad62304977 # Fused triton kernel for relu(x @ W1.T)^2 @ W2.T return FusedLinearReLUSquareFunction.apply(x, c_fc, c_proj) class Block(nn.Module): def __init__(self, dim: int, head_dim: int, num_heads: int, has_attn: bool, has_mlp: bool, use_paired_head: bool): super().__init__() # skip attention of blocks.6 (the 7th layer) by @YouJiacheng if has_attn: if use_paired_head: self.attn = PairedHeadCausalSelfAttention(dim, head_dim, num_heads) else: self.attn = CausalSelfAttention(dim, head_dim, num_heads) else: self.attn = None # skip MLP blocks for first MLP layer by @EmelyanenkoK self.mlp = MLP() if has_mlp else None def forward(self, x: Tensor, attn_args: AttnArgs, qkvo_w: Tensor = None, c_fc: Tensor = None, c_proj: Tensor = None): if self.attn is not None: x = x + self.attn(norm(x), attn_args, qkvo_w) if self.mlp is not None: x = x + self.mlp(norm(x), c_fc, c_proj) return x # ----------------------------------------------------------------------------- # The main model def next_multiple_of_n(v: float | int, *, n: int): return next(x for x in range(n, int(v) + 1 + n, n) if x >= v) @dataclass class ForwardScheduleConfig: mtp_weights: torch.Tensor ws_short: int ws_long: int class GPT(nn.Module): def __init__(self, vocab_size: int, num_layers: int, num_heads: int, head_dim: int, model_dim: int, max_seq_len: int): super().__init__() self.num_layers = num_layers vocab_size = next_multiple_of_n(vocab_size, n=128) self.smear_gate = nn.Linear(12, 1, bias=False) nn.init.zeros_(self.smear_gate.weight) self.smear_gate.weight.label = 'smear_gate' self.skip_gate = nn.Linear(12, 1, bias=False) nn.init.zeros_(self.skip_gate.weight) self.skip_gate.weight.label = 'skip_gate' # token value embeddings by @KoszarskyB - inspired by @Grad62304977's value residual implementation following https://arxiv.org/abs/2410.17897 # value embedding code simplification inspired by @ragulpr https://github.com/KellerJordan/modded-nanogpt/pull/78 self.value_embeds = nn.ModuleList([nn.Embedding(vocab_size, model_dim) for _ in range(3)]) for embed in self.value_embeds: nn.init.zeros_(embed.weight) for i, ve in enumerate(self.value_embeds): ve.weight.label = f've{i}' # ve0, ve1, ve2 # parameter banks for attention and value embedding gate weights self.attn_gate_bank = nn.Parameter(torch.zeros(10, num_heads, 12)) # 10 layers self.attn_gate_bank.label = 'attn_gate_bank' self.ve_gate_bank = nn.Parameter(torch.zeros(5, num_heads, 12)) # 5 layers self.ve_gate_bank.label = 've_gate_bank' # ----------------------------------- # Parameter banks for sharded optimization, by @chrisjmccormick # Identify which layers have attention/MLP # Attention is skipped in layer 6 by @YouJiacheng self.attn_layer_indices = [i for i in range(num_layers) if i != 6] # All layers have MLP (At 11 layers--dropped first layer @EmelyanenkoK) self.mlp_layer_indices = list(range(num_layers)) hdim = num_heads * head_dim mlp_hdim = 4 * model_dim # Create index mappings: layer_idx -> bank_idx self.layer_to_attn_idx = {layer_idx: bank_idx for bank_idx, layer_idx in enumerate(self.attn_layer_indices)} self.layer_to_mlp_idx = {layer_idx: bank_idx for bank_idx, layer_idx in enumerate(self.mlp_layer_indices)} # Attention bank: stores QKVO weights for all attention layers # merged QKVO weights: suggested by many, implemented by @fernbear.bsky.social, and further improved by @YouJiacheng # https://x.com/hi_tysam/status/1879699187107033311 # Simplified layout by @chrisjmccormick # Shape: (num_attn_layers, 4*model_dim, hdim) = (10, 3072, 768) # Reshape for sharding: (40, 768, 768) for even distribution across 8 GPUs self.attn_bank = nn.Parameter(torch.empty(len(self.attn_layer_indices), 4 * model_dim, hdim)) self.attn_bank.label = 'attn' self.attn_bank.reshape = (len(self.attn_layer_indices) * 4, hdim, hdim) # (40, 768, 768) # MLP bank: stores c_fc and c_proj for all MLP layers # Shape: (num_mlp_layers + padding, 2, mlp_hdim, model_dim) = (12, 2, 3072, 768) # We add 1 padding layer (index 11) to get 12*2=24 matrices for even distribution across 8 GPUs # Reshape for sharding: (24, 3072, 768) num_mlp_with_padding = len(self.mlp_layer_indices) + 1 # 11 + 1 = 12 self.mlp_bank = nn.Parameter(torch.empty(num_mlp_with_padding, 2, mlp_hdim, model_dim)) self.mlp_bank.label = 'mlp' self.mlp_bank.reshape = (num_mlp_with_padding * 2, mlp_hdim, model_dim) # (24, 3072, 768) # improved init scale by @YouJiacheng # Attention uses dim^-0.5, MLP uses 0.5 * dim^-0.5 attn_std = model_dim ** -0.5 attn_bound = (3 ** 0.5) * attn_std mlp_std = 0.5 * (model_dim ** -0.5) mlp_bound = (3 ** 0.5) * mlp_std with torch.no_grad(): # Init attention bank (QKV uniform, O zero) self.attn_bank[:, :model_dim * 3, :].uniform_(-attn_bound, attn_bound) self.attn_bank[:, model_dim * 3:, :].zero_() # Init MLP bank (c_fc uniform, c_proj zero) self.mlp_bank[:, 0, :, :].uniform_(-mlp_bound, mlp_bound) # c_fc self.mlp_bank[:, 1, :, :].zero_() # c_proj - zero init suggested by @Grad62304977 # Create blocks with has_attn/has_mlp flags self.paired_head_layers = [0, 2, 5, 9] self.blocks = nn.ModuleList([ Block(model_dim, head_dim, num_heads, has_attn=(i in self.layer_to_attn_idx), has_mlp=(i in self.layer_to_mlp_idx), use_paired_head=(i in self.paired_head_layers)) for i in range(num_layers) ]) self.yarn = Yarn(head_dim, max_seq_len) self.yarn_paired_head = YarnPairedHead(head_dim, max_seq_len) # there are only 50257 unique GPT-2 tokens; we extend to nearest multiple of 128 for efficiency. # suggested to me by @Grad62304977. this originates from Karpathy's experiments. use_fp8 = not os.environ.get("DISABLE_FP8", False) # Transposed weight storage for faster gradient accumulation self.lm_head = CastedLinearT(model_dim, vocab_size, use_fp8=use_fp8, x_s=100/448, w_s=1.6/448, grad_s=0.75/448) nn.init.normal_(self.lm_head.weight, mean=0, std=0.005) self.lm_head.weight.label = 'lm_head' self.embed = nn.Embedding(vocab_size, model_dim) self.embed.weight.label = 'embed' with torch.no_grad(): self.embed.weight.copy_(self.lm_head.weight.T) self.bigram_embed = nn.Embedding(args.bigram_vocab_size, model_dim) self.bigram_embed.weight.label = 'bigram_embed' nn.init.zeros_(self.bigram_embed.weight) # x0_lambdas separated out for different optimizer treatment (no beta smoothing) self.x0_lambdas = nn.Parameter(torch.zeros(num_layers)) self.x0_lambdas.label = 'x0_lambdas' pad = (-num_layers * 3 - 3) % dist.get_world_size() # updated: 3*num_layers instead of 4* self.scalars = nn.Parameter( torch.cat( [ 1.1 * torch.ones(num_layers), # resid lambdas. 1.1 init such that layer i weight is i^(num_layers-i). *[torch.tensor([0.5, 1.0]) for _ in range(num_layers)], # SA lambdas 0.1 * torch.ones(num_layers), # bigram lambdas torch.zeros(1), # smear_lambda 0.5*torch.ones(1), # backout_lambda -1.5 * torch.ones(1), # skip_lambda -> σ(-1.5) ≈ 0.18 torch.ones(pad), ] ) ) self.scalars.label = 'scalars' def forward(self, input_seq: Tensor, target_seq: Tensor, seqlens: Tensor, bigram_input_seq: Tensor, schedule_cfg: ForwardScheduleConfig): assert input_seq.ndim == 1 # unpack schedule_cfg mtp_weights, ws_short, ws_long = schedule_cfg.mtp_weights, schedule_cfg.ws_short, schedule_cfg.ws_long # set configs skip_connections = [] skip_in = [3] # long attention window on layer 3 skip_out = [6] # no attn op on layer 6 x_backout = None backout_layer = 7 # set lambdas resid_lambdas = self.scalars[: 1 * self.num_layers] x0_lambdas = self.x0_lambdas sa_lambdas = self.scalars[1 * self.num_layers: 3 * self.num_layers].view(-1, 2) bigram_lambdas = self.scalars[3 * self.num_layers: 4 * self.num_layers] smear_lambda = self.scalars[4 * self.num_layers] backout_lambda = self.scalars[4 * self.num_layers+1] skip_lambda = self.scalars[4 * self.num_layers+2] # set block masks and key shift short_bm = ws_short * args.block_size long_bm = ws_long * args.block_size bm_sizes = [short_bm, short_bm, short_bm, long_bm, short_bm, short_bm, None, short_bm, short_bm, short_bm, long_bm] assert len(bm_sizes) == self.num_layers key_offset = [b==long_bm for b in bm_sizes] # apply partial key offset to long windows # Embedding lookup - embed is synced from lm_head during tied phase by optimizer x = self.embed(input_seq) x0_bigram = self.bigram_embed(bigram_input_seq)[None] # Value embeddings - always computed (not precomputed) ve = [value_embed(input_seq) for value_embed in self.value_embeds] # 012 ... 012 structure on token value embeddings by @YouJiacheng, improved on @leloykun's U-net structure # dropping first layer updates this to .12 ... 012 ve = [ve[1], ve[2]] + [None] * (self.num_layers - 5) + [ve[0], ve[1], ve[2]] assert len(ve) == self.num_layers # smear token embed forward 1 position @classiclarryd smear_gate_out = smear_lambda * torch.sigmoid(self.smear_gate(x[1:, :self.smear_gate.weight.size(-1)])) x = torch.cat([x[:1], x[1:] + smear_gate_out * x[:-1]]) x = x0 = norm(x[None]) # unbind gate banks to avoid select_backwards kernel ag = [w.bfloat16() for w in self.attn_gate_bank.unbind(0)] veg = [w.bfloat16() for w in self.ve_gate_bank.unbind(0)] attn_gates = ag[:6] + [None] + ag[6:] ve_gates = [veg[0], veg[1]] + [None] * (self.num_layers - 5) + [veg[2], veg[3], veg[4]] assert len(attn_gates) == self.num_layers assert len(ve_gates) == self.num_layers # unbind weight banks to avoid select_backwards kernel attn_weights = self.attn_bank.unbind(0) # tuple of [4*dim, hdim] tensors mlp_fcs = self.mlp_bank[:, 0, :, :].unbind(0) # tuple of [mlp_hdim, dim] tensors mlp_projs = self.mlp_bank[:, 1, :, :].unbind(0) # tuple of [mlp_hdim, dim] tensors for i in range(self.num_layers): yarn = self.yarn_paired_head if i in self.paired_head_layers else self.yarn attn_args = AttnArgs( ve=ve[i], sa_lambdas=sa_lambdas[i], seqlens=seqlens, bm_size=bm_sizes[i], yarn=yarn, key_offset=key_offset[i], attn_gate_w=attn_gates[i], ve_gate_w=ve_gates[i] ) if i in skip_out: skip_gate_out = torch.sigmoid(skip_lambda) * 2 * torch.sigmoid(self.skip_gate(x0[..., :self.skip_gate.weight.size(-1)])) x = x + skip_gate_out * skip_connections.pop() if i == 0: x = (resid_lambdas[0] + x0_lambdas[0]) * x + bigram_lambdas[0] * x0_bigram else: x = resid_lambdas[i] * x + x0_lambdas[i] * x0 + bigram_lambdas[i] * x0_bigram # Get weights for this layer from banks qkvo_w = attn_weights[self.layer_to_attn_idx[i]] if i in self.layer_to_attn_idx else None c_fc = mlp_fcs[self.layer_to_mlp_idx[i]] if i in self.layer_to_mlp_idx else None c_proj = mlp_projs[self.layer_to_mlp_idx[i]] if i in self.layer_to_mlp_idx else None x = self.blocks[i](x, attn_args, qkvo_w, c_fc, c_proj) if i in skip_in: skip_connections.append(x) if i == backout_layer: x_backout = x # back out contributions from first 7 layers that are only required for downstream context and not direct prediction x -= backout_lambda * x_backout x = norm(x) logits = self.lm_head(x) # @Grad62304977 added tanh softcapping following Gemma 2 paper, @KoszarskyB reduced it from 30 to 15 # @YouJiacheng shifted it by +15 (2*sigmoid(2*x)=tanh(x)+1). @classiclarryd updated to 23*sigmoid((logits+5)/7.5) if self.training: losses = FusedSoftcappedCrossEntropy.apply(logits.view(-1, logits.size(-1)), target_seq, mtp_weights) loss = losses.sum() else: logits = 23 * torch.sigmoid((logits + 5) / 7.5) logits_for_loss = logits.float() loss = F.cross_entropy(logits_for_loss.view(-1, logits_for_loss.size(-1)), target_seq, reduction="mean") return loss # ----------------------------------------------------------------------------- # Distributed data loader def _load_data_shard(file: Path): header = torch.from_file(str(file), False, 256, dtype=torch.int32) # header is 256 int32 assert header[0] == 20240520, "magic number mismatch in the data .bin file" assert header[1] == 1, "unsupported version" num_tokens = int(header[2]) # number of tokens (claimed) with file.open("rb", buffering=0) as f: tokens = torch.empty(num_tokens, dtype=torch.uint16, pin_memory=True) # avoid pin_memory copy by @YouJiacheng f.seek(256 * 4) nbytes = f.readinto(tokens.numpy()) # avoid bytes->array copy by @YouJiacheng assert nbytes == 2 * num_tokens, "number of tokens read does not match header" return tokens BOS_ID = 50256 class BOSFinder: # Helper for getting sequences that start at the beginning of documents by @varunneal based on work by @classiclarryd def __init__(self, tokens: Tensor, world_size: int = 1, quickload: bool = False): # Precompute BOS positions once per shard self.tokens=tokens self.size = tokens.numel() self.quickload = quickload if quickload: # only scan first 4 million tokens, then kickoff async thread to scan rest self.bos_idx = (tokens[:4_000_000] == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self.thread = None self.ready = threading.Event() self.start() else: self.bos_idx = (tokens == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self.i = 0 self.world_size = world_size self.batch_iter = 0 def _load(self): self.bos_idx_async = (self.tokens == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self.ready.set() def start(self): self.ready.clear() self.thread = threading.Thread(target=self._load) self.thread.start() def get(self): if self.thread: self.ready.wait() self.thread.join() self.bos_idx = self.bos_idx_async def next_batch(self, num_tokens_local: int, max_seq_len: int): # if quickload was used, repoint to the full dataset after 5 batches if self.quickload and self.batch_iter==5: self.get() n = len(self.bos_idx) starts = [[] for _ in range(self.world_size)] ends = [[] for _ in range(self.world_size)] idx = self.i for r in range(self.world_size): cur_len = 0 while cur_len <= num_tokens_local: if idx >= n: raise StopIteration(f"Insufficient BOS ahead; hit tail of shard.") cur = self.bos_idx[idx] starts[r].append(cur) end = min(self.bos_idx[idx + 1] if idx + 1 < n else self.size, cur + max_seq_len, cur + num_tokens_local - cur_len + 1) ends[r].append(end) cur_len += end - cur idx += 1 assert cur_len == num_tokens_local + 1 self.i = idx self.batch_iter+=1 return starts, ends class DataPreloader: # Helper for asynchronously loading next shard and indexing bos tokens def __init__(self, file_iter, world_size: int = 1): self.file_iter = file_iter self.world_size = world_size self.thread = None self.data = None self.ready = threading.Event() def _load(self): tokens = _load_data_shard(next(self.file_iter)) self.data = (tokens, BOSFinder(tokens, self.world_size)) self.ready.set() def start(self): self.ready.clear() self.thread = threading.Thread(target=self._load) self.thread.start() def get(self): if self.thread: self.ready.wait() self.thread.join() return self.data def get_bigram_hash(x): """ Computes bigram hash for each position using [prev_token, curr_token]. Multiply by arbitary large ints to get even spread over int32 range. Position 0 is mapped to the reserved index (vocab_size - 1). BOS_tokens within the batch will hash based on last token of prior doc. Masking this ran slower and showed no improvement. """ rand_int_1 = 36313 rand_int_2 = 27191 mod = args.bigram_vocab_size-1 x = x.to(torch.int32).clone() x[0] = mod x[1:] = torch.bitwise_xor(rand_int_1 * x[1:], rand_int_2 * x[:-1]) % mod return x def distributed_data_generator(filename_pattern: str, num_tokens: int, max_seq_len: int, grad_accum_steps: int = 1, align_to_bos: bool = True): # align_to_bos: each sequence begins with Beginning of Sequence token, sequences truncated to max_seq_len rank = dist.get_rank() if dist.is_initialized() else 0 world_size = dist.get_world_size() if dist.is_initialized() else 1 assert num_tokens % (world_size * grad_accum_steps) == 0, "Batch size must be divisible by world size" num_tokens = num_tokens // grad_accum_steps files = [Path(file) for file in sorted(glob.glob(filename_pattern))] if not files: raise FileNotFoundError(f"No files found for pattern: {filename_pattern}") file_iter = iter(files) # Use itertools.cycle(files) for multi-epoch training tokens = _load_data_shard(next(file_iter)) if align_to_bos: finder = BOSFinder(tokens, world_size=world_size, quickload=True) preloader = DataPreloader(file_iter, world_size) preloader.start() else: pos = 0 # for unaligned case while True: num_tokens_local = num_tokens // world_size max_num_docs = next_multiple_of_n(num_tokens_local // 300, n=128) # median doc length is ~400 if align_to_bos: try: seq_starts, seq_ends = finder.next_batch(num_tokens_local, max_seq_len) start_idxs, end_idxs = torch.tensor(seq_starts[rank]), torch.tensor(seq_ends[rank]) except StopIteration: # This shard is exhausted, load the next one in the next loop iteration. tokens, finder = preloader.get() preloader.start() continue buf = torch.cat([tokens[i:j] for i, j in zip(start_idxs, end_idxs)]) _inputs = buf[:-1] _targets = buf[1:] end_idxs[-1] -= 1 # last document was too long to account for _targets offset cum_lengths = (end_idxs - start_idxs).cumsum(0) else: if pos + num_tokens + 1 >= len(tokens): # should not occur for val data tokens, pos = _load_data_shard(next(file_iter)), 0 pos_local = pos + rank * num_tokens_local buf = tokens[pos_local: pos_local + num_tokens_local + 1] _inputs = buf[:-1].view(num_tokens_local, ) _targets = buf[1:].view(num_tokens_local, ) cum_lengths = torch.nonzero(_inputs == BOS_ID)[:, 0] pos += num_tokens _cum_lengths = torch.full((max_num_docs,), num_tokens_local) _cum_lengths[0] = 0 _cum_lengths[1:len(cum_lengths) + 1] = cum_lengths # Cast to int32 on CPU before transfer to avoid dtype conversion during .to() _inputs = _inputs.to(dtype=torch.int32) _targets = _targets.to(dtype=torch.int64) _cum_lengths = _cum_lengths.to(dtype=torch.int32) _bigram_inputs = get_bigram_hash(_inputs) new_params = yield ( _inputs.to(device="cuda", non_blocking=True), _targets.to(device="cuda", non_blocking=True), _cum_lengths.to(device="cuda", non_blocking=True), _bigram_inputs.to(device="cuda", non_blocking=True) ) if new_params is not None: # makes it possible for generator to receive new (num_tokens, max_seq_len, grad_accum_steps) via .send() new_num_tokens, new_max_seq_len, new_grad_accum_steps = new_params assert new_num_tokens % (world_size * new_grad_accum_steps) == 0, "Num tokens must be divisible by world size" num_tokens = new_num_tokens // new_grad_accum_steps max_seq_len = new_max_seq_len # ----------------------------------------------------------------------------- # Training Management def get_bs(step: int): if step >= args.num_scheduled_iterations: return args.train_bs_extension x = step / args.num_scheduled_iterations bs_idx = int(len(args.train_bs_schedule) * x) return args.train_bs_schedule[bs_idx] def get_ws(step: int): # set short window size to half of long window size # Higher ws on "extension" steps if step >= args.num_scheduled_iterations: return args.ws_final // 2, args.ws_final x = step / args.num_scheduled_iterations assert 0 <= x < 1 ws_idx = int(len(args.ws_schedule) * x) return args.ws_schedule[ws_idx] // 2, args.ws_schedule[ws_idx] # learning rate schedule: tied to batch size schedule, with cooldown at the end. def get_lr(step: int): if step > args.num_scheduled_iterations: return 0.1 lr_max = 1.0 x = step / args.num_scheduled_iterations if x > 1/3: lr_max = 1.52 # (16/8)**0.6 if x > 2/3: lr_max = 1.73 # (24/8)**0.5 if x >= 1 - args.cooldown_frac: w = (1 - x) / args.cooldown_frac lr = lr_max * w + (1 - w) * 0.1 return lr return lr_max def get_muon_momentum(step: int, muon_warmup_steps=300, muon_cooldown_steps=50, momentum_min=0.85, momentum_max=0.95): # warmup phase: linearly increase momentum from min to max # cooldown phase: linearly decrease momentum from max to min momentum_cd_start = args.num_iterations - muon_cooldown_steps if step < muon_warmup_steps: frac = step / muon_warmup_steps momentum = momentum_min + frac * (momentum_max - momentum_min) elif step > momentum_cd_start: frac = (step - momentum_cd_start) / muon_cooldown_steps momentum = momentum_max - frac * (momentum_max - momentum_min) else: momentum = momentum_max return momentum class TrainingManager(): """ Manages the NorMuonAndAdam for all parameters with explicit ordering. Notable Features: 1. Scalars are given higher momentum terms to smooth learning @ChrisJMcCormick 2. Adam optimizers are only stepped on odd steps @classiclarryd 3. Explicit scatter_order and work_order for communication scheduling (no backward hooks) 4. Muon has a linear momentum warmup and cooldown schedule 5. Learning rates follow a linear decay schedule 6. Embed is tied to lm_head until split step (2/3 of training), then untied @classiclarryd Manages model architecture, data, and target that changes during training Notable Features: 1. Multi Token Prediction schedule of [1, 0.5, 0.25->0] -> [1, 0.5->0] -> [1] @varunneal 2. Sliding Attention window schedule of [1,3] -> [3,7] -> [5,11] -> [6,13] 3. YaRN updates to RoPE on window changes 4. Split embed and lm_head at 2/3 of training (weights and optimizer state copied) 5. Batch size schedule of 8 -> 16 -> 24 6. Post training extension of long windows from 13 to 20 """ def __init__(self, model): self.mtp_weights_schedule = self._build_mtp_schedule() self.model = model # - Ordering dictates when to launch reduce/reduce_scatter operations # - "sharded" parameters use reduce_scatter/all_gather and "replicated" ones use all_reduce # - lr_mul and wd_mul are per-parameter learning rate and weight decay multipliers self.param_table = { "attn": {"optim": "normuon", "comms": "sharded", "adam_betas": None}, "mlp": {"optim": "normuon", "comms": "sharded", "adam_betas": None}, "scalars": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99], "lr_mul": 5.0, "wd_mul": 0.0}, "ve0": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve1": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "ve2": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "bigram_embed": {"optim": "adam", "comms": "sharded", "adam_betas": [0.75, 0.95], "lr_mul": 75., "wd_mul": 5.0}, "smear_gate": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99], "lr_mul": 0.01, "wd_mul": 0.0}, "skip_gate": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99], "lr_mul": 0.05, "wd_mul": 0.0}, "attn_gate_bank": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99]}, "ve_gate_bank": {"optim": "adam", "comms": "replicated", "adam_betas": [0.9, 0.99]}, "x0_lambdas": {"optim": "adam", "comms": "replicated", "adam_betas": [0.65, 0.95], "lr_mul": 5.0, "wd_mul": 0.0}, "lm_head": {"optim": "adam", "comms": "sharded", "adam_betas": [0.5, 0.95], "wd_mul": 150.}, "embed": {"optim": "adam", "comms": "sharded", "adam_betas": [0.5, 0.95], "wd_mul": 150.}, } # - Process smaller/faster params first while large reduces complete # - lm_head must complete before embed sync (when tied) self.work_order = [ "scalars", "smear_gate", "skip_gate", "attn_gate_bank", "ve_gate_bank", "x0_lambdas", # Small, fast "ve0", "ve1", "ve2", "bigram_embed", # Medium "lm_head", "embed", # lm_head must complete before embed sync (when tied) "attn", "mlp", # Large, polar express - process last to maximize overlap ] adam_defaults = dict( lr=0.008, eps=1e-10, weight_decay=0.005, ) normuon_defaults = dict( lr=0.023, momentum=0.95, beta2=0.95, weight_decay=1.2, ) self.optimizer = NorMuonAndAdam( model.named_parameters(), param_table=self.param_table, scatter_order=list(self.param_table.keys()), # Dict order defines scatter priority work_order=self.work_order, adam_defaults=adam_defaults, normuon_defaults=normuon_defaults, ) # Split embed from lm_head at 2/3 of training (on an odd step so Adam updates) self.split_step = math.ceil(args.split_embed_frac * args.num_scheduled_iterations) | 1 self.reset() def _build_mtp_schedule(self): # Precompute MTP weights for all steps to avoid tensor allocation during training # Schedule: [1, 0.5, 0.25->0] -> [1, 0.5->0] -> [1] mtp_weights_schedule = [] for s in range(args.num_iterations + 1): x = s / args.num_scheduled_iterations if x < 1/3: w = [1.0, 0.5, 0.25 * (1 - 3*x)] elif x < 2/3: w = [1.0, 0.5 * (1 - (3*x - 1))] else: w = [1.0] mtp_weights_schedule.append(torch.tensor(w, device=device)) return mtp_weights_schedule def apply_final_ws_ext(self): self.ws_long = args.ws_validate_post_yarn_ext def get_forward_args(self): return ForwardScheduleConfig( mtp_weights = self.mtp_weights, ws_short = self.ws_short, ws_long = self.ws_long ) def _is_adam_step(self, step: int): """Adam params are only updated on odd steps.""" return step % 2 == 1 def get_transition_steps(self): transition_steps = [] ws_short, ws_long = get_ws(0) for step in range(1, args.num_iterations): ws_short, new_ws_long = get_ws(step) if new_ws_long != ws_long: transition_steps.append(step) ws_long = new_ws_long return transition_steps def advance_schedule(self, step: int): self.ws_short, new_ws_long = get_ws(step) if new_ws_long != self.ws_long: self.model.yarn.apply(self.ws_long, new_ws_long) self.model.yarn_paired_head.apply(self.ws_long, new_ws_long) new_batch_size = get_bs(step) if new_batch_size != self.batch_size: self.train_loader_send_args = (new_batch_size, args.train_max_seq_len, grad_accum_steps) self.batch_size = new_batch_size else: self.train_loader_send_args = None self.ws_long = new_ws_long self.mtp_weights = self.mtp_weights_schedule[step] def step_optimizers(self, step: int): step_lr = get_lr(step) muon_momentum = get_muon_momentum(step) do_adam = self._is_adam_step(step) # Update learning rates and momentum for all params for param, p_cfg in self.optimizer.param_cfgs.items(): p_cfg.lr = p_cfg.initial_lr * step_lr if p_cfg.optim == "normuon": p_cfg.momentum = muon_momentum # Step optimizer with do_adam flag self.optimizer.step(do_adam=do_adam) # At split step: copy lm_head optimizer state to embed and mark as split if step == self.split_step: self.optimizer.copy_lm_state_to_embed() def reset(self, state=None): if state is not None: self.optimizer.load_state_dict(state) # Reset NorMuon momentum buffers and split_embed state self.optimizer.reset() self.ws_short, self.ws_long = get_ws(0) self.batch_size = get_bs(0) self.model.yarn.reset() self.model.yarn_paired_head.reset() def get_state(self): return copy.deepcopy(self.optimizer.state_dict()) # ----------------------------------------------------------------------------- # int main @dataclass class Hyperparameters: # data train_files: str = "data/fineweb10B/fineweb_train_*.bin" # input .bin to train on val_files: str = "data/fineweb10B/fineweb_val_*.bin" # input .bin to eval validation loss on val_tokens: int = 10485760 # how many tokens of validation data? it's important to keep this fixed for consistent comparisons # batch sizes train_bs_schedule: tuple = (8 * 2048 * 8, 16 * 2048 * 8, 24 * 2048 * 8) train_bs_extension: int = 24 * 2048 * 8 train_max_seq_len: int = 128 * 16 val_batch_size: int = 4 * 64 * 1024 * 8 # optimization num_scheduled_iterations: int = 1560 # number of steps to complete lr and ws schedule num_extension_iterations: int = 40 # number of steps to continue training at final lr and ws num_iterations: int = num_scheduled_iterations + num_extension_iterations cooldown_frac: float = 0.55 # fraction of num_scheduled_iterations spent cooling down the learning rate split_embed_frac: float = 2/3 # fraction of training when embeddings split from lm_head # evaluation and logging run_id: str = f"{uuid.uuid4()}" val_loss_every: int = 250 # every how many steps to evaluate val loss? 0 for only at the end save_checkpoint: bool = False # attention masking block_size: int = 128 ws_schedule: tuple = (3, 7, 11) ws_final: int = 13 # increase final validation ws, used for YaRN extension and short window size @classiclarryd ws_validate_post_yarn_ext: int = 20 # extend long windows out even further after applying YaRN # bigram hash embedding bigram_vocab_size = 50304 * 5 args = Hyperparameters() data_path = os.environ.get("DATA_PATH", ".") args.train_files = os.path.join(data_path, args.train_files) args.val_files = os.path.join(data_path, args.val_files) # torchrun sets these env variables rank = int(os.environ["RANK"]) world_size = int(os.environ["WORLD_SIZE"]) assert 8 % world_size == 0, "world_size must be a divisor of 8" grad_accum_steps = 8 // world_size assert torch.cuda.is_available() device = torch.device("cuda", int(os.environ["LOCAL_RANK"])) torch.cuda.set_device(device) dist.init_process_group(backend="nccl", device_id=device) dist.barrier() master_process = (rank == 0) # this process will do logging, checkpointing etc. # begin logging logfile = None if master_process: run_id = args.run_id os.makedirs("logs", exist_ok=True) logfile = f"logs/{run_id}.txt" print(logfile) def print0(s, console=False): if master_process: with open(logfile, "a") as f: if console: print(s) print(s, file=f) # begin by printing this file (the Python code) print0(code) print0("="*100) # log information about the hardware/software environment this is running on print0(f"Running Python {sys.version}") print0(f"Running PyTorch {torch.version.__version__} compiled for CUDA {torch.version.cuda}") print0(f"Running Triton version {triton.__version__}") def nvidia_smi(): import subprocess # avoid top level import return subprocess.run(["nvidia-smi"], stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True).stdout print0(nvidia_smi()) print0("="*100) model: nn.Module = GPT( vocab_size=50257, num_layers=11, num_heads=6, head_dim=128, model_dim=768, max_seq_len=args.val_batch_size // (grad_accum_steps * world_size) ).cuda() for m in model.modules(): if isinstance(m, (nn.Embedding, nn.Linear)): m.weight.data = m.weight.data.bfloat16() model.attn_gate_bank.data = model.attn_gate_bank.data.bfloat16() model.ve_gate_bank.data = model.ve_gate_bank.data.bfloat16() model.attn_bank.data = model.attn_bank.data.bfloat16() model.mlp_bank.data = model.mlp_bank.data.bfloat16() for param in model.parameters(): dist.broadcast(param.detach(), 0) model: nn.Module = torch.compile(model, dynamic=False, fullgraph=True) training_manager = TrainingManager(model) ######################################## # Warmup kernels # ######################################## print0("Compiling model and warming up kernels (~7 minutes on first execution)", console=True) # Warmup the training kernels, then re-initialize the state so we aren't cheating initial_state = dict(model=copy.deepcopy(model.state_dict()), optimizer=training_manager.get_state()) # save the initial state train_loader = distributed_data_generator(args.train_files, args.train_bs_schedule[0], args.train_max_seq_len, grad_accum_steps=grad_accum_steps) val_loader = distributed_data_generator(args.val_files, args.val_batch_size, -1, grad_accum_steps=grad_accum_steps, align_to_bos=False) transition_steps = training_manager.get_transition_steps() # first few steps plus transitions warmup_steps = sorted({0, 1, 2} | set(s + offset for s in transition_steps for offset in [-1, 0, 1] if s + offset >= 0)) print0(f"Sampling steps {warmup_steps} for warmup", console=True) for step in warmup_steps: training_manager.advance_schedule(step) model.eval() with torch.no_grad(): inputs, targets, cum_seqlens, bigram_inputs = next(val_loader) model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) model.train() for idx in range(grad_accum_steps): send_args = training_manager.train_loader_send_args inputs, targets, cum_seqlens, bigram_inputs = train_loader.send(send_args) (model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) / grad_accum_steps).backward() training_manager.step_optimizers(step) print0("Resetting Model", console=True) model.zero_grad(set_to_none=True) model.load_state_dict(initial_state["model"]) training_manager.reset(initial_state["optimizer"]) del val_loader, train_loader, initial_state model.train() ######################################## # Training and validation # ######################################## train_loader = distributed_data_generator(args.train_files, args.train_bs_schedule[0], args.train_max_seq_len, grad_accum_steps=grad_accum_steps) gc.collect() training_time_ms = 0 # start the clock torch.cuda.synchronize() t0 = time.perf_counter() # begin training train_steps = args.num_iterations for step in range(train_steps + 1): last_step = (step == train_steps) training_manager.advance_schedule(step) # --------------- VALIDATION SECTION ----------------- if last_step or (args.val_loss_every > 0 and step % args.val_loss_every == 0): if last_step: training_manager.apply_final_ws_ext() # stop the clock torch.cuda.synchronize() training_time_ms += 1000 * (time.perf_counter() - t0) model.eval() assert args.val_tokens % args.val_batch_size == 0 val_steps = grad_accum_steps * args.val_tokens // args.val_batch_size val_loader = distributed_data_generator(args.val_files, args.val_batch_size, -1, grad_accum_steps=grad_accum_steps, align_to_bos=False) val_loss = 0 with torch.no_grad(): for _ in range(val_steps): inputs, targets, cum_seqlens, bigram_inputs = next(val_loader) val_loss += model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) val_loss /= val_steps del val_loader dist.reduce(val_loss, 0, op=dist.ReduceOp.AVG) print0(f"step:{step}/{train_steps} val_loss:{val_loss:.4f} train_time:{training_time_ms:.0f}ms step_avg:{training_time_ms/max(step, 1):.2f}ms", console=True) model.train() # start the clock again torch.cuda.synchronize() t0 = time.perf_counter() if last_step: if master_process and args.save_checkpoint: log = dict(step=step, code=code, model=model.state_dict(), optimizer=training_manager.get_state()) os.makedirs(f"logs/{run_id}", exist_ok=True) torch.save(log, f"logs/{run_id}/state_step{step:06d}.pt") # the last step only has the validation loop, so break to avoid training break # --------------- TRAINING SECTION ----------------- for idx in range(grad_accum_steps): inputs, targets, cum_seqlens, bigram_inputs = train_loader.send(training_manager.train_loader_send_args) (model(inputs, targets, cum_seqlens, bigram_inputs, training_manager.get_forward_args()) / grad_accum_steps).backward() training_manager.step_optimizers(step) # logging approx_training_time_ms = training_time_ms + 1000 * (time.perf_counter() - t0) print0(f"step:{step+1}/{train_steps} train_time:{approx_training_time_ms:.0f}ms step_avg:{approx_training_time_ms/(step + 1):.2f}ms", console=True) print0(f"peak memory allocated: {torch.cuda.max_memory_allocated() // 1024 // 1024} MiB " f"reserved: {torch.cuda.max_memory_reserved() // 1024 // 1024} MiB", console=True) dist.destroy_process_group() ---------------------------------------- # triton_kernels.py ---------------------------------------- import torch import triton import triton.language as tl from triton.tools.tensor_descriptor import TensorDescriptor # ----------------------------------------------------------------------------- # Triton kernel for symmetric matrix multiplication by @byronxu99 def _get_autotune_configs(): return [ triton.Config( { "BLOCK_SIZE_M": bm, "BLOCK_SIZE_N": bn, "BLOCK_SIZE_K": bk, "GROUP_SIZE_M": 8, "LOWER_UPPER": 1, }, num_stages=stages, num_warps=warps, ) for bm in [64, 128] for bn in [64, 128, 256] for bk in [64, 128] for stages, warps in [(3, 4), (3, 8), (4, 4)] if bm // bn <= 2 and bn // bm <= 2 ] @triton.jit def _pid_to_block( pid, M, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, GROUP_SIZE_M: tl.constexpr, ): # Split output matrix into blocks of size (BLOCK_SIZE_M, BLOCK_SIZE_N) num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) num_pid_n = tl.cdiv(M, BLOCK_SIZE_N) # Map PID to a single matrix in batch batch_idx = pid // (num_pid_m * num_pid_n) pid = pid % (num_pid_m * num_pid_n) # Map PID to 2D grid of blocks pid_m = pid // num_pid_n pid_n = pid % num_pid_n pid_m, pid_n = tl.swizzle2d(pid_m, pid_n, num_pid_m, num_pid_n, GROUP_SIZE_M) m_idx = pid_m * BLOCK_SIZE_M n_idx = pid_n * BLOCK_SIZE_N return batch_idx, m_idx, n_idx @triton.autotune( configs=_get_autotune_configs(), key=["M", "K", "a_stride_r", "a_stride_c", "c_stride_r", "c_stride_c"], ) @triton.jit def XXT_kernel( A_ptr, C_ptr, M, K, a_stride_b, a_stride_r, a_stride_c, c_stride_b, c_stride_r, c_stride_c, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, LOWER_UPPER: tl.constexpr, ): pid = tl.program_id(axis=0) batch_idx, m_idx, n_idx = _pid_to_block( pid, M, BLOCK_SIZE_M, BLOCK_SIZE_N, GROUP_SIZE_M ) # Skip blocks that don't need to be computed skip_block_below_diag = (LOWER_UPPER == 0) and (n_idx + BLOCK_SIZE_N <= m_idx) skip_block_above_diag = (LOWER_UPPER != 0) and (m_idx + BLOCK_SIZE_M <= n_idx) if skip_block_below_diag or skip_block_above_diag: return # Index into one matrix of batch A_ptr += batch_idx * a_stride_b C_ptr += batch_idx * c_stride_b # Create pointer arrays for A and A.T offs_m = (m_idx + tl.arange(0, BLOCK_SIZE_M)) % M offs_n = (n_idx + tl.arange(0, BLOCK_SIZE_N)) % M offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = A_ptr + (offs_m[:, None] * a_stride_r + offs_k[None, :] * a_stride_c) at_ptrs = A_ptr + (offs_k[:, None] * a_stride_c + offs_n[None, :] * a_stride_r) accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) # Accumulate over blocks of K for k in tl.range(0, tl.cdiv(K, BLOCK_SIZE_K)): a = tl.load(a_ptrs, mask=offs_k[None, :] < K - k * BLOCK_SIZE_K, other=0.0) at = tl.load(at_ptrs, mask=offs_k[:, None] < K - k * BLOCK_SIZE_K, other=0.0) accumulator = tl.dot(a, at, accumulator) a_ptrs += BLOCK_SIZE_K * a_stride_c at_ptrs += BLOCK_SIZE_K * a_stride_c out_dtype = C_ptr.dtype.element_ty output = accumulator.to(out_dtype) # Store block of C offs_cm = m_idx + tl.arange(0, BLOCK_SIZE_M) offs_cn = n_idx + tl.arange(0, BLOCK_SIZE_N) c_ptrs = C_ptr + (offs_cm[:, None] * c_stride_r + offs_cn[None, :] * c_stride_c) c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < M) tl.store(c_ptrs, output, mask=c_mask) # Store block of C mirrored across the diagonal c_ptrs_t = C_ptr + (offs_cn[:, None] * c_stride_r + offs_cm[None, :] * c_stride_c) c_mask_t = (offs_cn[:, None] < M) & (offs_cm[None, :] < M) tl.store(c_ptrs_t, output.T, mask=c_mask_t) def XXT(A: torch.Tensor, out: torch.Tensor): """ Launch Triton kernel to compute C = A @ A.T """ assert A.ndim == 2 or A.ndim == 3 M, K = A.shape[-2:] assert out.size(-2) == M, "Output matrix has incorrect shape" assert out.size(-1) == M, "Output matrix has incorrect shape" batch_size = A.size(0) if A.ndim == 3 else 1 input_batch_stride = A.stride(0) if A.ndim == 3 else 0 output_batch_stride = out.stride(0) if out.ndim == 3 else 0 grid = lambda meta: ( batch_size * triton.cdiv(M, meta["BLOCK_SIZE_M"]) * triton.cdiv(M, meta["BLOCK_SIZE_N"]), ) XXT_kernel[grid]( A_ptr=A, C_ptr=out, M=M, K=K, a_stride_b=input_batch_stride, a_stride_r=A.stride(-2), a_stride_c=A.stride(-1), c_stride_b=output_batch_stride, c_stride_r=out.stride(-2), c_stride_c=out.stride(-1), ) return out @triton.autotune( configs=_get_autotune_configs(), key=["M", "a_stride_r", "a_stride_c", "c_stride_r", "c_stride_c"], ) @triton.jit def ba_plus_cAA_kernel( A_ptr, C_ptr, M, a_stride_b, a_stride_r, a_stride_c, c_stride_b, c_stride_r, c_stride_c, alpha, beta, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, LOWER_UPPER: tl.constexpr, ): # This is mostly duplicated from XXT_kernel, but also loads and adds a block of A # Performance is slightly slower than XXT_kernel, so we use two separate kernels pid = tl.program_id(axis=0) batch_idx, m_idx, n_idx = _pid_to_block( pid, M, BLOCK_SIZE_M, BLOCK_SIZE_N, GROUP_SIZE_M ) # Skip blocks that don't need to be computed skip_block_below_diag = (LOWER_UPPER == 0) and (n_idx + BLOCK_SIZE_N <= m_idx) skip_block_above_diag = (LOWER_UPPER != 0) and (m_idx + BLOCK_SIZE_M <= n_idx) if skip_block_below_diag or skip_block_above_diag: return # Index into one matrix of batch A_ptr += batch_idx * a_stride_b C_ptr += batch_idx * c_stride_b # Create pointer arrays for A and A.T offs_m = (m_idx + tl.arange(0, BLOCK_SIZE_M)) % M offs_n = (n_idx + tl.arange(0, BLOCK_SIZE_N)) % M offs_k = tl.arange(0, BLOCK_SIZE_K) a_ptrs = A_ptr + (offs_m[:, None] * a_stride_r + offs_k[None, :] * a_stride_c) at_ptrs = A_ptr + (offs_k[:, None] * a_stride_c + offs_n[None, :] * a_stride_r) accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) # Accumulate over blocks of K for k in tl.range(0, tl.cdiv(M, BLOCK_SIZE_K)): a = tl.load(a_ptrs, mask=offs_k[None, :] < M - k * BLOCK_SIZE_K, other=0.0) at = tl.load(at_ptrs, mask=offs_k[:, None] < M - k * BLOCK_SIZE_K, other=0.0) accumulator = tl.dot(a, at, accumulator) a_ptrs += BLOCK_SIZE_K * a_stride_c at_ptrs += BLOCK_SIZE_K * a_stride_c # Load block of A to add (corresponds to the current block of C) offs_am = m_idx + tl.arange(0, BLOCK_SIZE_M) offs_an = n_idx + tl.arange(0, BLOCK_SIZE_N) a_add_ptrs = A_ptr + (offs_am[:, None] * a_stride_r + offs_an[None, :] * a_stride_c) a_add_mask = (offs_am[:, None] < M) & (offs_an[None, :] < M) a_add = tl.load(a_add_ptrs, mask=a_add_mask, other=0.0).to(tl.float32) # Apply alpha and beta accumulator *= alpha accumulator += a_add * beta out_dtype = C_ptr.dtype.element_ty output = accumulator.to(out_dtype) # Store block of C offs_cm = m_idx + tl.arange(0, BLOCK_SIZE_M) offs_cn = n_idx + tl.arange(0, BLOCK_SIZE_N) c_ptrs = C_ptr + (offs_cm[:, None] * c_stride_r + offs_cn[None, :] * c_stride_c) c_mask = (offs_cm[:, None] < M) & (offs_cn[None, :] < M) tl.store(c_ptrs, output, mask=c_mask) # Store block of C mirrored across the diagonal c_ptrs_t = C_ptr + (offs_cn[:, None] * c_stride_r + offs_cm[None, :] * c_stride_c) c_mask_t = (offs_cn[:, None] < M) & (offs_cm[None, :] < M) tl.store(c_ptrs_t, output.T, mask=c_mask_t) def ba_plus_cAA(A: torch.Tensor, alpha: float, beta: float, out: torch.Tensor): """ Launch Triton kernel to compute C = alpha * A @ A.T + beta * A """ assert A.ndim == 2 or A.ndim == 3 M, K = A.shape[-2:] assert M == K, "Input matrix must be square" assert out.size(-2) == M assert out.size(-1) == M batch_size = A.size(0) if A.ndim == 3 else 1 input_batch_stride = A.stride(0) if A.ndim == 3 else 0 output_batch_stride = out.stride(0) if out.ndim == 3 else 0 grid = lambda meta: ( batch_size * triton.cdiv(M, meta["BLOCK_SIZE_M"]) * triton.cdiv(M, meta["BLOCK_SIZE_N"]), ) ba_plus_cAA_kernel[grid]( A_ptr=A, C_ptr=out, M=M, a_stride_b=input_batch_stride, a_stride_r=A.stride(-2), a_stride_c=A.stride(-1), c_stride_b=output_batch_stride, c_stride_r=out.stride(-2), c_stride_c=out.stride(-1), alpha=alpha, beta=beta, ) return out # ----------------------------------------------------------------------------- # Triton kernel for MLP: relu(x @ W1.T)^2, by @andrewbriand, @jrauvola @triton.jit def linear_relu_square_kernel(a_desc, b_desc, c_desc, aux_desc, M, N, K, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, NUM_SMS: tl.constexpr, FORWARD: tl.constexpr, ): dtype = tl.bfloat16 start_pid = tl.program_id(axis=0) num_pid_m = tl.cdiv(M, BLOCK_SIZE_M) num_pid_n = tl.cdiv(N, BLOCK_SIZE_N) k_tiles = tl.cdiv(K, BLOCK_SIZE_K) num_tiles = num_pid_m * num_pid_n tile_id_c = start_pid - NUM_SMS num_pid_in_group = GROUP_SIZE_M * num_pid_n for tile_id in tl.range(start_pid, num_tiles, NUM_SMS, flatten=True): pid_m = tile_id // num_pid_n pid_n = tile_id % num_pid_n offs_am = pid_m * BLOCK_SIZE_M offs_bn = pid_n * BLOCK_SIZE_N accumulator = tl.zeros((BLOCK_SIZE_M, BLOCK_SIZE_N), dtype=tl.float32) for ki in range(k_tiles): offs_k = ki * BLOCK_SIZE_K a = a_desc.load([offs_am, offs_k]) b = b_desc.load([offs_bn, offs_k]) accumulator = tl.dot(a, b.T, accumulator) tile_id_c += NUM_SMS pid_m = tile_id // num_pid_n pid_n = tile_id % num_pid_n offs_am_c = pid_m * BLOCK_SIZE_M offs_bn_c = pid_n * BLOCK_SIZE_N acc = tl.reshape(accumulator, (BLOCK_SIZE_M, 2, BLOCK_SIZE_N // 2)) acc = tl.permute(acc, (0, 2, 1)) acc0, acc1 = tl.split(acc) c0 = acc0.to(dtype) if not FORWARD: c0_pre = aux_desc.load([offs_am_c, offs_bn_c]) c0 = 2 * c0 * tl.where(c0_pre > 0, c0_pre, 0) c_desc.store([offs_am_c, offs_bn_c], c0) if FORWARD: c0_post = tl.maximum(c0, 0) c0_post = c0_post * c0_post aux_desc.store([offs_am_c, offs_bn_c], c0_post) c1 = acc1.to(dtype) if not FORWARD: c1_pre = aux_desc.load([offs_am_c, offs_bn_c + BLOCK_SIZE_N // 2]) c1 = 2 * c1 * tl.where(c1_pre > 0, c1_pre, 0) c_desc.store([offs_am_c, offs_bn_c + BLOCK_SIZE_N // 2], c1) if FORWARD: c1_post = tl.maximum(c1, 0) c1_post = c1_post * c1_post aux_desc.store([offs_am_c, offs_bn_c + BLOCK_SIZE_N // 2], c1_post) def linear_relu_square(a, b, aux=None): M, K = a.shape N, K = b.shape dtype = a.dtype c = torch.empty((M, N), device=a.device, dtype=dtype) FORWARD = False if aux is None: FORWARD = True aux = torch.empty((M, N), device=a.device, dtype=dtype) NUM_SMS = torch.cuda.get_device_properties("cuda").multi_processor_count BLOCK_SIZE_M = 128 BLOCK_SIZE_N = 256 BLOCK_SIZE_K = 64 num_stages = 4 if FORWARD else 3 num_warps = 8 a_desc = TensorDescriptor.from_tensor(a, [BLOCK_SIZE_M, BLOCK_SIZE_K]) b_desc = TensorDescriptor.from_tensor(b, [BLOCK_SIZE_N, BLOCK_SIZE_K]) c_desc = TensorDescriptor.from_tensor(c, [BLOCK_SIZE_M, BLOCK_SIZE_N // 2]) aux_desc = TensorDescriptor.from_tensor(aux, [BLOCK_SIZE_M, BLOCK_SIZE_N // 2]) def grid(META): return (min( NUM_SMS, triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(N, BLOCK_SIZE_N), ), ) linear_relu_square_kernel[grid]( a_desc, b_desc, c_desc, aux_desc, M, N, K, BLOCK_SIZE_M=BLOCK_SIZE_M, BLOCK_SIZE_N=BLOCK_SIZE_N, BLOCK_SIZE_K=BLOCK_SIZE_K, GROUP_SIZE_M=1, NUM_SMS=NUM_SMS, FORWARD=FORWARD, num_stages=num_stages, num_warps=num_warps ) if FORWARD: return c, aux else: return c class FusedLinearReLUSquareFunction(torch.autograd.Function): @staticmethod def forward(ctx, x, W1, W2): pre, post = linear_relu_square(x.view((-1, x.shape[-1])), W1) x3 = post @ W2 ctx.save_for_backward(x, W1, W2, pre, post) return x3.view(x.shape) @staticmethod def backward(ctx, grad_output): x, W1, W2, pre, post = ctx.saved_tensors dW2 = post.T @ grad_output dpre = linear_relu_square(grad_output.view((-1, grad_output.shape[-1])), W2, aux=pre) dW1 = dpre.T @ x dx = dpre @ W1 return dx.view(x.shape), dW1, dW2 # ----------------------------------------------------------------------------- # Fused Softcapped Cross Entropy @triton.jit def fused_softcapped_entropy_fwd_kernel( logits_ptr, losses_ptr, lse_ptr, targets_ptr, mtp_weights_ptr, stride_logits_n, stride_logits_v, n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE: tl.constexpr ): row_idx = tl.program_id(0).to(tl.int64) logits_row_ptr = logits_ptr + row_idx * stride_logits_n max_val = -float('inf') sum_exp = 0.0 for off in range(0, n_cols, BLOCK_SIZE): cols = off + tl.arange(0, BLOCK_SIZE) mask = cols < n_cols val = tl.load(logits_row_ptr + cols, mask=mask, other=-float('inf')).to(tl.float32) z = A * tl.sigmoid((val + B) / C) z = tl.where(mask, z, -float('inf')) curr_max = tl.max(z, axis=0) new_max = tl.maximum(max_val, curr_max) sum_exp = sum_exp * tl.exp(max_val - new_max) + tl.sum(tl.exp(z - new_max), axis=0) max_val = new_max lse = max_val + tl.log(sum_exp) tl.store(lse_ptr + row_idx, lse) total_loss = 0.0 for k in range(n_predict): target_idx = row_idx + k if target_idx < n_rows: weight = tl.load(mtp_weights_ptr + k) if weight > 0: target = tl.load(targets_ptr + target_idx).to(tl.int32) if target >= 0 and target < n_cols: val_target = tl.load(logits_row_ptr + target).to(tl.float32) z_target = A * tl.sigmoid((val_target + B) / C) total_loss += weight * (lse - z_target) tl.store(losses_ptr + row_idx, total_loss) @triton.jit def fused_softcapped_entropy_bwd_kernel( grad_input_ptr, grad_output_ptr, lse_ptr, logits_ptr, targets_ptr, mtp_weights_ptr, stride_logits_n, stride_logits_v, stride_grad_n, stride_grad_v, n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE: tl.constexpr ): row_idx = tl.program_id(0).to(tl.int64) logits_row_ptr = logits_ptr + row_idx * stride_logits_n grad_row_ptr = grad_input_ptr + row_idx * stride_grad_n lse = tl.load(lse_ptr + row_idx) grad_loss = tl.load(grad_output_ptr + row_idx) S_w = 0.0 for k in range(n_predict): if row_idx + k < n_rows: S_w += tl.load(mtp_weights_ptr + k) for off in range(0, n_cols, BLOCK_SIZE): cols = off + tl.arange(0, BLOCK_SIZE) mask = cols < n_cols val = tl.load(logits_row_ptr + cols, mask=mask, other=0.0).to(tl.float32) u = (val + B) / C sigmoid_u = tl.sigmoid(u) z = A * sigmoid_u p = tl.exp(z - lse) term1 = S_w * p term2 = tl.zeros([BLOCK_SIZE], dtype=tl.float32) for k in range(n_predict): if row_idx + k < n_rows: target = tl.load(targets_ptr + row_idx + k).to(tl.int32) weight = tl.load(mtp_weights_ptr + k) term2 += tl.where(cols == target, weight, 0.0) grad_z = grad_loss * (term1 - term2) dz_dx = (1.0 / C) * z * (1.0 - sigmoid_u) grad_x = grad_z * dz_dx tl.store(grad_row_ptr + cols, grad_x.to(tl.bfloat16), mask=mask) class FusedSoftcappedCrossEntropy(torch.autograd.Function): @staticmethod def forward(ctx, logits, targets, mtp_weights, A=23.0, B=5.0, C=7.5): n_rows, n_cols = logits.shape if mtp_weights is None: mtp_weights = torch.tensor([1.0], device=logits.device, dtype=torch.float32) n_predict = mtp_weights.shape[0] losses = torch.empty(n_rows, dtype=torch.float32, device=logits.device) lse = torch.empty(n_rows, dtype=torch.float32, device=logits.device) logits = logits.contiguous() targets = targets.contiguous() mtp_weights = mtp_weights.contiguous() grid = (n_rows,) fused_softcapped_entropy_fwd_kernel[grid]( logits, losses, lse, targets, mtp_weights, logits.stride(0), logits.stride(1), n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE=1024, num_warps=8, num_stages=4 ) ctx.save_for_backward(logits, targets, mtp_weights, lse) ctx.params = (A, B, C) return losses @staticmethod def backward(ctx, grad_output): logits, targets, mtp_weights, lse = ctx.saved_tensors A, B, C = ctx.params n_rows, n_cols = logits.shape n_predict = mtp_weights.shape[0] grad_input = torch.empty((n_rows, n_cols), dtype=torch.bfloat16, device=logits.device) grad_output = grad_output.contiguous() grid = (n_rows,) fused_softcapped_entropy_bwd_kernel[grid]( grad_input, grad_output, lse, logits, targets, mtp_weights, logits.stride(0), logits.stride(1), grad_input.stride(0), grad_input.stride(1), n_rows, n_cols, n_predict, A, B, C, BLOCK_SIZE=1024, num_warps=8, num_stages=4 ) return grad_input, None, None, None, None, None ==================================================================================================== Running Python 3.10.12 (main, May 27 2025, 17:12:29) [GCC 11.4.0] Running PyTorch 2.10.0.dev20251210+cu126 compiled for CUDA 12.6 Running Triton version 3.6.0 Mon Jan 26 02:56:39 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 32C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:62:00.0 Off | 0 | | N/A 32C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:63:00.0 Off | 0 | | N/A 33C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:64:00.0 Off | 0 | | N/A 31C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:6A:00.0 Off | 0 | | N/A 31C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 34C P0 123W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:6C:00.0 Off | 0 | | N/A 33C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:6D:00.0 Off | 0 | | N/A 32C P0 116W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 274175 C /usr/bin/python3 1510MiB | | 1 N/A N/A 274176 C /usr/bin/python3 1510MiB | | 2 N/A N/A 274177 C /usr/bin/python3 1510MiB | | 3 N/A N/A 274178 C /usr/bin/python3 1510MiB | | 4 N/A N/A 274179 C /usr/bin/python3 1510MiB | | 5 N/A N/A 274180 C /usr/bin/python3 1510MiB | | 6 N/A N/A 274181 C /usr/bin/python3 1510MiB | | 7 N/A N/A 274182 C /usr/bin/python3 1510MiB | +-----------------------------------------------------------------------------------------+ ==================================================================================================== Compiling model and warming up kernels (~7 minutes on first execution) Sampling steps [0, 1, 2, 519, 520, 521, 1039, 1040, 1041, 1559, 1560, 1561] for warmup Resetting Model step:0/1600 val_loss:10.8263 train_time:0ms step_avg:0.03ms step:1/1600 train_time:86ms step_avg:85.52ms step:2/1600 train_time:109ms step_avg:54.30ms step:3/1600 train_time:129ms step_avg:43.05ms step:4/1600 train_time:153ms step_avg:38.27ms step:5/1600 train_time:183ms step_avg:36.63ms step:6/1600 train_time:291ms step_avg:48.53ms step:7/1600 train_time:309ms step_avg:44.11ms step:8/1600 train_time:328ms step_avg:41.04ms step:9/1600 train_time:359ms step_avg:39.85ms step:10/1600 train_time:396ms step_avg:39.60ms step:11/1600 train_time:427ms step_avg:38.78ms step:12/1600 train_time:464ms step_avg:38.68ms step:13/1600 train_time:494ms step_avg:38.03ms step:14/1600 train_time:532ms step_avg:37.98ms step:15/1600 train_time:563ms step_avg:37.53ms step:16/1600 train_time:600ms step_avg:37.53ms step:17/1600 train_time:631ms step_avg:37.13ms step:18/1600 train_time:669ms step_avg:37.16ms step:19/1600 train_time:700ms step_avg:36.83ms step:20/1600 train_time:737ms step_avg:36.84ms step:21/1600 train_time:767ms step_avg:36.55ms step:22/1600 train_time:805ms step_avg:36.58ms step:23/1600 train_time:835ms step_avg:36.32ms step:24/1600 train_time:873ms step_avg:36.36ms step:25/1600 train_time:903ms step_avg:36.14ms step:26/1600 train_time:941ms step_avg:36.20ms step:27/1600 train_time:972ms step_avg:36.00ms step:28/1600 train_time:1010ms step_avg:36.06ms step:29/1600 train_time:1040ms step_avg:35.88ms step:30/1600 train_time:1078ms step_avg:35.92ms step:31/1600 train_time:1108ms step_avg:35.75ms step:32/1600 train_time:1146ms step_avg:35.80ms step:33/1600 train_time:1176ms step_avg:35.64ms step:34/1600 train_time:1213ms step_avg:35.68ms step:35/1600 train_time:1244ms step_avg:35.55ms step:36/1600 train_time:1283ms step_avg:35.63ms step:37/1600 train_time:1314ms step_avg:35.51ms step:38/1600 train_time:1351ms step_avg:35.56ms step:39/1600 train_time:1382ms step_avg:35.44ms step:40/1600 train_time:1419ms step_avg:35.48ms step:41/1600 train_time:1450ms step_avg:35.36ms step:42/1600 train_time:1487ms step_avg:35.40ms step:43/1600 train_time:1518ms step_avg:35.29ms step:44/1600 train_time:1555ms step_avg:35.34ms step:45/1600 train_time:1585ms step_avg:35.23ms step:46/1600 train_time:1623ms step_avg:35.28ms step:47/1600 train_time:1654ms step_avg:35.18ms step:48/1600 train_time:1691ms step_avg:35.24ms step:49/1600 train_time:1722ms step_avg:35.14ms step:50/1600 train_time:1760ms step_avg:35.19ms step:51/1600 train_time:1790ms step_avg:35.10ms step:52/1600 train_time:1827ms step_avg:35.14ms step:53/1600 train_time:1858ms step_avg:35.05ms step:54/1600 train_time:1895ms step_avg:35.09ms step:55/1600 train_time:1925ms step_avg:35.01ms step:56/1600 train_time:1963ms step_avg:35.06ms step:57/1600 train_time:1994ms step_avg:34.98ms step:58/1600 train_time:2031ms step_avg:35.01ms step:59/1600 train_time:2061ms step_avg:34.93ms step:60/1600 train_time:2098ms step_avg:34.97ms step:61/1600 train_time:2129ms step_avg:34.90ms step:62/1600 train_time:2167ms step_avg:34.94ms step:63/1600 train_time:2197ms step_avg:34.87ms step:64/1600 train_time:2234ms step_avg:34.90ms step:65/1600 train_time:2264ms step_avg:34.83ms step:66/1600 train_time:2302ms step_avg:34.88ms step:67/1600 train_time:2333ms step_avg:34.82ms step:68/1600 train_time:2371ms step_avg:34.86ms step:69/1600 train_time:2401ms step_avg:34.80ms step:70/1600 train_time:2439ms step_avg:34.84ms step:71/1600 train_time:2469ms step_avg:34.78ms step:72/1600 train_time:2507ms step_avg:34.82ms step:73/1600 train_time:2538ms step_avg:34.76ms step:74/1600 train_time:2575ms step_avg:34.79ms step:75/1600 train_time:2605ms step_avg:34.74ms step:76/1600 train_time:2643ms step_avg:34.77ms step:77/1600 train_time:2673ms step_avg:34.72ms step:78/1600 train_time:2711ms step_avg:34.75ms step:79/1600 train_time:2742ms step_avg:34.70ms step:80/1600 train_time:2780ms step_avg:34.75ms step:81/1600 train_time:2810ms step_avg:34.70ms step:82/1600 train_time:2848ms step_avg:34.73ms step:83/1600 train_time:2878ms step_avg:34.68ms step:84/1600 train_time:2915ms step_avg:34.70ms step:85/1600 train_time:2946ms step_avg:34.66ms step:86/1600 train_time:2984ms step_avg:34.69ms step:87/1600 train_time:3014ms step_avg:34.64ms step:88/1600 train_time:3051ms step_avg:34.67ms step:89/1600 train_time:3082ms step_avg:34.63ms step:90/1600 train_time:3120ms step_avg:34.67ms step:91/1600 train_time:3151ms step_avg:34.63ms step:92/1600 train_time:3189ms step_avg:34.66ms step:93/1600 train_time:3219ms step_avg:34.61ms step:94/1600 train_time:3256ms step_avg:34.64ms step:95/1600 train_time:3287ms step_avg:34.59ms step:96/1600 train_time:3324ms step_avg:34.63ms step:97/1600 train_time:3355ms step_avg:34.59ms step:98/1600 train_time:3392ms step_avg:34.62ms step:99/1600 train_time:3423ms step_avg:34.57ms step:100/1600 train_time:3460ms step_avg:34.60ms step:101/1600 train_time:3491ms step_avg:34.56ms step:102/1600 train_time:3529ms step_avg:34.59ms step:103/1600 train_time:3559ms step_avg:34.55ms step:104/1600 train_time:3596ms step_avg:34.58ms step:105/1600 train_time:3627ms step_avg:34.54ms step:106/1600 train_time:3664ms step_avg:34.57ms step:107/1600 train_time:3695ms step_avg:34.53ms step:108/1600 train_time:3732ms step_avg:34.55ms step:109/1600 train_time:3763ms step_avg:34.52ms step:110/1600 train_time:3801ms step_avg:34.55ms step:111/1600 train_time:3832ms step_avg:34.52ms step:112/1600 train_time:3869ms step_avg:34.54ms step:113/1600 train_time:3900ms step_avg:34.51ms step:114/1600 train_time:3937ms step_avg:34.54ms step:115/1600 train_time:3968ms step_avg:34.51ms step:116/1600 train_time:4006ms step_avg:34.53ms step:117/1600 train_time:4036ms step_avg:34.50ms step:118/1600 train_time:4074ms step_avg:34.53ms step:119/1600 train_time:4104ms step_avg:34.49ms step:120/1600 train_time:4141ms step_avg:34.51ms step:121/1600 train_time:4172ms step_avg:34.48ms step:122/1600 train_time:4209ms step_avg:34.50ms step:123/1600 train_time:4240ms step_avg:34.47ms step:124/1600 train_time:4277ms step_avg:34.49ms step:125/1600 train_time:4308ms step_avg:34.46ms step:126/1600 train_time:4345ms step_avg:34.49ms step:127/1600 train_time:4376ms step_avg:34.46ms step:128/1600 train_time:4413ms step_avg:34.48ms step:129/1600 train_time:4444ms step_avg:34.45ms step:130/1600 train_time:4482ms step_avg:34.48ms step:131/1600 train_time:4512ms step_avg:34.44ms step:132/1600 train_time:4550ms step_avg:34.47ms step:133/1600 train_time:4580ms step_avg:34.44ms step:134/1600 train_time:4617ms step_avg:34.45ms step:135/1600 train_time:4647ms step_avg:34.43ms step:136/1600 train_time:4685ms step_avg:34.45ms step:137/1600 train_time:4716ms step_avg:34.42ms step:138/1600 train_time:4753ms step_avg:34.44ms step:139/1600 train_time:4784ms step_avg:34.42ms step:140/1600 train_time:4822ms step_avg:34.44ms step:141/1600 train_time:4852ms step_avg:34.41ms step:142/1600 train_time:4890ms step_avg:34.43ms step:143/1600 train_time:4920ms step_avg:34.41ms step:144/1600 train_time:4957ms step_avg:34.42ms step:145/1600 train_time:4988ms step_avg:34.40ms step:146/1600 train_time:5025ms step_avg:34.42ms step:147/1600 train_time:5056ms step_avg:34.39ms step:148/1600 train_time:5093ms step_avg:34.41ms step:149/1600 train_time:5124ms step_avg:34.39ms step:150/1600 train_time:5161ms step_avg:34.41ms step:151/1600 train_time:5191ms step_avg:34.38ms step:152/1600 train_time:5229ms step_avg:34.40ms step:153/1600 train_time:5260ms step_avg:34.38ms step:154/1600 train_time:5297ms step_avg:34.40ms step:155/1600 train_time:5328ms step_avg:34.37ms step:156/1600 train_time:5365ms step_avg:34.39ms step:157/1600 train_time:5395ms step_avg:34.36ms step:158/1600 train_time:5432ms step_avg:34.38ms step:159/1600 train_time:5463ms step_avg:34.36ms step:160/1600 train_time:5501ms step_avg:34.38ms step:161/1600 train_time:5531ms step_avg:34.36ms step:162/1600 train_time:5568ms step_avg:34.37ms step:163/1600 train_time:5600ms step_avg:34.35ms step:164/1600 train_time:5637ms step_avg:34.37ms step:165/1600 train_time:5668ms step_avg:34.35ms step:166/1600 train_time:5706ms step_avg:34.37ms step:167/1600 train_time:5736ms step_avg:34.35ms step:168/1600 train_time:5773ms step_avg:34.36ms step:169/1600 train_time:5804ms step_avg:34.34ms step:170/1600 train_time:5841ms step_avg:34.36ms step:171/1600 train_time:5872ms step_avg:34.34ms step:172/1600 train_time:5909ms step_avg:34.36ms step:173/1600 train_time:5940ms step_avg:34.33ms step:174/1600 train_time:5977ms step_avg:34.35ms step:175/1600 train_time:6008ms step_avg:34.33ms step:176/1600 train_time:6046ms step_avg:34.35ms step:177/1600 train_time:6076ms step_avg:34.33ms step:178/1600 train_time:6113ms step_avg:34.34ms step:179/1600 train_time:6144ms step_avg:34.32ms step:180/1600 train_time:6181ms step_avg:34.34ms step:181/1600 train_time:6211ms step_avg:34.32ms step:182/1600 train_time:6249ms step_avg:34.33ms step:183/1600 train_time:6279ms step_avg:34.31ms step:184/1600 train_time:6316ms step_avg:34.33ms step:185/1600 train_time:6347ms step_avg:34.31ms step:186/1600 train_time:6384ms step_avg:34.32ms step:187/1600 train_time:6415ms step_avg:34.30ms step:188/1600 train_time:6452ms step_avg:34.32ms step:189/1600 train_time:6482ms step_avg:34.30ms step:190/1600 train_time:6520ms step_avg:34.31ms step:191/1600 train_time:6550ms step_avg:34.29ms step:192/1600 train_time:6588ms step_avg:34.31ms step:193/1600 train_time:6618ms step_avg:34.29ms step:194/1600 train_time:6655ms step_avg:34.31ms step:195/1600 train_time:6686ms step_avg:34.29ms step:196/1600 train_time:6723ms step_avg:34.30ms step:197/1600 train_time:6754ms step_avg:34.28ms step:198/1600 train_time:6791ms step_avg:34.30ms step:199/1600 train_time:6821ms step_avg:34.28ms step:200/1600 train_time:6859ms step_avg:34.29ms step:201/1600 train_time:6890ms step_avg:34.28ms step:202/1600 train_time:6927ms step_avg:34.29ms step:203/1600 train_time:6957ms step_avg:34.27ms step:204/1600 train_time:6995ms step_avg:34.29ms step:205/1600 train_time:7025ms step_avg:34.27ms step:206/1600 train_time:7062ms step_avg:34.28ms step:207/1600 train_time:7093ms step_avg:34.27ms step:208/1600 train_time:7130ms step_avg:34.28ms step:209/1600 train_time:7161ms step_avg:34.26ms step:210/1600 train_time:7198ms step_avg:34.28ms step:211/1600 train_time:7229ms step_avg:34.26ms step:212/1600 train_time:7266ms step_avg:34.27ms step:213/1600 train_time:7297ms step_avg:34.26ms step:214/1600 train_time:7334ms step_avg:34.27ms step:215/1600 train_time:7364ms step_avg:34.25ms step:216/1600 train_time:7402ms step_avg:34.27ms step:217/1600 train_time:7432ms step_avg:34.25ms step:218/1600 train_time:7470ms step_avg:34.26ms step:219/1600 train_time:7500ms step_avg:34.25ms step:220/1600 train_time:7538ms step_avg:34.26ms step:221/1600 train_time:7568ms step_avg:34.25ms step:222/1600 train_time:7606ms step_avg:34.26ms step:223/1600 train_time:7636ms step_avg:34.24ms step:224/1600 train_time:7673ms step_avg:34.25ms step:225/1600 train_time:7704ms step_avg:34.24ms step:226/1600 train_time:7741ms step_avg:34.25ms step:227/1600 train_time:7772ms step_avg:34.24ms step:228/1600 train_time:7809ms step_avg:34.25ms step:229/1600 train_time:7840ms step_avg:34.24ms step:230/1600 train_time:7877ms step_avg:34.25ms step:231/1600 train_time:7908ms step_avg:34.23ms step:232/1600 train_time:7946ms step_avg:34.25ms step:233/1600 train_time:7976ms step_avg:34.23ms step:234/1600 train_time:8014ms step_avg:34.25ms step:235/1600 train_time:8045ms step_avg:34.23ms step:236/1600 train_time:8082ms step_avg:34.25ms step:237/1600 train_time:8113ms step_avg:34.23ms step:238/1600 train_time:8150ms step_avg:34.24ms step:239/1600 train_time:8180ms step_avg:34.23ms step:240/1600 train_time:8217ms step_avg:34.24ms step:241/1600 train_time:8248ms step_avg:34.22ms step:242/1600 train_time:8285ms step_avg:34.24ms step:243/1600 train_time:8316ms step_avg:34.22ms step:244/1600 train_time:8353ms step_avg:34.23ms step:245/1600 train_time:8383ms step_avg:34.22ms step:246/1600 train_time:8421ms step_avg:34.23ms step:247/1600 train_time:8451ms step_avg:34.21ms step:248/1600 train_time:8489ms step_avg:34.23ms step:249/1600 train_time:8519ms step_avg:34.21ms step:250/1600 train_time:8556ms step_avg:34.22ms step:250/1600 val_loss:4.5872 train_time:8604ms step_avg:34.41ms step:251/1600 train_time:8624ms step_avg:34.36ms step:252/1600 train_time:8644ms step_avg:34.30ms step:253/1600 train_time:8661ms step_avg:34.24ms step:254/1600 train_time:8696ms step_avg:34.24ms step:255/1600 train_time:8730ms step_avg:34.23ms step:256/1600 train_time:8768ms step_avg:34.25ms step:257/1600 train_time:8800ms step_avg:34.24ms step:258/1600 train_time:8838ms step_avg:34.26ms step:259/1600 train_time:8868ms step_avg:34.24ms step:260/1600 train_time:8906ms step_avg:34.25ms step:261/1600 train_time:8937ms step_avg:34.24ms step:262/1600 train_time:8974ms step_avg:34.25ms step:263/1600 train_time:9004ms step_avg:34.24ms step:264/1600 train_time:9041ms step_avg:34.25ms step:265/1600 train_time:9072ms step_avg:34.23ms step:266/1600 train_time:9109ms step_avg:34.25ms step:267/1600 train_time:9140ms step_avg:34.23ms step:268/1600 train_time:9177ms step_avg:34.24ms step:269/1600 train_time:9207ms step_avg:34.23ms step:270/1600 train_time:9245ms step_avg:34.24ms step:271/1600 train_time:9275ms step_avg:34.23ms step:272/1600 train_time:9312ms step_avg:34.24ms step:273/1600 train_time:9342ms step_avg:34.22ms step:274/1600 train_time:9379ms step_avg:34.23ms step:275/1600 train_time:9410ms step_avg:34.22ms step:276/1600 train_time:9447ms step_avg:34.23ms step:277/1600 train_time:9477ms step_avg:34.21ms step:278/1600 train_time:9514ms step_avg:34.22ms step:279/1600 train_time:9545ms step_avg:34.21ms step:280/1600 train_time:9582ms step_avg:34.22ms step:281/1600 train_time:9612ms step_avg:34.21ms step:282/1600 train_time:9650ms step_avg:34.22ms step:283/1600 train_time:9680ms step_avg:34.21ms step:284/1600 train_time:9717ms step_avg:34.22ms step:285/1600 train_time:9748ms step_avg:34.20ms step:286/1600 train_time:9786ms step_avg:34.22ms step:287/1600 train_time:9817ms step_avg:34.21ms step:288/1600 train_time:9855ms step_avg:34.22ms step:289/1600 train_time:9885ms step_avg:34.21ms step:290/1600 train_time:9923ms step_avg:34.22ms step:291/1600 train_time:9954ms step_avg:34.21ms step:292/1600 train_time:9991ms step_avg:34.22ms step:293/1600 train_time:10021ms step_avg:34.20ms step:294/1600 train_time:10059ms step_avg:34.21ms step:295/1600 train_time:10089ms step_avg:34.20ms step:296/1600 train_time:10127ms step_avg:34.21ms step:297/1600 train_time:10157ms step_avg:34.20ms step:298/1600 train_time:10195ms step_avg:34.21ms step:299/1600 train_time:10225ms step_avg:34.20ms step:300/1600 train_time:10262ms step_avg:34.21ms step:301/1600 train_time:10293ms step_avg:34.20ms step:302/1600 train_time:10330ms step_avg:34.21ms step:303/1600 train_time:10361ms step_avg:34.19ms step:304/1600 train_time:10398ms step_avg:34.20ms step:305/1600 train_time:10428ms step_avg:34.19ms step:306/1600 train_time:10465ms step_avg:34.20ms step:307/1600 train_time:10496ms step_avg:34.19ms step:308/1600 train_time:10533ms step_avg:34.20ms step:309/1600 train_time:10563ms step_avg:34.18ms step:310/1600 train_time:10600ms step_avg:34.19ms step:311/1600 train_time:10631ms step_avg:34.18ms step:312/1600 train_time:10669ms step_avg:34.19ms step:313/1600 train_time:10699ms step_avg:34.18ms step:314/1600 train_time:10736ms step_avg:34.19ms step:315/1600 train_time:10767ms step_avg:34.18ms step:316/1600 train_time:10804ms step_avg:34.19ms step:317/1600 train_time:10835ms step_avg:34.18ms step:318/1600 train_time:10873ms step_avg:34.19ms step:319/1600 train_time:10903ms step_avg:34.18ms step:320/1600 train_time:10941ms step_avg:34.19ms step:321/1600 train_time:10971ms step_avg:34.18ms step:322/1600 train_time:11008ms step_avg:34.19ms step:323/1600 train_time:11039ms step_avg:34.18ms step:324/1600 train_time:11076ms step_avg:34.19ms step:325/1600 train_time:11107ms step_avg:34.17ms step:326/1600 train_time:11144ms step_avg:34.18ms step:327/1600 train_time:11174ms step_avg:34.17ms step:328/1600 train_time:11212ms step_avg:34.18ms step:329/1600 train_time:11242ms step_avg:34.17ms step:330/1600 train_time:11279ms step_avg:34.18ms step:331/1600 train_time:11310ms step_avg:34.17ms step:332/1600 train_time:11347ms step_avg:34.18ms step:333/1600 train_time:11378ms step_avg:34.17ms step:334/1600 train_time:11415ms step_avg:34.18ms step:335/1600 train_time:11445ms step_avg:34.17ms step:336/1600 train_time:11482ms step_avg:34.17ms step:337/1600 train_time:11513ms step_avg:34.16ms step:338/1600 train_time:11550ms step_avg:34.17ms step:339/1600 train_time:11580ms step_avg:34.16ms step:340/1600 train_time:11617ms step_avg:34.17ms step:341/1600 train_time:11648ms step_avg:34.16ms step:342/1600 train_time:11686ms step_avg:34.17ms step:343/1600 train_time:11716ms step_avg:34.16ms step:344/1600 train_time:11754ms step_avg:34.17ms step:345/1600 train_time:11784ms step_avg:34.16ms step:346/1600 train_time:11821ms step_avg:34.16ms step:347/1600 train_time:11852ms step_avg:34.15ms step:348/1600 train_time:11889ms step_avg:34.16ms step:349/1600 train_time:11920ms step_avg:34.15ms step:350/1600 train_time:11957ms step_avg:34.16ms step:351/1600 train_time:11988ms step_avg:34.15ms step:352/1600 train_time:12025ms step_avg:34.16ms step:353/1600 train_time:12056ms step_avg:34.15ms step:354/1600 train_time:12093ms step_avg:34.16ms step:355/1600 train_time:12123ms step_avg:34.15ms step:356/1600 train_time:12160ms step_avg:34.16ms step:357/1600 train_time:12191ms step_avg:34.15ms step:358/1600 train_time:12228ms step_avg:34.16ms step:359/1600 train_time:12259ms step_avg:34.15ms step:360/1600 train_time:12296ms step_avg:34.16ms step:361/1600 train_time:12327ms step_avg:34.15ms step:362/1600 train_time:12364ms step_avg:34.16ms step:363/1600 train_time:12395ms step_avg:34.15ms step:364/1600 train_time:12432ms step_avg:34.15ms step:365/1600 train_time:12463ms step_avg:34.14ms step:366/1600 train_time:12500ms step_avg:34.15ms step:367/1600 train_time:12531ms step_avg:34.14ms step:368/1600 train_time:12568ms step_avg:34.15ms step:369/1600 train_time:12598ms step_avg:34.14ms step:370/1600 train_time:12636ms step_avg:34.15ms step:371/1600 train_time:12666ms step_avg:34.14ms step:372/1600 train_time:12704ms step_avg:34.15ms step:373/1600 train_time:12734ms step_avg:34.14ms step:374/1600 train_time:12772ms step_avg:34.15ms step:375/1600 train_time:12803ms step_avg:34.14ms step:376/1600 train_time:12840ms step_avg:34.15ms step:377/1600 train_time:12870ms step_avg:34.14ms step:378/1600 train_time:12908ms step_avg:34.15ms step:379/1600 train_time:12938ms step_avg:34.14ms step:380/1600 train_time:12975ms step_avg:34.15ms step:381/1600 train_time:13005ms step_avg:34.13ms step:382/1600 train_time:13043ms step_avg:34.14ms step:383/1600 train_time:13073ms step_avg:34.13ms step:384/1600 train_time:13111ms step_avg:34.14ms step:385/1600 train_time:13141ms step_avg:34.13ms step:386/1600 train_time:13178ms step_avg:34.14ms step:387/1600 train_time:13209ms step_avg:34.13ms step:388/1600 train_time:13246ms step_avg:34.14ms step:389/1600 train_time:13277ms step_avg:34.13ms step:390/1600 train_time:13314ms step_avg:34.14ms step:391/1600 train_time:13345ms step_avg:34.13ms step:392/1600 train_time:13382ms step_avg:34.14ms step:393/1600 train_time:13413ms step_avg:34.13ms step:394/1600 train_time:13451ms step_avg:34.14ms step:395/1600 train_time:13481ms step_avg:34.13ms step:396/1600 train_time:13518ms step_avg:34.14ms step:397/1600 train_time:13549ms step_avg:34.13ms step:398/1600 train_time:13586ms step_avg:34.14ms step:399/1600 train_time:13617ms step_avg:34.13ms step:400/1600 train_time:13654ms step_avg:34.14ms step:401/1600 train_time:13685ms step_avg:34.13ms step:402/1600 train_time:13722ms step_avg:34.13ms step:403/1600 train_time:13753ms step_avg:34.13ms step:404/1600 train_time:13790ms step_avg:34.13ms step:405/1600 train_time:13821ms step_avg:34.12ms step:406/1600 train_time:13858ms step_avg:34.13ms step:407/1600 train_time:13888ms step_avg:34.12ms step:408/1600 train_time:13926ms step_avg:34.13ms step:409/1600 train_time:13956ms step_avg:34.12ms step:410/1600 train_time:13993ms step_avg:34.13ms step:411/1600 train_time:14024ms step_avg:34.12ms step:412/1600 train_time:14062ms step_avg:34.13ms step:413/1600 train_time:14092ms step_avg:34.12ms step:414/1600 train_time:14129ms step_avg:34.13ms step:415/1600 train_time:14160ms step_avg:34.12ms step:416/1600 train_time:14197ms step_avg:34.13ms step:417/1600 train_time:14228ms step_avg:34.12ms step:418/1600 train_time:14266ms step_avg:34.13ms step:419/1600 train_time:14297ms step_avg:34.12ms step:420/1600 train_time:14334ms step_avg:34.13ms step:421/1600 train_time:14365ms step_avg:34.12ms step:422/1600 train_time:14403ms step_avg:34.13ms step:423/1600 train_time:14433ms step_avg:34.12ms step:424/1600 train_time:14471ms step_avg:34.13ms step:425/1600 train_time:14501ms step_avg:34.12ms step:426/1600 train_time:14538ms step_avg:34.13ms step:427/1600 train_time:14569ms step_avg:34.12ms step:428/1600 train_time:14606ms step_avg:34.13ms step:429/1600 train_time:14637ms step_avg:34.12ms step:430/1600 train_time:14674ms step_avg:34.13ms step:431/1600 train_time:14705ms step_avg:34.12ms step:432/1600 train_time:14742ms step_avg:34.13ms step:433/1600 train_time:14773ms step_avg:34.12ms step:434/1600 train_time:14810ms step_avg:34.13ms step:435/1600 train_time:14841ms step_avg:34.12ms step:436/1600 train_time:14878ms step_avg:34.12ms step:437/1600 train_time:14909ms step_avg:34.12ms step:438/1600 train_time:14946ms step_avg:34.12ms step:439/1600 train_time:14977ms step_avg:34.12ms step:440/1600 train_time:15014ms step_avg:34.12ms step:441/1600 train_time:15045ms step_avg:34.12ms step:442/1600 train_time:15082ms step_avg:34.12ms step:443/1600 train_time:15112ms step_avg:34.11ms step:444/1600 train_time:15150ms step_avg:34.12ms step:445/1600 train_time:15181ms step_avg:34.11ms step:446/1600 train_time:15218ms step_avg:34.12ms step:447/1600 train_time:15248ms step_avg:34.11ms step:448/1600 train_time:15286ms step_avg:34.12ms step:449/1600 train_time:15316ms step_avg:34.11ms step:450/1600 train_time:15353ms step_avg:34.12ms step:451/1600 train_time:15384ms step_avg:34.11ms step:452/1600 train_time:15421ms step_avg:34.12ms step:453/1600 train_time:15452ms step_avg:34.11ms step:454/1600 train_time:15489ms step_avg:34.12ms step:455/1600 train_time:15520ms step_avg:34.11ms step:456/1600 train_time:15557ms step_avg:34.12ms step:457/1600 train_time:15588ms step_avg:34.11ms step:458/1600 train_time:15625ms step_avg:34.12ms step:459/1600 train_time:15656ms step_avg:34.11ms step:460/1600 train_time:15693ms step_avg:34.11ms step:461/1600 train_time:15723ms step_avg:34.11ms step:462/1600 train_time:15761ms step_avg:34.12ms step:463/1600 train_time:15791ms step_avg:34.11ms step:464/1600 train_time:15829ms step_avg:34.11ms step:465/1600 train_time:15859ms step_avg:34.11ms step:466/1600 train_time:15896ms step_avg:34.11ms step:467/1600 train_time:15927ms step_avg:34.10ms step:468/1600 train_time:15964ms step_avg:34.11ms step:469/1600 train_time:15994ms step_avg:34.10ms step:470/1600 train_time:16031ms step_avg:34.11ms step:471/1600 train_time:16062ms step_avg:34.10ms step:472/1600 train_time:16099ms step_avg:34.11ms step:473/1600 train_time:16130ms step_avg:34.10ms step:474/1600 train_time:16168ms step_avg:34.11ms step:475/1600 train_time:16198ms step_avg:34.10ms step:476/1600 train_time:16235ms step_avg:34.11ms step:477/1600 train_time:16266ms step_avg:34.10ms step:478/1600 train_time:16303ms step_avg:34.11ms step:479/1600 train_time:16334ms step_avg:34.10ms step:480/1600 train_time:16371ms step_avg:34.11ms step:481/1600 train_time:16402ms step_avg:34.10ms step:482/1600 train_time:16439ms step_avg:34.11ms step:483/1600 train_time:16469ms step_avg:34.10ms step:484/1600 train_time:16507ms step_avg:34.10ms step:485/1600 train_time:16537ms step_avg:34.10ms step:486/1600 train_time:16574ms step_avg:34.10ms step:487/1600 train_time:16604ms step_avg:34.10ms step:488/1600 train_time:16642ms step_avg:34.10ms step:489/1600 train_time:16672ms step_avg:34.09ms step:490/1600 train_time:16709ms step_avg:34.10ms step:491/1600 train_time:16740ms step_avg:34.09ms step:492/1600 train_time:16777ms step_avg:34.10ms step:493/1600 train_time:16808ms step_avg:34.09ms step:494/1600 train_time:16845ms step_avg:34.10ms step:495/1600 train_time:16876ms step_avg:34.09ms step:496/1600 train_time:16914ms step_avg:34.10ms step:497/1600 train_time:16944ms step_avg:34.09ms step:498/1600 train_time:16981ms step_avg:34.10ms step:499/1600 train_time:17012ms step_avg:34.09ms step:500/1600 train_time:17049ms step_avg:34.10ms step:500/1600 val_loss:4.2438 train_time:17097ms step_avg:34.19ms step:501/1600 train_time:17119ms step_avg:34.17ms step:502/1600 train_time:17139ms step_avg:34.14ms step:503/1600 train_time:17157ms step_avg:34.11ms step:504/1600 train_time:17189ms step_avg:34.11ms step:505/1600 train_time:17221ms step_avg:34.10ms step:506/1600 train_time:17259ms step_avg:34.11ms step:507/1600 train_time:17290ms step_avg:34.10ms step:508/1600 train_time:17328ms step_avg:34.11ms step:509/1600 train_time:17358ms step_avg:34.10ms step:510/1600 train_time:17396ms step_avg:34.11ms step:511/1600 train_time:17427ms step_avg:34.10ms step:512/1600 train_time:17464ms step_avg:34.11ms step:513/1600 train_time:17494ms step_avg:34.10ms step:514/1600 train_time:17532ms step_avg:34.11ms step:515/1600 train_time:17562ms step_avg:34.10ms step:516/1600 train_time:17599ms step_avg:34.11ms step:517/1600 train_time:17630ms step_avg:34.10ms step:518/1600 train_time:17667ms step_avg:34.11ms step:519/1600 train_time:17697ms step_avg:34.10ms step:520/1600 train_time:17734ms step_avg:34.10ms step:521/1600 train_time:17806ms step_avg:34.18ms step:522/1600 train_time:17863ms step_avg:34.22ms step:523/1600 train_time:17924ms step_avg:34.27ms step:524/1600 train_time:17982ms step_avg:34.32ms step:525/1600 train_time:18043ms step_avg:34.37ms step:526/1600 train_time:18102ms step_avg:34.41ms step:527/1600 train_time:18165ms step_avg:34.47ms step:528/1600 train_time:18223ms step_avg:34.51ms step:529/1600 train_time:18287ms step_avg:34.57ms step:530/1600 train_time:18347ms step_avg:34.62ms step:531/1600 train_time:18409ms step_avg:34.67ms step:532/1600 train_time:18468ms step_avg:34.71ms step:533/1600 train_time:18531ms step_avg:34.77ms step:534/1600 train_time:18590ms step_avg:34.81ms step:535/1600 train_time:18654ms step_avg:34.87ms step:536/1600 train_time:18713ms step_avg:34.91ms step:537/1600 train_time:18775ms step_avg:34.96ms step:538/1600 train_time:18834ms step_avg:35.01ms step:539/1600 train_time:18896ms step_avg:35.06ms step:540/1600 train_time:18957ms step_avg:35.10ms step:541/1600 train_time:19018ms step_avg:35.15ms step:542/1600 train_time:19076ms step_avg:35.20ms step:543/1600 train_time:19139ms step_avg:35.25ms step:544/1600 train_time:19199ms step_avg:35.29ms step:545/1600 train_time:19262ms step_avg:35.34ms step:546/1600 train_time:19322ms step_avg:35.39ms step:547/1600 train_time:19384ms step_avg:35.44ms step:548/1600 train_time:19444ms step_avg:35.48ms step:549/1600 train_time:19506ms step_avg:35.53ms step:550/1600 train_time:19565ms step_avg:35.57ms step:551/1600 train_time:19627ms step_avg:35.62ms step:552/1600 train_time:19686ms step_avg:35.66ms step:553/1600 train_time:19748ms step_avg:35.71ms step:554/1600 train_time:19807ms step_avg:35.75ms step:555/1600 train_time:19870ms step_avg:35.80ms step:556/1600 train_time:19930ms step_avg:35.85ms step:557/1600 train_time:19992ms step_avg:35.89ms step:558/1600 train_time:20052ms step_avg:35.94ms step:559/1600 train_time:20115ms step_avg:35.98ms step:560/1600 train_time:20175ms step_avg:36.03ms step:561/1600 train_time:20239ms step_avg:36.08ms step:562/1600 train_time:20298ms step_avg:36.12ms step:563/1600 train_time:20360ms step_avg:36.16ms step:564/1600 train_time:20419ms step_avg:36.20ms step:565/1600 train_time:20482ms step_avg:36.25ms step:566/1600 train_time:20542ms step_avg:36.29ms step:567/1600 train_time:20604ms step_avg:36.34ms step:568/1600 train_time:20666ms step_avg:36.38ms step:569/1600 train_time:20729ms step_avg:36.43ms step:570/1600 train_time:20785ms step_avg:36.46ms step:571/1600 train_time:20847ms step_avg:36.51ms step:572/1600 train_time:20906ms step_avg:36.55ms step:573/1600 train_time:20973ms step_avg:36.60ms step:574/1600 train_time:21029ms step_avg:36.64ms step:575/1600 train_time:21091ms step_avg:36.68ms step:576/1600 train_time:21150ms step_avg:36.72ms step:577/1600 train_time:21211ms step_avg:36.76ms step:578/1600 train_time:21272ms step_avg:36.80ms step:579/1600 train_time:21334ms step_avg:36.85ms step:580/1600 train_time:21394ms step_avg:36.89ms step:581/1600 train_time:21458ms step_avg:36.93ms step:582/1600 train_time:21516ms step_avg:36.97ms step:583/1600 train_time:21579ms step_avg:37.01ms step:584/1600 train_time:21638ms step_avg:37.05ms step:585/1600 train_time:21701ms step_avg:37.10ms step:586/1600 train_time:21760ms step_avg:37.13ms step:587/1600 train_time:21822ms step_avg:37.18ms step:588/1600 train_time:21883ms step_avg:37.22ms step:589/1600 train_time:21944ms step_avg:37.26ms step:590/1600 train_time:22003ms step_avg:37.29ms step:591/1600 train_time:22066ms step_avg:37.34ms step:592/1600 train_time:22125ms step_avg:37.37ms step:593/1600 train_time:22187ms step_avg:37.41ms step:594/1600 train_time:22246ms step_avg:37.45ms step:595/1600 train_time:22307ms step_avg:37.49ms step:596/1600 train_time:22367ms step_avg:37.53ms step:597/1600 train_time:22429ms step_avg:37.57ms step:598/1600 train_time:22487ms step_avg:37.60ms step:599/1600 train_time:22550ms step_avg:37.65ms step:600/1600 train_time:22609ms step_avg:37.68ms step:601/1600 train_time:22673ms step_avg:37.73ms step:602/1600 train_time:22733ms step_avg:37.76ms step:603/1600 train_time:22796ms step_avg:37.80ms step:604/1600 train_time:22855ms step_avg:37.84ms step:605/1600 train_time:22917ms step_avg:37.88ms step:606/1600 train_time:22977ms step_avg:37.92ms step:607/1600 train_time:23039ms step_avg:37.96ms step:608/1600 train_time:23098ms step_avg:37.99ms step:609/1600 train_time:23161ms step_avg:38.03ms step:610/1600 train_time:23221ms step_avg:38.07ms step:611/1600 train_time:23283ms step_avg:38.11ms step:612/1600 train_time:23342ms step_avg:38.14ms step:613/1600 train_time:23404ms step_avg:38.18ms step:614/1600 train_time:23464ms step_avg:38.21ms step:615/1600 train_time:23525ms step_avg:38.25ms step:616/1600 train_time:23583ms step_avg:38.28ms step:617/1600 train_time:23646ms step_avg:38.32ms step:618/1600 train_time:23704ms step_avg:38.36ms step:619/1600 train_time:23767ms step_avg:38.40ms step:620/1600 train_time:23825ms step_avg:38.43ms step:621/1600 train_time:23887ms step_avg:38.47ms step:622/1600 train_time:23947ms step_avg:38.50ms step:623/1600 train_time:24009ms step_avg:38.54ms step:624/1600 train_time:24069ms step_avg:38.57ms step:625/1600 train_time:24131ms step_avg:38.61ms step:626/1600 train_time:24191ms step_avg:38.64ms step:627/1600 train_time:24254ms step_avg:38.68ms step:628/1600 train_time:24313ms step_avg:38.72ms step:629/1600 train_time:24376ms step_avg:38.75ms step:630/1600 train_time:24435ms step_avg:38.79ms step:631/1600 train_time:24497ms step_avg:38.82ms step:632/1600 train_time:24557ms step_avg:38.86ms step:633/1600 train_time:24620ms step_avg:38.89ms step:634/1600 train_time:24680ms step_avg:38.93ms step:635/1600 train_time:24742ms step_avg:38.96ms step:636/1600 train_time:24801ms step_avg:38.99ms step:637/1600 train_time:24863ms step_avg:39.03ms step:638/1600 train_time:24922ms step_avg:39.06ms step:639/1600 train_time:24987ms step_avg:39.10ms step:640/1600 train_time:25044ms step_avg:39.13ms step:641/1600 train_time:25107ms step_avg:39.17ms step:642/1600 train_time:25165ms step_avg:39.20ms step:643/1600 train_time:25228ms step_avg:39.23ms step:644/1600 train_time:25287ms step_avg:39.26ms step:645/1600 train_time:25348ms step_avg:39.30ms step:646/1600 train_time:25407ms step_avg:39.33ms step:647/1600 train_time:25470ms step_avg:39.37ms step:648/1600 train_time:25529ms step_avg:39.40ms step:649/1600 train_time:25593ms step_avg:39.43ms step:650/1600 train_time:25653ms step_avg:39.47ms step:651/1600 train_time:25714ms step_avg:39.50ms step:652/1600 train_time:25777ms step_avg:39.53ms step:653/1600 train_time:25838ms step_avg:39.57ms step:654/1600 train_time:25897ms step_avg:39.60ms step:655/1600 train_time:25961ms step_avg:39.63ms step:656/1600 train_time:26019ms step_avg:39.66ms step:657/1600 train_time:26081ms step_avg:39.70ms step:658/1600 train_time:26141ms step_avg:39.73ms step:659/1600 train_time:26204ms step_avg:39.76ms step:660/1600 train_time:26263ms step_avg:39.79ms step:661/1600 train_time:26325ms step_avg:39.83ms step:662/1600 train_time:26384ms step_avg:39.86ms step:663/1600 train_time:26446ms step_avg:39.89ms step:664/1600 train_time:26505ms step_avg:39.92ms step:665/1600 train_time:26568ms step_avg:39.95ms step:666/1600 train_time:26626ms step_avg:39.98ms step:667/1600 train_time:26688ms step_avg:40.01ms step:668/1600 train_time:26747ms step_avg:40.04ms step:669/1600 train_time:26809ms step_avg:40.07ms step:670/1600 train_time:26868ms step_avg:40.10ms step:671/1600 train_time:26932ms step_avg:40.14ms step:672/1600 train_time:26992ms step_avg:40.17ms step:673/1600 train_time:27054ms step_avg:40.20ms step:674/1600 train_time:27114ms step_avg:40.23ms step:675/1600 train_time:27177ms step_avg:40.26ms step:676/1600 train_time:27236ms step_avg:40.29ms step:677/1600 train_time:27298ms step_avg:40.32ms step:678/1600 train_time:27357ms step_avg:40.35ms step:679/1600 train_time:27420ms step_avg:40.38ms step:680/1600 train_time:27480ms step_avg:40.41ms step:681/1600 train_time:27543ms step_avg:40.44ms step:682/1600 train_time:27602ms step_avg:40.47ms step:683/1600 train_time:27664ms step_avg:40.50ms step:684/1600 train_time:27723ms step_avg:40.53ms step:685/1600 train_time:27786ms step_avg:40.56ms step:686/1600 train_time:27844ms step_avg:40.59ms step:687/1600 train_time:27906ms step_avg:40.62ms step:688/1600 train_time:27966ms step_avg:40.65ms step:689/1600 train_time:28028ms step_avg:40.68ms step:690/1600 train_time:28087ms step_avg:40.71ms step:691/1600 train_time:28150ms step_avg:40.74ms step:692/1600 train_time:28209ms step_avg:40.76ms step:693/1600 train_time:28272ms step_avg:40.80ms step:694/1600 train_time:28331ms step_avg:40.82ms step:695/1600 train_time:28395ms step_avg:40.86ms step:696/1600 train_time:28454ms step_avg:40.88ms step:697/1600 train_time:28518ms step_avg:40.91ms step:698/1600 train_time:28577ms step_avg:40.94ms step:699/1600 train_time:28639ms step_avg:40.97ms step:700/1600 train_time:28698ms step_avg:41.00ms step:701/1600 train_time:28760ms step_avg:41.03ms step:702/1600 train_time:28819ms step_avg:41.05ms step:703/1600 train_time:28882ms step_avg:41.08ms step:704/1600 train_time:28941ms step_avg:41.11ms step:705/1600 train_time:29003ms step_avg:41.14ms step:706/1600 train_time:29063ms step_avg:41.17ms step:707/1600 train_time:29125ms step_avg:41.20ms step:708/1600 train_time:29184ms step_avg:41.22ms step:709/1600 train_time:29247ms step_avg:41.25ms step:710/1600 train_time:29305ms step_avg:41.27ms step:711/1600 train_time:29368ms step_avg:41.30ms step:712/1600 train_time:29427ms step_avg:41.33ms step:713/1600 train_time:29490ms step_avg:41.36ms step:714/1600 train_time:29549ms step_avg:41.39ms step:715/1600 train_time:29611ms step_avg:41.41ms step:716/1600 train_time:29672ms step_avg:41.44ms step:717/1600 train_time:29735ms step_avg:41.47ms step:718/1600 train_time:29796ms step_avg:41.50ms step:719/1600 train_time:29857ms step_avg:41.53ms step:720/1600 train_time:29916ms step_avg:41.55ms step:721/1600 train_time:29978ms step_avg:41.58ms step:722/1600 train_time:30039ms step_avg:41.60ms step:723/1600 train_time:30101ms step_avg:41.63ms step:724/1600 train_time:30159ms step_avg:41.66ms step:725/1600 train_time:30222ms step_avg:41.68ms step:726/1600 train_time:30282ms step_avg:41.71ms step:727/1600 train_time:30345ms step_avg:41.74ms step:728/1600 train_time:30403ms step_avg:41.76ms step:729/1600 train_time:30466ms step_avg:41.79ms step:730/1600 train_time:30524ms step_avg:41.81ms step:731/1600 train_time:30587ms step_avg:41.84ms step:732/1600 train_time:30646ms step_avg:41.87ms step:733/1600 train_time:30708ms step_avg:41.89ms step:734/1600 train_time:30767ms step_avg:41.92ms step:735/1600 train_time:30828ms step_avg:41.94ms step:736/1600 train_time:30889ms step_avg:41.97ms step:737/1600 train_time:30951ms step_avg:42.00ms step:738/1600 train_time:31009ms step_avg:42.02ms step:739/1600 train_time:31072ms step_avg:42.05ms step:740/1600 train_time:31132ms step_avg:42.07ms step:741/1600 train_time:31195ms step_avg:42.10ms step:742/1600 train_time:31257ms step_avg:42.13ms step:743/1600 train_time:31318ms step_avg:42.15ms step:744/1600 train_time:31376ms step_avg:42.17ms step:745/1600 train_time:31439ms step_avg:42.20ms step:746/1600 train_time:31500ms step_avg:42.23ms step:747/1600 train_time:31562ms step_avg:42.25ms step:748/1600 train_time:31621ms step_avg:42.27ms step:749/1600 train_time:31684ms step_avg:42.30ms step:750/1600 train_time:31744ms step_avg:42.33ms step:750/1600 val_loss:3.8925 train_time:31789ms step_avg:42.39ms step:751/1600 train_time:31810ms step_avg:42.36ms step:752/1600 train_time:31870ms step_avg:42.38ms step:753/1600 train_time:31936ms step_avg:42.41ms step:754/1600 train_time:31999ms step_avg:42.44ms step:755/1600 train_time:32061ms step_avg:42.46ms step:756/1600 train_time:32120ms step_avg:42.49ms step:757/1600 train_time:32181ms step_avg:42.51ms step:758/1600 train_time:32240ms step_avg:42.53ms step:759/1600 train_time:32301ms step_avg:42.56ms step:760/1600 train_time:32359ms step_avg:42.58ms step:761/1600 train_time:32422ms step_avg:42.60ms step:762/1600 train_time:32480ms step_avg:42.62ms step:763/1600 train_time:32541ms step_avg:42.65ms step:764/1600 train_time:32602ms step_avg:42.67ms step:765/1600 train_time:32661ms step_avg:42.69ms step:766/1600 train_time:32720ms step_avg:42.72ms step:767/1600 train_time:32784ms step_avg:42.74ms step:768/1600 train_time:32845ms step_avg:42.77ms step:769/1600 train_time:32909ms step_avg:42.79ms step:770/1600 train_time:32969ms step_avg:42.82ms step:771/1600 train_time:33032ms step_avg:42.84ms step:772/1600 train_time:33092ms step_avg:42.87ms step:773/1600 train_time:33154ms step_avg:42.89ms step:774/1600 train_time:33212ms step_avg:42.91ms step:775/1600 train_time:33275ms step_avg:42.94ms step:776/1600 train_time:33334ms step_avg:42.96ms step:777/1600 train_time:33396ms step_avg:42.98ms step:778/1600 train_time:33455ms step_avg:43.00ms step:779/1600 train_time:33518ms step_avg:43.03ms step:780/1600 train_time:33576ms step_avg:43.05ms step:781/1600 train_time:33638ms step_avg:43.07ms step:782/1600 train_time:33697ms step_avg:43.09ms step:783/1600 train_time:33759ms step_avg:43.12ms step:784/1600 train_time:33819ms step_avg:43.14ms step:785/1600 train_time:33882ms step_avg:43.16ms step:786/1600 train_time:33941ms step_avg:43.18ms step:787/1600 train_time:34003ms step_avg:43.21ms step:788/1600 train_time:34063ms step_avg:43.23ms step:789/1600 train_time:34126ms step_avg:43.25ms step:790/1600 train_time:34185ms step_avg:43.27ms step:791/1600 train_time:34248ms step_avg:43.30ms step:792/1600 train_time:34307ms step_avg:43.32ms step:793/1600 train_time:34370ms step_avg:43.34ms step:794/1600 train_time:34429ms step_avg:43.36ms step:795/1600 train_time:34491ms step_avg:43.39ms step:796/1600 train_time:34550ms step_avg:43.40ms step:797/1600 train_time:34613ms step_avg:43.43ms step:798/1600 train_time:34673ms step_avg:43.45ms step:799/1600 train_time:34735ms step_avg:43.47ms step:800/1600 train_time:34795ms step_avg:43.49ms step:801/1600 train_time:34858ms step_avg:43.52ms step:802/1600 train_time:34916ms step_avg:43.54ms step:803/1600 train_time:34979ms step_avg:43.56ms step:804/1600 train_time:35037ms step_avg:43.58ms step:805/1600 train_time:35101ms step_avg:43.60ms step:806/1600 train_time:35161ms step_avg:43.62ms step:807/1600 train_time:35222ms step_avg:43.65ms step:808/1600 train_time:35279ms step_avg:43.66ms step:809/1600 train_time:35343ms step_avg:43.69ms step:810/1600 train_time:35403ms step_avg:43.71ms step:811/1600 train_time:35465ms step_avg:43.73ms step:812/1600 train_time:35524ms step_avg:43.75ms step:813/1600 train_time:35586ms step_avg:43.77ms step:814/1600 train_time:35645ms step_avg:43.79ms step:815/1600 train_time:35708ms step_avg:43.81ms step:816/1600 train_time:35767ms step_avg:43.83ms step:817/1600 train_time:35830ms step_avg:43.86ms step:818/1600 train_time:35890ms step_avg:43.88ms step:819/1600 train_time:35953ms step_avg:43.90ms step:820/1600 train_time:36014ms step_avg:43.92ms step:821/1600 train_time:36076ms step_avg:43.94ms step:822/1600 train_time:36134ms step_avg:43.96ms step:823/1600 train_time:36197ms step_avg:43.98ms step:824/1600 train_time:36256ms step_avg:44.00ms step:825/1600 train_time:36318ms step_avg:44.02ms step:826/1600 train_time:36378ms step_avg:44.04ms step:827/1600 train_time:36439ms step_avg:44.06ms step:828/1600 train_time:36498ms step_avg:44.08ms step:829/1600 train_time:36561ms step_avg:44.10ms step:830/1600 train_time:36620ms step_avg:44.12ms step:831/1600 train_time:36683ms step_avg:44.14ms step:832/1600 train_time:36742ms step_avg:44.16ms step:833/1600 train_time:36807ms step_avg:44.19ms step:834/1600 train_time:36865ms step_avg:44.20ms step:835/1600 train_time:36928ms step_avg:44.23ms step:836/1600 train_time:36988ms step_avg:44.24ms step:837/1600 train_time:37050ms step_avg:44.26ms step:838/1600 train_time:37110ms step_avg:44.28ms step:839/1600 train_time:37172ms step_avg:44.31ms step:840/1600 train_time:37231ms step_avg:44.32ms step:841/1600 train_time:37294ms step_avg:44.34ms step:842/1600 train_time:37352ms step_avg:44.36ms step:843/1600 train_time:37415ms step_avg:44.38ms step:844/1600 train_time:37474ms step_avg:44.40ms step:845/1600 train_time:37537ms step_avg:44.42ms step:846/1600 train_time:37596ms step_avg:44.44ms step:847/1600 train_time:37659ms step_avg:44.46ms step:848/1600 train_time:37718ms step_avg:44.48ms step:849/1600 train_time:37780ms step_avg:44.50ms step:850/1600 train_time:37839ms step_avg:44.52ms step:851/1600 train_time:37900ms step_avg:44.54ms step:852/1600 train_time:37959ms step_avg:44.55ms step:853/1600 train_time:38022ms step_avg:44.57ms step:854/1600 train_time:38081ms step_avg:44.59ms step:855/1600 train_time:38145ms step_avg:44.61ms step:856/1600 train_time:38204ms step_avg:44.63ms step:857/1600 train_time:38267ms step_avg:44.65ms step:858/1600 train_time:38327ms step_avg:44.67ms step:859/1600 train_time:38389ms step_avg:44.69ms step:860/1600 train_time:38449ms step_avg:44.71ms step:861/1600 train_time:38511ms step_avg:44.73ms step:862/1600 train_time:38570ms step_avg:44.75ms step:863/1600 train_time:38634ms step_avg:44.77ms step:864/1600 train_time:38694ms step_avg:44.78ms step:865/1600 train_time:38757ms step_avg:44.81ms step:866/1600 train_time:38815ms step_avg:44.82ms step:867/1600 train_time:38878ms step_avg:44.84ms step:868/1600 train_time:38937ms step_avg:44.86ms step:869/1600 train_time:38999ms step_avg:44.88ms step:870/1600 train_time:39058ms step_avg:44.89ms step:871/1600 train_time:39121ms step_avg:44.92ms step:872/1600 train_time:39179ms step_avg:44.93ms step:873/1600 train_time:39242ms step_avg:44.95ms step:874/1600 train_time:39301ms step_avg:44.97ms step:875/1600 train_time:39363ms step_avg:44.99ms step:876/1600 train_time:39423ms step_avg:45.00ms step:877/1600 train_time:39484ms step_avg:45.02ms step:878/1600 train_time:39545ms step_avg:45.04ms step:879/1600 train_time:39608ms step_avg:45.06ms step:880/1600 train_time:39667ms step_avg:45.08ms step:881/1600 train_time:39735ms step_avg:45.10ms step:882/1600 train_time:39789ms step_avg:45.11ms step:883/1600 train_time:39851ms step_avg:45.13ms step:884/1600 train_time:39911ms step_avg:45.15ms step:885/1600 train_time:39974ms step_avg:45.17ms step:886/1600 train_time:40033ms step_avg:45.18ms step:887/1600 train_time:40095ms step_avg:45.20ms step:888/1600 train_time:40154ms step_avg:45.22ms step:889/1600 train_time:40218ms step_avg:45.24ms step:890/1600 train_time:40284ms step_avg:45.26ms step:891/1600 train_time:40343ms step_avg:45.28ms step:892/1600 train_time:40401ms step_avg:45.29ms step:893/1600 train_time:40461ms step_avg:45.31ms step:894/1600 train_time:40519ms step_avg:45.32ms step:895/1600 train_time:40582ms step_avg:45.34ms step:896/1600 train_time:40640ms step_avg:45.36ms step:897/1600 train_time:40703ms step_avg:45.38ms step:898/1600 train_time:40762ms step_avg:45.39ms step:899/1600 train_time:40824ms step_avg:45.41ms step:900/1600 train_time:40884ms step_avg:45.43ms step:901/1600 train_time:40947ms step_avg:45.45ms step:902/1600 train_time:41006ms step_avg:45.46ms step:903/1600 train_time:41069ms step_avg:45.48ms step:904/1600 train_time:41129ms step_avg:45.50ms step:905/1600 train_time:41191ms step_avg:45.52ms step:906/1600 train_time:41251ms step_avg:45.53ms step:907/1600 train_time:41314ms step_avg:45.55ms step:908/1600 train_time:41373ms step_avg:45.57ms step:909/1600 train_time:41436ms step_avg:45.58ms step:910/1600 train_time:41494ms step_avg:45.60ms step:911/1600 train_time:41557ms step_avg:45.62ms step:912/1600 train_time:41617ms step_avg:45.63ms step:913/1600 train_time:41679ms step_avg:45.65ms step:914/1600 train_time:41738ms step_avg:45.66ms step:915/1600 train_time:41801ms step_avg:45.68ms step:916/1600 train_time:41860ms step_avg:45.70ms step:917/1600 train_time:41922ms step_avg:45.72ms step:918/1600 train_time:41980ms step_avg:45.73ms step:919/1600 train_time:42043ms step_avg:45.75ms step:920/1600 train_time:42105ms step_avg:45.77ms step:921/1600 train_time:42166ms step_avg:45.78ms step:922/1600 train_time:42226ms step_avg:45.80ms step:923/1600 train_time:42289ms step_avg:45.82ms step:924/1600 train_time:42349ms step_avg:45.83ms step:925/1600 train_time:42412ms step_avg:45.85ms step:926/1600 train_time:42471ms step_avg:45.86ms step:927/1600 train_time:42534ms step_avg:45.88ms step:928/1600 train_time:42593ms step_avg:45.90ms step:929/1600 train_time:42655ms step_avg:45.92ms step:930/1600 train_time:42714ms step_avg:45.93ms step:931/1600 train_time:42777ms step_avg:45.95ms step:932/1600 train_time:42836ms step_avg:45.96ms step:933/1600 train_time:42898ms step_avg:45.98ms step:934/1600 train_time:42957ms step_avg:45.99ms step:935/1600 train_time:43019ms step_avg:46.01ms step:936/1600 train_time:43078ms step_avg:46.02ms step:937/1600 train_time:43141ms step_avg:46.04ms step:938/1600 train_time:43200ms step_avg:46.05ms step:939/1600 train_time:43262ms step_avg:46.07ms step:940/1600 train_time:43321ms step_avg:46.09ms step:941/1600 train_time:43384ms step_avg:46.10ms step:942/1600 train_time:43444ms step_avg:46.12ms step:943/1600 train_time:43507ms step_avg:46.14ms step:944/1600 train_time:43566ms step_avg:46.15ms step:945/1600 train_time:43630ms step_avg:46.17ms step:946/1600 train_time:43689ms step_avg:46.18ms step:947/1600 train_time:43752ms step_avg:46.20ms step:948/1600 train_time:43810ms step_avg:46.21ms step:949/1600 train_time:43873ms step_avg:46.23ms step:950/1600 train_time:43933ms step_avg:46.25ms step:951/1600 train_time:43997ms step_avg:46.26ms step:952/1600 train_time:44055ms step_avg:46.28ms step:953/1600 train_time:44118ms step_avg:46.29ms step:954/1600 train_time:44176ms step_avg:46.31ms step:955/1600 train_time:44239ms step_avg:46.32ms step:956/1600 train_time:44298ms step_avg:46.34ms step:957/1600 train_time:44362ms step_avg:46.36ms step:958/1600 train_time:44419ms step_avg:46.37ms step:959/1600 train_time:44484ms step_avg:46.39ms step:960/1600 train_time:44541ms step_avg:46.40ms step:961/1600 train_time:44605ms step_avg:46.42ms step:962/1600 train_time:44664ms step_avg:46.43ms step:963/1600 train_time:44727ms step_avg:46.45ms step:964/1600 train_time:44787ms step_avg:46.46ms step:965/1600 train_time:44849ms step_avg:46.48ms step:966/1600 train_time:44908ms step_avg:46.49ms step:967/1600 train_time:44972ms step_avg:46.51ms step:968/1600 train_time:45031ms step_avg:46.52ms step:969/1600 train_time:45094ms step_avg:46.54ms step:970/1600 train_time:45153ms step_avg:46.55ms step:971/1600 train_time:45216ms step_avg:46.57ms step:972/1600 train_time:45275ms step_avg:46.58ms step:973/1600 train_time:45337ms step_avg:46.60ms step:974/1600 train_time:45398ms step_avg:46.61ms step:975/1600 train_time:45461ms step_avg:46.63ms step:976/1600 train_time:45518ms step_avg:46.64ms step:977/1600 train_time:45580ms step_avg:46.65ms step:978/1600 train_time:45639ms step_avg:46.67ms step:979/1600 train_time:45700ms step_avg:46.68ms step:980/1600 train_time:45759ms step_avg:46.69ms step:981/1600 train_time:45822ms step_avg:46.71ms step:982/1600 train_time:45882ms step_avg:46.72ms step:983/1600 train_time:45944ms step_avg:46.74ms step:984/1600 train_time:46004ms step_avg:46.75ms step:985/1600 train_time:46067ms step_avg:46.77ms step:986/1600 train_time:46126ms step_avg:46.78ms step:987/1600 train_time:46189ms step_avg:46.80ms step:988/1600 train_time:46249ms step_avg:46.81ms step:989/1600 train_time:46311ms step_avg:46.83ms step:990/1600 train_time:46371ms step_avg:46.84ms step:991/1600 train_time:46433ms step_avg:46.85ms step:992/1600 train_time:46492ms step_avg:46.87ms step:993/1600 train_time:46556ms step_avg:46.88ms step:994/1600 train_time:46616ms step_avg:46.90ms step:995/1600 train_time:46678ms step_avg:46.91ms step:996/1600 train_time:46736ms step_avg:46.92ms step:997/1600 train_time:46798ms step_avg:46.94ms step:998/1600 train_time:46857ms step_avg:46.95ms step:999/1600 train_time:46919ms step_avg:46.97ms step:1000/1600 train_time:46978ms step_avg:46.98ms step:1000/1600 val_loss:3.5970 train_time:47024ms step_avg:47.02ms step:1001/1600 train_time:47045ms step_avg:47.00ms step:1002/1600 train_time:47103ms step_avg:47.01ms step:1003/1600 train_time:47167ms step_avg:47.03ms step:1004/1600 train_time:47226ms step_avg:47.04ms step:1005/1600 train_time:47289ms step_avg:47.05ms step:1006/1600 train_time:47348ms step_avg:47.07ms step:1007/1600 train_time:47410ms step_avg:47.08ms step:1008/1600 train_time:47469ms step_avg:47.09ms step:1009/1600 train_time:47530ms step_avg:47.11ms step:1010/1600 train_time:47588ms step_avg:47.12ms step:1011/1600 train_time:47653ms step_avg:47.13ms step:1012/1600 train_time:47712ms step_avg:47.15ms step:1013/1600 train_time:47773ms step_avg:47.16ms step:1014/1600 train_time:47833ms step_avg:47.17ms step:1015/1600 train_time:47895ms step_avg:47.19ms step:1016/1600 train_time:47954ms step_avg:47.20ms step:1017/1600 train_time:48018ms step_avg:47.22ms step:1018/1600 train_time:48079ms step_avg:47.23ms step:1019/1600 train_time:48143ms step_avg:47.25ms step:1020/1600 train_time:48203ms step_avg:47.26ms step:1021/1600 train_time:48265ms step_avg:47.27ms step:1022/1600 train_time:48323ms step_avg:47.28ms step:1023/1600 train_time:48386ms step_avg:47.30ms step:1024/1600 train_time:48444ms step_avg:47.31ms step:1025/1600 train_time:48507ms step_avg:47.32ms step:1026/1600 train_time:48565ms step_avg:47.33ms step:1027/1600 train_time:48627ms step_avg:47.35ms step:1028/1600 train_time:48685ms step_avg:47.36ms step:1029/1600 train_time:48748ms step_avg:47.37ms step:1030/1600 train_time:48806ms step_avg:47.38ms step:1031/1600 train_time:48868ms step_avg:47.40ms step:1032/1600 train_time:48927ms step_avg:47.41ms step:1033/1600 train_time:48991ms step_avg:47.43ms step:1034/1600 train_time:49050ms step_avg:47.44ms step:1035/1600 train_time:49114ms step_avg:47.45ms step:1036/1600 train_time:49175ms step_avg:47.47ms step:1037/1600 train_time:49239ms step_avg:47.48ms step:1038/1600 train_time:49297ms step_avg:47.49ms step:1039/1600 train_time:49359ms step_avg:47.51ms step:1040/1600 train_time:49418ms step_avg:47.52ms step:1041/1600 train_time:49493ms step_avg:47.54ms step:1042/1600 train_time:49573ms step_avg:47.57ms step:1043/1600 train_time:49664ms step_avg:47.62ms step:1044/1600 train_time:49748ms step_avg:47.65ms step:1045/1600 train_time:49835ms step_avg:47.69ms step:1046/1600 train_time:49920ms step_avg:47.72ms step:1047/1600 train_time:50009ms step_avg:47.76ms step:1048/1600 train_time:50094ms step_avg:47.80ms step:1049/1600 train_time:50183ms step_avg:47.84ms step:1050/1600 train_time:50269ms step_avg:47.88ms step:1051/1600 train_time:50358ms step_avg:47.91ms step:1052/1600 train_time:50443ms step_avg:47.95ms step:1053/1600 train_time:50532ms step_avg:47.99ms step:1054/1600 train_time:50616ms step_avg:48.02ms step:1055/1600 train_time:50705ms step_avg:48.06ms step:1056/1600 train_time:50790ms step_avg:48.10ms step:1057/1600 train_time:50878ms step_avg:48.13ms step:1058/1600 train_time:50963ms step_avg:48.17ms step:1059/1600 train_time:51051ms step_avg:48.21ms step:1060/1600 train_time:51137ms step_avg:48.24ms step:1061/1600 train_time:51225ms step_avg:48.28ms step:1062/1600 train_time:51311ms step_avg:48.32ms step:1063/1600 train_time:51398ms step_avg:48.35ms step:1064/1600 train_time:51484ms step_avg:48.39ms step:1065/1600 train_time:51572ms step_avg:48.42ms step:1066/1600 train_time:51656ms step_avg:48.46ms step:1067/1600 train_time:51745ms step_avg:48.50ms step:1068/1600 train_time:51830ms step_avg:48.53ms step:1069/1600 train_time:51918ms step_avg:48.57ms step:1070/1600 train_time:52004ms step_avg:48.60ms step:1071/1600 train_time:52092ms step_avg:48.64ms step:1072/1600 train_time:52178ms step_avg:48.67ms step:1073/1600 train_time:52267ms step_avg:48.71ms step:1074/1600 train_time:52352ms step_avg:48.74ms step:1075/1600 train_time:52440ms step_avg:48.78ms step:1076/1600 train_time:52526ms step_avg:48.82ms step:1077/1600 train_time:52615ms step_avg:48.85ms step:1078/1600 train_time:52701ms step_avg:48.89ms step:1079/1600 train_time:52790ms step_avg:48.92ms step:1080/1600 train_time:52874ms step_avg:48.96ms step:1081/1600 train_time:52962ms step_avg:48.99ms step:1082/1600 train_time:53049ms step_avg:49.03ms step:1083/1600 train_time:53137ms step_avg:49.06ms step:1084/1600 train_time:53224ms step_avg:49.10ms step:1085/1600 train_time:53311ms step_avg:49.13ms step:1086/1600 train_time:53396ms step_avg:49.17ms step:1087/1600 train_time:53485ms step_avg:49.20ms step:1088/1600 train_time:53570ms step_avg:49.24ms step:1089/1600 train_time:53657ms step_avg:49.27ms step:1090/1600 train_time:53743ms step_avg:49.31ms step:1091/1600 train_time:53831ms step_avg:49.34ms step:1092/1600 train_time:53916ms step_avg:49.37ms step:1093/1600 train_time:54004ms step_avg:49.41ms step:1094/1600 train_time:54090ms step_avg:49.44ms step:1095/1600 train_time:54179ms step_avg:49.48ms step:1096/1600 train_time:54264ms step_avg:49.51ms step:1097/1600 train_time:54352ms step_avg:49.55ms step:1098/1600 train_time:54438ms step_avg:49.58ms step:1099/1600 train_time:54527ms step_avg:49.61ms step:1100/1600 train_time:54611ms step_avg:49.65ms step:1101/1600 train_time:54699ms step_avg:49.68ms step:1102/1600 train_time:54787ms step_avg:49.72ms step:1103/1600 train_time:54873ms step_avg:49.75ms step:1104/1600 train_time:54959ms step_avg:49.78ms step:1105/1600 train_time:55046ms step_avg:49.82ms step:1106/1600 train_time:55133ms step_avg:49.85ms step:1107/1600 train_time:55221ms step_avg:49.88ms step:1108/1600 train_time:55306ms step_avg:49.92ms step:1109/1600 train_time:55395ms step_avg:49.95ms step:1110/1600 train_time:55480ms step_avg:49.98ms step:1111/1600 train_time:55568ms step_avg:50.02ms step:1112/1600 train_time:55654ms step_avg:50.05ms step:1113/1600 train_time:55741ms step_avg:50.08ms step:1114/1600 train_time:55826ms step_avg:50.11ms step:1115/1600 train_time:55914ms step_avg:50.15ms step:1116/1600 train_time:55999ms step_avg:50.18ms step:1117/1600 train_time:56089ms step_avg:50.21ms step:1118/1600 train_time:56174ms step_avg:50.25ms step:1119/1600 train_time:56262ms step_avg:50.28ms step:1120/1600 train_time:56347ms step_avg:50.31ms step:1121/1600 train_time:56436ms step_avg:50.34ms step:1122/1600 train_time:56522ms step_avg:50.38ms step:1123/1600 train_time:56611ms step_avg:50.41ms step:1124/1600 train_time:56697ms step_avg:50.44ms step:1125/1600 train_time:56784ms step_avg:50.47ms step:1126/1600 train_time:56871ms step_avg:50.51ms step:1127/1600 train_time:56958ms step_avg:50.54ms step:1128/1600 train_time:57042ms step_avg:50.57ms step:1129/1600 train_time:57130ms step_avg:50.60ms step:1130/1600 train_time:57215ms step_avg:50.63ms step:1131/1600 train_time:57304ms step_avg:50.67ms step:1132/1600 train_time:57389ms step_avg:50.70ms step:1133/1600 train_time:57477ms step_avg:50.73ms step:1134/1600 train_time:57574ms step_avg:50.77ms step:1135/1600 train_time:57664ms step_avg:50.81ms step:1136/1600 train_time:57745ms step_avg:50.83ms step:1137/1600 train_time:57832ms step_avg:50.86ms step:1138/1600 train_time:57917ms step_avg:50.89ms step:1139/1600 train_time:58005ms step_avg:50.93ms step:1140/1600 train_time:58088ms step_avg:50.95ms step:1141/1600 train_time:58175ms step_avg:50.99ms step:1142/1600 train_time:58260ms step_avg:51.02ms step:1143/1600 train_time:58347ms step_avg:51.05ms step:1144/1600 train_time:58432ms step_avg:51.08ms step:1145/1600 train_time:58521ms step_avg:51.11ms step:1146/1600 train_time:58606ms step_avg:51.14ms step:1147/1600 train_time:58694ms step_avg:51.17ms step:1148/1600 train_time:58779ms step_avg:51.20ms step:1149/1600 train_time:58867ms step_avg:51.23ms step:1150/1600 train_time:58953ms step_avg:51.26ms step:1151/1600 train_time:59041ms step_avg:51.30ms step:1152/1600 train_time:59126ms step_avg:51.32ms step:1153/1600 train_time:59215ms step_avg:51.36ms step:1154/1600 train_time:59305ms step_avg:51.39ms step:1155/1600 train_time:59391ms step_avg:51.42ms step:1156/1600 train_time:59476ms step_avg:51.45ms step:1157/1600 train_time:59563ms step_avg:51.48ms step:1158/1600 train_time:59648ms step_avg:51.51ms step:1159/1600 train_time:59736ms step_avg:51.54ms step:1160/1600 train_time:59821ms step_avg:51.57ms step:1161/1600 train_time:59911ms step_avg:51.60ms step:1162/1600 train_time:59996ms step_avg:51.63ms step:1163/1600 train_time:60083ms step_avg:51.66ms step:1164/1600 train_time:60167ms step_avg:51.69ms step:1165/1600 train_time:60257ms step_avg:51.72ms step:1166/1600 train_time:60341ms step_avg:51.75ms step:1167/1600 train_time:60429ms step_avg:51.78ms step:1168/1600 train_time:60514ms step_avg:51.81ms step:1169/1600 train_time:60603ms step_avg:51.84ms step:1170/1600 train_time:60688ms step_avg:51.87ms step:1171/1600 train_time:60776ms step_avg:51.90ms step:1172/1600 train_time:60861ms step_avg:51.93ms step:1173/1600 train_time:60950ms step_avg:51.96ms step:1174/1600 train_time:61035ms step_avg:51.99ms step:1175/1600 train_time:61122ms step_avg:52.02ms step:1176/1600 train_time:61207ms step_avg:52.05ms step:1177/1600 train_time:61296ms step_avg:52.08ms step:1178/1600 train_time:61381ms step_avg:52.11ms step:1179/1600 train_time:61469ms step_avg:52.14ms step:1180/1600 train_time:61554ms step_avg:52.16ms step:1181/1600 train_time:61642ms step_avg:52.20ms step:1182/1600 train_time:61728ms step_avg:52.22ms step:1183/1600 train_time:61817ms step_avg:52.25ms step:1184/1600 train_time:61902ms step_avg:52.28ms step:1185/1600 train_time:61990ms step_avg:52.31ms step:1186/1600 train_time:62075ms step_avg:52.34ms step:1187/1600 train_time:62163ms step_avg:52.37ms step:1188/1600 train_time:62249ms step_avg:52.40ms step:1189/1600 train_time:62336ms step_avg:52.43ms step:1190/1600 train_time:62422ms step_avg:52.46ms step:1191/1600 train_time:62511ms step_avg:52.49ms step:1192/1600 train_time:62596ms step_avg:52.51ms step:1193/1600 train_time:62685ms step_avg:52.54ms step:1194/1600 train_time:62771ms step_avg:52.57ms step:1195/1600 train_time:62858ms step_avg:52.60ms step:1196/1600 train_time:62944ms step_avg:52.63ms step:1197/1600 train_time:63032ms step_avg:52.66ms step:1198/1600 train_time:63117ms step_avg:52.68ms step:1199/1600 train_time:63205ms step_avg:52.71ms step:1200/1600 train_time:63291ms step_avg:52.74ms step:1201/1600 train_time:63378ms step_avg:52.77ms step:1202/1600 train_time:63463ms step_avg:52.80ms step:1203/1600 train_time:63551ms step_avg:52.83ms step:1204/1600 train_time:63637ms step_avg:52.85ms step:1205/1600 train_time:63725ms step_avg:52.88ms step:1206/1600 train_time:63810ms step_avg:52.91ms step:1207/1600 train_time:63898ms step_avg:52.94ms step:1208/1600 train_time:63984ms step_avg:52.97ms step:1209/1600 train_time:64073ms step_avg:53.00ms step:1210/1600 train_time:64158ms step_avg:53.02ms step:1211/1600 train_time:64247ms step_avg:53.05ms step:1212/1600 train_time:64332ms step_avg:53.08ms step:1213/1600 train_time:64419ms step_avg:53.11ms step:1214/1600 train_time:64505ms step_avg:53.13ms step:1215/1600 train_time:64593ms step_avg:53.16ms step:1216/1600 train_time:64679ms step_avg:53.19ms step:1217/1600 train_time:64773ms step_avg:53.22ms step:1218/1600 train_time:64857ms step_avg:53.25ms step:1219/1600 train_time:64945ms step_avg:53.28ms step:1220/1600 train_time:65031ms step_avg:53.30ms step:1221/1600 train_time:65118ms step_avg:53.33ms step:1222/1600 train_time:65203ms step_avg:53.36ms step:1223/1600 train_time:65292ms step_avg:53.39ms step:1224/1600 train_time:65376ms step_avg:53.41ms step:1225/1600 train_time:65466ms step_avg:53.44ms step:1226/1600 train_time:65551ms step_avg:53.47ms step:1227/1600 train_time:65639ms step_avg:53.50ms step:1228/1600 train_time:65725ms step_avg:53.52ms step:1229/1600 train_time:65813ms step_avg:53.55ms step:1230/1600 train_time:65898ms step_avg:53.58ms step:1231/1600 train_time:65987ms step_avg:53.60ms step:1232/1600 train_time:66072ms step_avg:53.63ms step:1233/1600 train_time:66160ms step_avg:53.66ms step:1234/1600 train_time:66246ms step_avg:53.68ms step:1235/1600 train_time:66334ms step_avg:53.71ms step:1236/1600 train_time:66419ms step_avg:53.74ms step:1237/1600 train_time:66508ms step_avg:53.77ms step:1238/1600 train_time:66594ms step_avg:53.79ms step:1239/1600 train_time:66681ms step_avg:53.82ms step:1240/1600 train_time:66768ms step_avg:53.85ms step:1241/1600 train_time:66857ms step_avg:53.87ms step:1242/1600 train_time:66942ms step_avg:53.90ms step:1243/1600 train_time:67031ms step_avg:53.93ms step:1244/1600 train_time:67115ms step_avg:53.95ms step:1245/1600 train_time:67204ms step_avg:53.98ms step:1246/1600 train_time:67289ms step_avg:54.00ms step:1247/1600 train_time:67377ms step_avg:54.03ms step:1248/1600 train_time:67463ms step_avg:54.06ms step:1249/1600 train_time:67551ms step_avg:54.08ms step:1250/1600 train_time:67637ms step_avg:54.11ms step:1250/1600 val_loss:3.4166 train_time:67710ms step_avg:54.17ms step:1251/1600 train_time:67731ms step_avg:54.14ms step:1252/1600 train_time:67817ms step_avg:54.17ms step:1253/1600 train_time:67909ms step_avg:54.20ms step:1254/1600 train_time:67995ms step_avg:54.22ms step:1255/1600 train_time:68082ms step_avg:54.25ms step:1256/1600 train_time:68167ms step_avg:54.27ms step:1257/1600 train_time:68254ms step_avg:54.30ms step:1258/1600 train_time:68339ms step_avg:54.32ms step:1259/1600 train_time:68426ms step_avg:54.35ms step:1260/1600 train_time:68511ms step_avg:54.37ms step:1261/1600 train_time:68598ms step_avg:54.40ms step:1262/1600 train_time:68684ms step_avg:54.42ms step:1263/1600 train_time:68774ms step_avg:54.45ms step:1264/1600 train_time:68862ms step_avg:54.48ms step:1265/1600 train_time:68952ms step_avg:54.51ms step:1266/1600 train_time:69039ms step_avg:54.53ms step:1267/1600 train_time:69127ms step_avg:54.56ms step:1268/1600 train_time:69211ms step_avg:54.58ms step:1269/1600 train_time:69300ms step_avg:54.61ms step:1270/1600 train_time:69383ms step_avg:54.63ms step:1271/1600 train_time:69470ms step_avg:54.66ms step:1272/1600 train_time:69554ms step_avg:54.68ms step:1273/1600 train_time:69642ms step_avg:54.71ms step:1274/1600 train_time:69727ms step_avg:54.73ms step:1275/1600 train_time:69817ms step_avg:54.76ms step:1276/1600 train_time:69904ms step_avg:54.78ms step:1277/1600 train_time:69994ms step_avg:54.81ms step:1278/1600 train_time:70079ms step_avg:54.83ms step:1279/1600 train_time:70167ms step_avg:54.86ms step:1280/1600 train_time:70252ms step_avg:54.88ms step:1281/1600 train_time:70339ms step_avg:54.91ms step:1282/1600 train_time:70424ms step_avg:54.93ms step:1283/1600 train_time:70513ms step_avg:54.96ms step:1284/1600 train_time:70597ms step_avg:54.98ms step:1285/1600 train_time:70685ms step_avg:55.01ms step:1286/1600 train_time:70772ms step_avg:55.03ms step:1287/1600 train_time:70860ms step_avg:55.06ms step:1288/1600 train_time:70946ms step_avg:55.08ms step:1289/1600 train_time:71035ms step_avg:55.11ms step:1290/1600 train_time:71120ms step_avg:55.13ms step:1291/1600 train_time:71208ms step_avg:55.16ms step:1292/1600 train_time:71292ms step_avg:55.18ms step:1293/1600 train_time:71381ms step_avg:55.21ms step:1294/1600 train_time:71464ms step_avg:55.23ms step:1295/1600 train_time:71552ms step_avg:55.25ms step:1296/1600 train_time:71637ms step_avg:55.28ms step:1297/1600 train_time:71726ms step_avg:55.30ms step:1298/1600 train_time:71812ms step_avg:55.32ms step:1299/1600 train_time:71900ms step_avg:55.35ms step:1300/1600 train_time:71986ms step_avg:55.37ms step:1301/1600 train_time:72075ms step_avg:55.40ms step:1302/1600 train_time:72161ms step_avg:55.42ms step:1303/1600 train_time:72250ms step_avg:55.45ms step:1304/1600 train_time:72334ms step_avg:55.47ms step:1305/1600 train_time:72422ms step_avg:55.50ms step:1306/1600 train_time:72506ms step_avg:55.52ms step:1307/1600 train_time:72594ms step_avg:55.54ms step:1308/1600 train_time:72680ms step_avg:55.57ms step:1309/1600 train_time:72770ms step_avg:55.59ms step:1310/1600 train_time:72855ms step_avg:55.61ms step:1311/1600 train_time:72945ms step_avg:55.64ms step:1312/1600 train_time:73030ms step_avg:55.66ms step:1313/1600 train_time:73118ms step_avg:55.69ms step:1314/1600 train_time:73204ms step_avg:55.71ms step:1315/1600 train_time:73291ms step_avg:55.73ms step:1316/1600 train_time:73377ms step_avg:55.76ms step:1317/1600 train_time:73465ms step_avg:55.78ms step:1318/1600 train_time:73550ms step_avg:55.80ms step:1319/1600 train_time:73638ms step_avg:55.83ms step:1320/1600 train_time:73723ms step_avg:55.85ms step:1321/1600 train_time:73811ms step_avg:55.88ms step:1322/1600 train_time:73899ms step_avg:55.90ms step:1323/1600 train_time:73987ms step_avg:55.92ms step:1324/1600 train_time:74072ms step_avg:55.95ms step:1325/1600 train_time:74160ms step_avg:55.97ms step:1326/1600 train_time:74244ms step_avg:55.99ms step:1327/1600 train_time:74332ms step_avg:56.02ms step:1328/1600 train_time:74419ms step_avg:56.04ms step:1329/1600 train_time:74506ms step_avg:56.06ms step:1330/1600 train_time:74591ms step_avg:56.08ms step:1331/1600 train_time:74679ms step_avg:56.11ms step:1332/1600 train_time:74764ms step_avg:56.13ms step:1333/1600 train_time:74854ms step_avg:56.15ms step:1334/1600 train_time:74939ms step_avg:56.18ms step:1335/1600 train_time:75027ms step_avg:56.20ms step:1336/1600 train_time:75112ms step_avg:56.22ms step:1337/1600 train_time:75201ms step_avg:56.25ms step:1338/1600 train_time:75286ms step_avg:56.27ms step:1339/1600 train_time:75373ms step_avg:56.29ms step:1340/1600 train_time:75458ms step_avg:56.31ms step:1341/1600 train_time:75546ms step_avg:56.34ms step:1342/1600 train_time:75631ms step_avg:56.36ms step:1343/1600 train_time:75719ms step_avg:56.38ms step:1344/1600 train_time:75805ms step_avg:56.40ms step:1345/1600 train_time:75894ms step_avg:56.43ms step:1346/1600 train_time:75977ms step_avg:56.45ms step:1347/1600 train_time:76066ms step_avg:56.47ms step:1348/1600 train_time:76152ms step_avg:56.49ms step:1349/1600 train_time:76240ms step_avg:56.52ms step:1350/1600 train_time:76325ms step_avg:56.54ms step:1351/1600 train_time:76414ms step_avg:56.56ms step:1352/1600 train_time:76499ms step_avg:56.58ms step:1353/1600 train_time:76587ms step_avg:56.61ms step:1354/1600 train_time:76672ms step_avg:56.63ms step:1355/1600 train_time:76760ms step_avg:56.65ms step:1356/1600 train_time:76846ms step_avg:56.67ms step:1357/1600 train_time:76935ms step_avg:56.69ms step:1358/1600 train_time:77020ms step_avg:56.72ms step:1359/1600 train_time:77108ms step_avg:56.74ms step:1360/1600 train_time:77193ms step_avg:56.76ms step:1361/1600 train_time:77280ms step_avg:56.78ms step:1362/1600 train_time:77366ms step_avg:56.80ms step:1363/1600 train_time:77454ms step_avg:56.83ms step:1364/1600 train_time:77539ms step_avg:56.85ms step:1365/1600 train_time:77628ms step_avg:56.87ms step:1366/1600 train_time:77717ms step_avg:56.89ms step:1367/1600 train_time:77804ms step_avg:56.92ms step:1368/1600 train_time:77889ms step_avg:56.94ms step:1369/1600 train_time:77975ms step_avg:56.96ms step:1370/1600 train_time:78061ms step_avg:56.98ms step:1371/1600 train_time:78149ms step_avg:57.00ms step:1372/1600 train_time:78233ms step_avg:57.02ms step:1373/1600 train_time:78322ms step_avg:57.04ms step:1374/1600 train_time:78407ms step_avg:57.07ms step:1375/1600 train_time:78495ms step_avg:57.09ms step:1376/1600 train_time:78580ms step_avg:57.11ms step:1377/1600 train_time:78669ms step_avg:57.13ms step:1378/1600 train_time:78754ms step_avg:57.15ms step:1379/1600 train_time:78844ms step_avg:57.18ms step:1380/1600 train_time:78929ms step_avg:57.20ms step:1381/1600 train_time:79016ms step_avg:57.22ms step:1382/1600 train_time:79102ms step_avg:57.24ms step:1383/1600 train_time:79191ms step_avg:57.26ms step:1384/1600 train_time:79276ms step_avg:57.28ms step:1385/1600 train_time:79365ms step_avg:57.30ms step:1386/1600 train_time:79450ms step_avg:57.32ms step:1387/1600 train_time:79541ms step_avg:57.35ms step:1388/1600 train_time:79624ms step_avg:57.37ms step:1389/1600 train_time:79712ms step_avg:57.39ms step:1390/1600 train_time:79797ms step_avg:57.41ms step:1391/1600 train_time:79885ms step_avg:57.43ms step:1392/1600 train_time:79970ms step_avg:57.45ms step:1393/1600 train_time:80058ms step_avg:57.47ms step:1394/1600 train_time:80145ms step_avg:57.49ms step:1395/1600 train_time:80233ms step_avg:57.51ms step:1396/1600 train_time:80319ms step_avg:57.53ms step:1397/1600 train_time:80409ms step_avg:57.56ms step:1398/1600 train_time:80493ms step_avg:57.58ms step:1399/1600 train_time:80582ms step_avg:57.60ms step:1400/1600 train_time:80667ms step_avg:57.62ms step:1401/1600 train_time:80755ms step_avg:57.64ms step:1402/1600 train_time:80841ms step_avg:57.66ms step:1403/1600 train_time:80929ms step_avg:57.68ms step:1404/1600 train_time:81014ms step_avg:57.70ms step:1405/1600 train_time:81102ms step_avg:57.72ms step:1406/1600 train_time:81188ms step_avg:57.74ms step:1407/1600 train_time:81276ms step_avg:57.77ms step:1408/1600 train_time:81362ms step_avg:57.79ms step:1409/1600 train_time:81451ms step_avg:57.81ms step:1410/1600 train_time:81535ms step_avg:57.83ms step:1411/1600 train_time:81624ms step_avg:57.85ms step:1412/1600 train_time:81709ms step_avg:57.87ms step:1413/1600 train_time:81797ms step_avg:57.89ms step:1414/1600 train_time:81882ms step_avg:57.91ms step:1415/1600 train_time:81970ms step_avg:57.93ms step:1416/1600 train_time:82055ms step_avg:57.95ms step:1417/1600 train_time:82143ms step_avg:57.97ms step:1418/1600 train_time:82227ms step_avg:57.99ms step:1419/1600 train_time:82317ms step_avg:58.01ms step:1420/1600 train_time:82403ms step_avg:58.03ms step:1421/1600 train_time:82492ms step_avg:58.05ms step:1422/1600 train_time:82576ms step_avg:58.07ms step:1423/1600 train_time:82665ms step_avg:58.09ms step:1424/1600 train_time:82750ms step_avg:58.11ms step:1425/1600 train_time:82838ms step_avg:58.13ms step:1426/1600 train_time:82924ms step_avg:58.15ms step:1427/1600 train_time:83012ms step_avg:58.17ms step:1428/1600 train_time:83097ms step_avg:58.19ms step:1429/1600 train_time:83185ms step_avg:58.21ms step:1430/1600 train_time:83271ms step_avg:58.23ms step:1431/1600 train_time:83359ms step_avg:58.25ms step:1432/1600 train_time:83445ms step_avg:58.27ms step:1433/1600 train_time:83532ms step_avg:58.29ms step:1434/1600 train_time:83617ms step_avg:58.31ms step:1435/1600 train_time:83706ms step_avg:58.33ms step:1436/1600 train_time:83792ms step_avg:58.35ms step:1437/1600 train_time:83880ms step_avg:58.37ms step:1438/1600 train_time:83966ms step_avg:58.39ms step:1439/1600 train_time:84054ms step_avg:58.41ms step:1440/1600 train_time:84139ms step_avg:58.43ms step:1441/1600 train_time:84227ms step_avg:58.45ms step:1442/1600 train_time:84312ms step_avg:58.47ms step:1443/1600 train_time:84401ms step_avg:58.49ms step:1444/1600 train_time:84485ms step_avg:58.51ms step:1445/1600 train_time:84574ms step_avg:58.53ms step:1446/1600 train_time:84661ms step_avg:58.55ms step:1447/1600 train_time:84750ms step_avg:58.57ms step:1448/1600 train_time:84835ms step_avg:58.59ms step:1449/1600 train_time:84922ms step_avg:58.61ms step:1450/1600 train_time:85006ms step_avg:58.63ms step:1451/1600 train_time:85095ms step_avg:58.65ms step:1452/1600 train_time:85181ms step_avg:58.66ms step:1453/1600 train_time:85269ms step_avg:58.68ms step:1454/1600 train_time:85354ms step_avg:58.70ms step:1455/1600 train_time:85442ms step_avg:58.72ms step:1456/1600 train_time:85527ms step_avg:58.74ms step:1457/1600 train_time:85616ms step_avg:58.76ms step:1458/1600 train_time:85702ms step_avg:58.78ms step:1459/1600 train_time:85790ms step_avg:58.80ms step:1460/1600 train_time:85876ms step_avg:58.82ms step:1461/1600 train_time:85964ms step_avg:58.84ms step:1462/1600 train_time:86049ms step_avg:58.86ms step:1463/1600 train_time:86138ms step_avg:58.88ms step:1464/1600 train_time:86224ms step_avg:58.90ms step:1465/1600 train_time:86313ms step_avg:58.92ms step:1466/1600 train_time:86398ms step_avg:58.93ms step:1467/1600 train_time:86486ms step_avg:58.95ms step:1468/1600 train_time:86571ms step_avg:58.97ms step:1469/1600 train_time:86659ms step_avg:58.99ms step:1470/1600 train_time:86745ms step_avg:59.01ms step:1471/1600 train_time:86832ms step_avg:59.03ms step:1472/1600 train_time:86919ms step_avg:59.05ms step:1473/1600 train_time:87006ms step_avg:59.07ms step:1474/1600 train_time:87091ms step_avg:59.09ms step:1475/1600 train_time:87181ms step_avg:59.11ms step:1476/1600 train_time:87265ms step_avg:59.12ms step:1477/1600 train_time:87354ms step_avg:59.14ms step:1478/1600 train_time:87438ms step_avg:59.16ms step:1479/1600 train_time:87527ms step_avg:59.18ms step:1480/1600 train_time:87612ms step_avg:59.20ms step:1481/1600 train_time:87701ms step_avg:59.22ms step:1482/1600 train_time:87785ms step_avg:59.23ms step:1483/1600 train_time:87874ms step_avg:59.25ms step:1484/1600 train_time:87960ms step_avg:59.27ms step:1485/1600 train_time:88047ms step_avg:59.29ms step:1486/1600 train_time:88132ms step_avg:59.31ms step:1487/1600 train_time:88221ms step_avg:59.33ms step:1488/1600 train_time:88305ms step_avg:59.35ms step:1489/1600 train_time:88394ms step_avg:59.36ms step:1490/1600 train_time:88480ms step_avg:59.38ms step:1491/1600 train_time:88569ms step_avg:59.40ms step:1492/1600 train_time:88653ms step_avg:59.42ms step:1493/1600 train_time:88742ms step_avg:59.44ms step:1494/1600 train_time:88826ms step_avg:59.46ms step:1495/1600 train_time:88916ms step_avg:59.48ms step:1496/1600 train_time:89001ms step_avg:59.49ms step:1497/1600 train_time:89089ms step_avg:59.51ms step:1498/1600 train_time:89174ms step_avg:59.53ms step:1499/1600 train_time:89262ms step_avg:59.55ms step:1500/1600 train_time:89347ms step_avg:59.56ms step:1500/1600 val_loss:3.3079 train_time:89418ms step_avg:59.61ms step:1501/1600 train_time:89439ms step_avg:59.59ms step:1502/1600 train_time:89526ms step_avg:59.60ms step:1503/1600 train_time:89618ms step_avg:59.63ms step:1504/1600 train_time:89704ms step_avg:59.64ms step:1505/1600 train_time:89791ms step_avg:59.66ms step:1506/1600 train_time:89875ms step_avg:59.68ms step:1507/1600 train_time:89963ms step_avg:59.70ms step:1508/1600 train_time:90047ms step_avg:59.71ms step:1509/1600 train_time:90133ms step_avg:59.73ms step:1510/1600 train_time:90218ms step_avg:59.75ms step:1511/1600 train_time:90305ms step_avg:59.77ms step:1512/1600 train_time:90390ms step_avg:59.78ms step:1513/1600 train_time:90480ms step_avg:59.80ms step:1514/1600 train_time:90568ms step_avg:59.82ms step:1515/1600 train_time:90659ms step_avg:59.84ms step:1516/1600 train_time:90744ms step_avg:59.86ms step:1517/1600 train_time:90832ms step_avg:59.88ms step:1518/1600 train_time:90917ms step_avg:59.89ms step:1519/1600 train_time:91004ms step_avg:59.91ms step:1520/1600 train_time:91088ms step_avg:59.93ms step:1521/1600 train_time:91175ms step_avg:59.94ms step:1522/1600 train_time:91260ms step_avg:59.96ms step:1523/1600 train_time:91348ms step_avg:59.98ms step:1524/1600 train_time:91433ms step_avg:60.00ms step:1525/1600 train_time:91523ms step_avg:60.01ms step:1526/1600 train_time:91609ms step_avg:60.03ms step:1527/1600 train_time:91698ms step_avg:60.05ms step:1528/1600 train_time:91783ms step_avg:60.07ms step:1529/1600 train_time:91871ms step_avg:60.09ms step:1530/1600 train_time:91957ms step_avg:60.10ms step:1531/1600 train_time:92045ms step_avg:60.12ms step:1532/1600 train_time:92130ms step_avg:60.14ms step:1533/1600 train_time:92216ms step_avg:60.15ms step:1534/1600 train_time:92301ms step_avg:60.17ms step:1535/1600 train_time:92390ms step_avg:60.19ms step:1536/1600 train_time:92476ms step_avg:60.21ms step:1537/1600 train_time:92565ms step_avg:60.22ms step:1538/1600 train_time:92651ms step_avg:60.24ms step:1539/1600 train_time:92739ms step_avg:60.26ms step:1540/1600 train_time:92825ms step_avg:60.28ms step:1541/1600 train_time:92914ms step_avg:60.29ms step:1542/1600 train_time:93000ms step_avg:60.31ms step:1543/1600 train_time:93088ms step_avg:60.33ms step:1544/1600 train_time:93171ms step_avg:60.34ms step:1545/1600 train_time:93258ms step_avg:60.36ms step:1546/1600 train_time:93344ms step_avg:60.38ms step:1547/1600 train_time:93432ms step_avg:60.40ms step:1548/1600 train_time:93518ms step_avg:60.41ms step:1549/1600 train_time:93607ms step_avg:60.43ms step:1550/1600 train_time:93693ms step_avg:60.45ms step:1551/1600 train_time:93781ms step_avg:60.47ms step:1552/1600 train_time:93867ms step_avg:60.48ms step:1553/1600 train_time:93957ms step_avg:60.50ms step:1554/1600 train_time:94042ms step_avg:60.52ms step:1555/1600 train_time:94128ms step_avg:60.53ms step:1556/1600 train_time:94212ms step_avg:60.55ms step:1557/1600 train_time:94300ms step_avg:60.57ms step:1558/1600 train_time:94388ms step_avg:60.58ms step:1559/1600 train_time:94474ms step_avg:60.60ms step:1560/1600 train_time:94560ms step_avg:60.62ms step:1561/1600 train_time:94657ms step_avg:60.64ms step:1562/1600 train_time:94742ms step_avg:60.65ms step:1563/1600 train_time:94830ms step_avg:60.67ms step:1564/1600 train_time:94915ms step_avg:60.69ms step:1565/1600 train_time:95003ms step_avg:60.71ms step:1566/1600 train_time:95089ms step_avg:60.72ms step:1567/1600 train_time:95177ms step_avg:60.74ms step:1568/1600 train_time:95262ms step_avg:60.75ms step:1569/1600 train_time:95351ms step_avg:60.77ms step:1570/1600 train_time:95437ms step_avg:60.79ms step:1571/1600 train_time:95527ms step_avg:60.81ms step:1572/1600 train_time:95612ms step_avg:60.82ms step:1573/1600 train_time:95702ms step_avg:60.84ms step:1574/1600 train_time:95788ms step_avg:60.86ms step:1575/1600 train_time:95875ms step_avg:60.87ms step:1576/1600 train_time:95962ms step_avg:60.89ms step:1577/1600 train_time:96054ms step_avg:60.91ms step:1578/1600 train_time:96137ms step_avg:60.92ms step:1579/1600 train_time:96225ms step_avg:60.94ms step:1580/1600 train_time:96310ms step_avg:60.96ms step:1581/1600 train_time:96398ms step_avg:60.97ms step:1582/1600 train_time:96485ms step_avg:60.99ms step:1583/1600 train_time:96573ms step_avg:61.01ms step:1584/1600 train_time:96659ms step_avg:61.02ms step:1585/1600 train_time:96749ms step_avg:61.04ms step:1586/1600 train_time:96834ms step_avg:61.06ms step:1587/1600 train_time:96923ms step_avg:61.07ms step:1588/1600 train_time:97008ms step_avg:61.09ms step:1589/1600 train_time:97096ms step_avg:61.11ms step:1590/1600 train_time:97183ms step_avg:61.12ms step:1591/1600 train_time:97271ms step_avg:61.14ms step:1592/1600 train_time:97356ms step_avg:61.15ms step:1593/1600 train_time:97446ms step_avg:61.17ms step:1594/1600 train_time:97531ms step_avg:61.19ms step:1595/1600 train_time:97620ms step_avg:61.20ms step:1596/1600 train_time:97706ms step_avg:61.22ms step:1597/1600 train_time:97794ms step_avg:61.24ms step:1598/1600 train_time:97879ms step_avg:61.25ms step:1599/1600 train_time:97968ms step_avg:61.27ms step:1600/1600 train_time:98053ms step_avg:61.28ms step:1600/1600 val_loss:3.2784 train_time:98124ms step_avg:61.33ms peak memory allocated: 30264 MiB reserved: 46238 MiB