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, pairwise 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 # ----------------------------------------------------------------------------- # Distributed training setup 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 grad_scale = 2 / grad_accum_steps # consistent grad magnitudes between different num_devices 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. # ----------------------------------------------------------------------------- # 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, paired=False): super().__init__() self.head_dim = head_dim self.max_seq_len = max_seq_len self.paired = paired 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) if not self.paired: 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 ) else: 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 = 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) if not self.paired: theta = torch.outer(t, self.angular_freq) self.factor1.copy_(theta.cos()) self.factor2.copy_(theta.sin()) else: 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, paired: bool = False): super().__init__() self.num_heads = num_heads self.head_dim = head_dim self.dim = dim self.hdim = num_heads * head_dim self.paired = paired 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) max_len = args.train_max_seq_len if self.training else (args.val_batch_size // (grad_accum_steps * world_size)) q, k = norm(q), norm(k) # QK norm @Grad62304977 if not self.paired: 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 else: # Paired heads: adjacent heads' queries attend to each other's keys. # Two copies of the input stream are interleaved to achieve this, which: # - doubles the length of each sequence # - halves the effective window size 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) seqlens = 2 * seqlens max_len = 2 * max_len # 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 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 self.attn = CausalSelfAttention(dim, head_dim, num_heads, paired=use_paired_head) if has_attn else 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 self.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.Parameter(torch.zeros(5 * self.vocab_size, model_dim, dtype=torch.bfloat16)) self.value_embeds.label = 'value_embed' # parameter banks for attention and value embedding gate weights self.attn_gate_bank = nn.Parameter(torch.zeros(10, num_heads, 12)) # 10 layers self.attn_gate_bank.label = 'attn_gate_bank' self.ve_gate_bank = nn.Parameter(torch.zeros(5, num_heads, 12)) # 5 unique gates self.ve_gate_bank.label = 've_gate_bank' # ----------------------------------- # Parameter banks for sharded optimization, by @chrisjmccormick # Identify which layers have attention/MLP # Attention is skipped in layer 6 by @YouJiacheng self.attn_layer_indices = [i for i in range(num_layers) if i != 6] # All layers have MLP (At 11 layers--dropped first layer @EmelyanenkoK) self.mlp_layer_indices = list(range(num_layers)) hdim = num_heads * head_dim mlp_hdim = 4 * model_dim # Create index mappings: layer_idx -> bank_idx self.layer_to_attn_idx = {layer_idx: bank_idx for bank_idx, layer_idx in enumerate(self.attn_layer_indices)} self.layer_to_mlp_idx = {layer_idx: bank_idx for bank_idx, layer_idx in enumerate(self.mlp_layer_indices)} # Attention bank: stores QKVO weights for all attention layers # merged QKVO weights: suggested by many, implemented by @fernbear.bsky.social, and further improved by @YouJiacheng # https://x.com/hi_tysam/status/1879699187107033311 # Simplified layout by @chrisjmccormick # Shape: (num_attn_layers, 4*model_dim, hdim) = (10, 3072, 768) # Reshape for sharding: (40, 768, 768) for even distribution across 8 GPUs self.attn_bank = nn.Parameter(torch.empty(len(self.attn_layer_indices), 4 * model_dim, hdim)) self.attn_bank.label = 'attn' self.attn_bank.reshape = (len(self.attn_layer_indices) * 4, hdim, hdim) # (40, 768, 768) # MLP bank: stores c_fc and c_proj for all MLP layers # Shape: (num_mlp_layers + padding, 2, mlp_hdim, model_dim) = (12, 2, 3072, 768) # We add 1 padding layer (index 11) to get 12*2=24 matrices for even distribution across 8 GPUs # Reshape for sharding: (24, 3072, 768) num_mlp_with_padding = len(self.mlp_layer_indices) + 1 # 11 + 1 = 12 self.mlp_bank = nn.Parameter(torch.empty(num_mlp_with_padding, 2, mlp_hdim, model_dim)) self.mlp_bank.label = 'mlp' self.mlp_bank.reshape = (num_mlp_with_padding * 2, mlp_hdim, model_dim) # (24, 3072, 768) # improved init scale by @YouJiacheng and @srashedll std = 0.5 * model_dim ** -0.5 bound = (3 ** 0.5) * std with torch.no_grad(): self.attn_bank.uniform_(-bound, bound) self.mlp_bank[:, 0, :, :].uniform_(-bound, 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 = Yarn(head_dim, max_seq_len, paired=True) # 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, self.vocab_size, use_fp8=use_fp8, x_s=100/448, w_s=1.6/448, grad_s=grad_scale * 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(self.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' @staticmethod @torch.compile(dynamic=False, fullgraph=True) def _compute_bigram_hash(x: Tensor, mod: int) -> Tensor: """ Computes bigram hash on GPU for each position using [prev_token, curr_token]. Mathematically identical to the CPU version but computed on device. """ rand_int_1 = 36313 rand_int_2 = 27191 result = torch.empty_like(x) result[0] = mod result[1:] = torch.bitwise_xor(rand_int_1 * x[1:], rand_int_2 * x[:-1]) % mod return result def forward(self, input_seq: Tensor, target_seq: Tensor, seqlens: 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 bm_sizes = [ws_short, ws_short, ws_short, ws_long, ws_short, ws_short, None, ws_short, ws_short, ws_short, ws_long] assert len(bm_sizes) == self.num_layers key_offset = [b==ws_long 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) # Compute bigram hash on GPU (moved from CPU data loader) bigram_seq = self._compute_bigram_hash(input_seq, args.bigram_vocab_size - 1) x0_bigram = self.bigram_embed(bigram_seq)[None] # Value embeddings - always computed (not precomputed) ve = self.value_embeds.view(5, self.vocab_size, -1)[:, input_seq] # 01 ... 234 structure on token value embeddings by @photomz ve = [ve[0], ve[1]] + [None] * (self.num_layers - 5) + [ve[2], ve[3], ve[4]] assert len(ve) == self.num_layers # smear token embed forward 1 position @classiclarryd smear_gate_out = smear_lambda * torch.sigmoid(self.smear_gate(x[1:, :self.smear_gate.weight.size(-1)])) x = torch.cat([x[:1], x[1:] + smear_gate_out * x[:-1]]) x = x0 = norm(x[None]) # unbind gate banks to avoid select_backwards kernel ag = [w.bfloat16() for w in self.attn_gate_bank.unbind(0)] veg = [w.bfloat16() for w in self.ve_gate_bank.unbind(0)] attn_gates = ag[:6] + [None] + ag[6:] ve_gates = [veg[0], veg[1]] + [None] * (self.num_layers - 5) + [veg[2], veg[3], veg[4]] assert len(attn_gates) == self.num_layers assert len(ve_gates) == self.num_layers # unbind weight banks to avoid select_backwards kernel attn_weights = self.attn_bank.unbind(0) # tuple of [4*dim, hdim] tensors mlp_fcs = self.mlp_bank[:, 0, :, :].unbind(0) # tuple of [mlp_hdim, dim] tensors mlp_projs = self.mlp_bank[:, 1, :, :].unbind(0) # tuple of [mlp_hdim, dim] tensors for i in range(self.num_layers): yarn = self.yarn_paired_head if i in self.paired_head_layers else self.yarn attn_args = AttnArgs( ve=ve[i], sa_lambdas=sa_lambdas[i], seqlens=seqlens, bm_size=bm_sizes[i], yarn=yarn, key_offset=key_offset[i], attn_gate_w=attn_gates[i], ve_gate_w=ve_gates[i] ) if i in skip_out: skip_gate_out = torch.sigmoid(skip_lambda) * 2 * torch.sigmoid(self.skip_gate(x0[..., :self.skip_gate.weight.size(-1)])) x = x + skip_gate_out * skip_connections.pop() if i == 0: x = (resid_lambdas[0] + x0_lambdas[0]) * x + bigram_lambdas[0] * x0_bigram else: x = resid_lambdas[i] * x + x0_lambdas[i] * x0 + bigram_lambdas[i] * x0_bigram # Get weights for this layer from banks qkvo_w = attn_weights[self.layer_to_attn_idx[i]] if i in self.layer_to_attn_idx else None c_fc = mlp_fcs[self.layer_to_mlp_idx[i]] if i in self.layer_to_mlp_idx else None c_proj = mlp_projs[self.layer_to_mlp_idx[i]] if i in self.layer_to_mlp_idx else None x = self.blocks[i](x, attn_args, qkvo_w, c_fc, c_proj) if i in skip_in: skip_connections.append(x) if i == backout_layer: x_backout = x # back out contributions from first 7 layers that are only required for downstream context and not direct prediction x -= backout_lambda * x_backout x = norm(x) logits = self.lm_head(x) # @Grad62304977 added tanh softcapping following Gemma 2 paper, @KoszarskyB reduced it from 30 to 15 # @YouJiacheng shifted it by +15 (2*sigmoid(2*x)=tanh(x)+1). @classiclarryd updated to 23*sigmoid((logits+5)/7.5) if self.training: losses = FusedSoftcappedCrossEntropy.apply(logits.view(-1, logits.size(-1)), target_seq, mtp_weights, 23.0, 5.0, 7.5) 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 Shard: def __init__(self, tokens: Tensor, world_size: int = 1): self.tokens = tokens self.size = tokens.numel() self.world_size = world_size self.i = 0 # Partial index now, full index async self.bos_idx = (tokens[:6_000_000] == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self._full_idx = None self._loader_thread = None self._ready = threading.Event() self._loader_thread = threading.Thread(target=self._scan) self._loader_thread.start() def _scan(self): self._full_idx = (self.tokens == BOS_ID).nonzero(as_tuple=True)[0].to(torch.int64).cpu().numpy() self._ready.set() def _maybe_switch(self): # Switch to full index as soon as async scan completes if self.bos_idx is not self._full_idx and self._ready.is_set(): self._loader_thread.join() self.bos_idx = self._full_idx def next_batch(self, num_tokens_local: int, max_seq_len: int): self._maybe_switch() 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 return starts, ends @staticmethod def load_async(file: Path, world_size: int = 1): """Returns getter function for async shard loading""" result = {} ready = threading.Event() def load(): tokens = _load_data_shard(file) result['shard'] = Shard(tokens, world_size) ready.set() thread = threading.Thread(target=load) thread.start() def get(): ready.wait() thread.join() return result['shard'] return get 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: shard = Shard(tokens, world_size) next_shard_getter = Shard.load_async(next(file_iter), world_size) 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 = shard.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. shard = next_shard_getter() tokens = shard.tokens try: next_shard_getter = Shard.load_async(next(file_iter), world_size) except StopIteration: next_shard_getter = None # no more shards to preload 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 hash computation moved to GPU in forward() 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), ) 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 @dataclass class Hyperparameters: # data data_path = os.environ.get("DATA_PATH", ".") train_files: str = os.path.join(data_path, "data/fineweb10B/fineweb_train_*.bin") # input .bin to train on val_files: str = os.path.join(data_path, "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_max_seq_len: int = 128 * 16 val_batch_size: int = 4 * 64 * 1024 * 8 # schedule num_scheduled_iterations: int = 1515 # 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 # 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 # bigram hash embedding bigram_vocab_size: int = 50304 * 5 args = Hyperparameters() @dataclass class TrainingStage: lr_mul: float batch_size: int window_sizes: tuple[int, int] # (short, long) in block units mtp_weights_start: list[float] mtp_weights_end: list[float] duration: float = None class TrainingSchedule: """ Training schedule initialized via TRAINING_STAGES 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 5. Batch size schedule of 8 -> 16 -> 24 6. Post training extension of long windows from 13 to 20 """ def __init__(self, stages: list[TrainingStage], scheduled_iterations: int, extension_iterations: int, cooldown_frac: float = 0.5, split_embed_stage: int = 2, ws_post_yarn_ext: int = 20): self.stages = stages self.scheduled_iterations = scheduled_iterations self.cooldown_frac = cooldown_frac # increase final validation ws, used for YaRN extension and short window size @classiclarryd self.ws_post_yarn_ext = ws_post_yarn_ext self.total_steps = self.scheduled_iterations + extension_iterations # Build stage boundaries (last is extension stage) ends = [0] + [round(c * scheduled_iterations) for c in accumulate(s.duration for s in stages[:-1])] + [self.total_steps] assert self.scheduled_iterations == ends[-2] self.boundaries = list(pairwise(ends)) # Split embed at specified stage (ensure odd step for Adam) self.split_step = self.boundaries[split_embed_stage][0] | 1 # Precompute MTP weights for all steps self.mtp_weights = [] for step in range(self.total_steps + 1): stage, t = self.lookup(step) w = [a + (b - a) * t for a, b in zip(stage.mtp_weights_start, stage.mtp_weights_end)] self.mtp_weights.append(torch.tensor(w, device=device)) def lookup(self, step: int) -> tuple[TrainingStage, float]: # Returns stage and % of the way through that stage for i, (start, end) in enumerate(self.boundaries): if step < end: t = (step - start) / (end - start) return self.stages[i], t return self.stages[-1], 1.0 def get_lr(self, step: int) -> float: # learning rate schedule: tied to batch size schedule, with cooldown at the end stage, _ = self.lookup(step) lr = stage.lr_mul cd_start = int(self.scheduled_iterations * (1 - self.cooldown_frac)) if step >= cd_start: t = min(1.0, (step - cd_start) / (self.scheduled_iterations - cd_start)) lr = lr * (1 - t) + 0.1 * t return lr # window_sizes are in units of `block_size` tokens (defined in TrainingManager) TRAINING_STAGES = [ TrainingStage(duration=1/3, batch_size=8 * 2048 * 8, window_sizes=(1, 3), lr_mul=1.0, mtp_weights_start=[1.0, 0.5, 0.25], mtp_weights_end=[1.0, 0.5, 0.0]), TrainingStage(duration=1/3, batch_size=16 * 2048 * 8, window_sizes=(3, 7), lr_mul=1.52, # (16/8)**0.6 mtp_weights_start=[1.0, 0.5], mtp_weights_end=[1.0, 0.0]), TrainingStage(duration=1/3, batch_size=24 * 2048 * 8, window_sizes=(5, 11), lr_mul=1.73, # (24/8)**0.5 mtp_weights_start=[1.0], mtp_weights_end=[1.0]), # extension stage TrainingStage(batch_size=24 * 2048 * 8, window_sizes=(6, 13), lr_mul=1.0, # lr_mul is not used mtp_weights_start=[1.0], mtp_weights_end=[1.0]), ] training_schedule = TrainingSchedule(TRAINING_STAGES, args.num_scheduled_iterations, args.num_extension_iterations, cooldown_frac=0.55) 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 = training_schedule.total_steps - 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. 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 """ def __init__(self, model): self.model = model self.block_size = 128 # - 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}, "value_embed": {"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 "value_embed", "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 = training_schedule.split_step self.reset() def apply_final_ws_ext(self): self.ws_long = training_schedule.ws_post_yarn_ext def get_forward_args(self): return ForwardScheduleConfig( mtp_weights = self.mtp_weights, ws_short = self.ws_short * self.block_size, ws_long = self.ws_long * self.block_size ) def _is_adam_step(self, step: int): """Adam params are only updated on odd steps.""" return step % 2 == 1 def get_transition_steps(self): return [start for start, _ in training_schedule.boundaries[1:]] def advance_schedule(self, step: int): stage, _ = training_schedule.lookup(step) self.ws_short, new_ws_long = stage.window_sizes if new_ws_long != self.ws_long: self.model.yarn.apply(self.ws_long * self.block_size, new_ws_long * self.block_size) self.model.yarn_paired_head.apply(self.ws_long * self.block_size, new_ws_long * self.block_size) new_batch_size = stage.batch_size 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 = training_schedule.mtp_weights[step] def step_optimizers(self, step: int): step_lr = training_schedule.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() stage, _ = training_schedule.lookup(0) self.ws_short, self.ws_long = stage.window_sizes self.batch_size = stage.batch_size self.model.yarn.reset() self.model.yarn_paired_head.reset() def get_state(self): return copy.deepcopy(self.optimizer.state_dict()) # ----------------------------------------------------------------------------- # int main # 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, TRAINING_STAGES[0].batch_size, 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 = next(val_loader) model(inputs, targets, cum_seqlens, 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 = train_loader.send(send_args) (model(inputs, targets, cum_seqlens, training_manager.get_forward_args()) * grad_scale).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, TRAINING_STAGES[0].batch_size, 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 = training_schedule.total_steps 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 = next(val_loader) val_loss += model(inputs, targets, cum_seqlens, 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 = train_loader.send(training_manager.train_loader_send_args) (model(inputs, targets, cum_seqlens, training_manager.get_forward_args()) * grad_scale).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 @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.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 # Hardcoded configs based on H100 autotuning if K == 768: BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K = 128, 128, 64 num_stages, num_warps = 4, 4 else: BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K = 64, 128, 128 num_stages, num_warps = 4, 4 grid = (batch_size * triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(M, 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), BLOCK_SIZE_M=BLOCK_SIZE_M, BLOCK_SIZE_N=BLOCK_SIZE_N, BLOCK_SIZE_K=BLOCK_SIZE_K, GROUP_SIZE_M=8, LOWER_UPPER=1, num_stages=num_stages, num_warps=num_warps, ) return out @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 # Hardcoded config based on H100 autotuning (M=768) BLOCK_SIZE_M, BLOCK_SIZE_N, BLOCK_SIZE_K = 128, 128, 64 num_stages, num_warps = 4, 4 grid = (batch_size * triton.cdiv(M, BLOCK_SIZE_M) * triton.cdiv(M, 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, BLOCK_SIZE_M=BLOCK_SIZE_M, BLOCK_SIZE_N=BLOCK_SIZE_N, BLOCK_SIZE_K=BLOCK_SIZE_K, GROUP_SIZE_M=8, LOWER_UPPER=1, num_stages=num_stages, num_warps=num_warps, ) 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.12.7 (main, Jan 31 2026, 04:21:49) [GCC 13.2.0] Running PyTorch 2.10.0.dev20251210+cu126 compiled for CUDA 12.6 Running Triton version 3.6.0 Sun Feb 1 06:03:44 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:63:00.0 Off | 0 | | N/A 32C P0 115W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 35C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:71:00.0 Off | 0 | | N/A 37C P0 123W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:79:00.0 Off | 0 | | N/A 33C P0 124W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:7F:00.0 Off | 0 | | N/A 31C P0 118W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:87:00.0 Off | 0 | | N/A 37C P0 118W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:8D:00.0 Off | 0 | | N/A 35C P0 122W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:95:00.0 Off | 0 | | N/A 33C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 12926 C /usr/local/bin/python 1510MiB | | 1 N/A N/A 12927 C /usr/local/bin/python 1510MiB | | 2 N/A N/A 12928 C /usr/local/bin/python 1510MiB | | 3 N/A N/A 12929 C /usr/local/bin/python 1510MiB | | 4 N/A N/A 12930 C /usr/local/bin/python 1510MiB | | 5 N/A N/A 12931 C /usr/local/bin/python 1510MiB | | 6 N/A N/A 12932 C /usr/local/bin/python 1510MiB | | 7 N/A N/A 12933 C /usr/local/bin/python 1510MiB | +-----------------------------------------------------------------------------------------+ ==================================================================================================== Compiling model and warming up kernels (~7 minutes on first execution) Sampling steps [0, 1, 2, 504, 505, 506, 1009, 1010, 1011, 1514, 1515, 1516] for warmup Resetting Model step:0/1555 val_loss:10.8301 train_time:0ms step_avg:0.03ms step:1/1555 train_time:80ms step_avg:80.22ms step:2/1555 train_time:104ms step_avg:51.92ms step:3/1555 train_time:126ms step_avg:42.13ms step:4/1555 train_time:149ms step_avg:37.27ms step:5/1555 train_time:179ms step_avg:35.85ms step:6/1555 train_time:217ms step_avg:36.19ms step:7/1555 train_time:248ms step_avg:35.40ms step:8/1555 train_time:285ms step_avg:35.67ms step:9/1555 train_time:316ms step_avg:35.11ms step:10/1555 train_time:354ms step_avg:35.40ms step:11/1555 train_time:384ms step_avg:34.94ms step:12/1555 train_time:422ms step_avg:35.20ms step:13/1555 train_time:453ms step_avg:34.86ms step:14/1555 train_time:491ms step_avg:35.06ms step:15/1555 train_time:522ms step_avg:34.80ms step:16/1555 train_time:560ms step_avg:35.00ms step:17/1555 train_time:591ms step_avg:34.76ms step:18/1555 train_time:628ms step_avg:34.91ms step:19/1555 train_time:659ms step_avg:34.69ms step:20/1555 train_time:697ms step_avg:34.85ms step:21/1555 train_time:728ms step_avg:34.68ms step:22/1555 train_time:766ms step_avg:34.81ms step:23/1555 train_time:797ms step_avg:34.65ms step:24/1555 train_time:835ms step_avg:34.78ms step:25/1555 train_time:865ms step_avg:34.62ms step:26/1555 train_time:903ms step_avg:34.75ms step:27/1555 train_time:934ms step_avg:34.60ms step:28/1555 train_time:972ms step_avg:34.71ms step:29/1555 train_time:1003ms step_avg:34.57ms step:30/1555 train_time:1040ms step_avg:34.67ms step:31/1555 train_time:1072ms step_avg:34.57ms step:32/1555 train_time:1109ms step_avg:34.66ms step:33/1555 train_time:1141ms step_avg:34.56ms step:34/1555 train_time:1179ms step_avg:34.68ms step:35/1555 train_time:1210ms step_avg:34.57ms step:36/1555 train_time:1248ms step_avg:34.65ms step:37/1555 train_time:1279ms step_avg:34.56ms step:38/1555 train_time:1317ms step_avg:34.65ms step:39/1555 train_time:1347ms step_avg:34.55ms step:40/1555 train_time:1385ms step_avg:34.63ms step:41/1555 train_time:1416ms step_avg:34.55ms step:42/1555 train_time:1455ms step_avg:34.63ms step:43/1555 train_time:1485ms step_avg:34.54ms step:44/1555 train_time:1523ms step_avg:34.62ms step:45/1555 train_time:1554ms step_avg:34.53ms step:46/1555 train_time:1592ms step_avg:34.60ms step:47/1555 train_time:1623ms step_avg:34.53ms step:48/1555 train_time:1661ms step_avg:34.60ms step:49/1555 train_time:1691ms step_avg:34.52ms step:50/1555 train_time:1729ms step_avg:34.58ms step:51/1555 train_time:1760ms step_avg:34.51ms step:52/1555 train_time:1798ms step_avg:34.57ms step:53/1555 train_time:1829ms step_avg:34.51ms step:54/1555 train_time:1866ms step_avg:34.56ms step:55/1555 train_time:1898ms step_avg:34.50ms step:56/1555 train_time:1935ms step_avg:34.56ms step:57/1555 train_time:1966ms step_avg:34.49ms step:58/1555 train_time:2004ms step_avg:34.55ms step:59/1555 train_time:2035ms step_avg:34.48ms step:60/1555 train_time:2073ms step_avg:34.55ms step:61/1555 train_time:2104ms step_avg:34.49ms step:62/1555 train_time:2141ms step_avg:34.54ms step:63/1555 train_time:2172ms step_avg:34.48ms step:64/1555 train_time:2210ms step_avg:34.53ms step:65/1555 train_time:2241ms step_avg:34.48ms step:66/1555 train_time:2279ms step_avg:34.52ms step:67/1555 train_time:2309ms step_avg:34.47ms step:68/1555 train_time:2348ms step_avg:34.52ms step:69/1555 train_time:2379ms step_avg:34.48ms step:70/1555 train_time:2417ms step_avg:34.53ms step:71/1555 train_time:2448ms step_avg:34.48ms step:72/1555 train_time:2486ms step_avg:34.52ms step:73/1555 train_time:2517ms step_avg:34.48ms step:74/1555 train_time:2555ms step_avg:34.53ms step:75/1555 train_time:2585ms step_avg:34.47ms step:76/1555 train_time:2624ms step_avg:34.52ms step:77/1555 train_time:2655ms step_avg:34.48ms step:78/1555 train_time:2693ms step_avg:34.52ms step:79/1555 train_time:2724ms step_avg:34.47ms step:80/1555 train_time:2761ms step_avg:34.52ms step:81/1555 train_time:2792ms step_avg:34.47ms step:82/1555 train_time:2829ms step_avg:34.51ms step:83/1555 train_time:2861ms step_avg:34.46ms step:84/1555 train_time:2898ms step_avg:34.50ms step:85/1555 train_time:2929ms step_avg:34.45ms step:86/1555 train_time:2966ms step_avg:34.49ms step:87/1555 train_time:2997ms step_avg:34.45ms step:88/1555 train_time:3035ms step_avg:34.49ms step:89/1555 train_time:3066ms step_avg:34.45ms step:90/1555 train_time:3104ms step_avg:34.49ms step:91/1555 train_time:3135ms step_avg:34.45ms step:92/1555 train_time:3173ms step_avg:34.49ms step:93/1555 train_time:3203ms step_avg:34.45ms step:94/1555 train_time:3241ms step_avg:34.48ms step:95/1555 train_time:3272ms step_avg:34.45ms step:96/1555 train_time:3310ms step_avg:34.48ms step:97/1555 train_time:3341ms step_avg:34.44ms step:98/1555 train_time:3379ms step_avg:34.48ms step:99/1555 train_time:3410ms step_avg:34.44ms step:100/1555 train_time:3447ms step_avg:34.47ms step:101/1555 train_time:3478ms step_avg:34.44ms step:102/1555 train_time:3516ms step_avg:34.47ms step:103/1555 train_time:3548ms step_avg:34.45ms step:104/1555 train_time:3586ms step_avg:34.48ms step:105/1555 train_time:3617ms step_avg:34.44ms step:106/1555 train_time:3654ms step_avg:34.47ms step:107/1555 train_time:3685ms step_avg:34.44ms step:108/1555 train_time:3723ms step_avg:34.47ms step:109/1555 train_time:3754ms step_avg:34.44ms step:110/1555 train_time:3791ms step_avg:34.47ms step:111/1555 train_time:3823ms step_avg:34.44ms step:112/1555 train_time:3861ms step_avg:34.47ms step:113/1555 train_time:3891ms step_avg:34.44ms step:114/1555 train_time:3929ms step_avg:34.46ms step:115/1555 train_time:3960ms step_avg:34.43ms step:116/1555 train_time:3997ms step_avg:34.46ms step:117/1555 train_time:4028ms step_avg:34.43ms step:118/1555 train_time:4066ms step_avg:34.46ms step:119/1555 train_time:4097ms step_avg:34.43ms step:120/1555 train_time:4135ms step_avg:34.46ms step:121/1555 train_time:4165ms step_avg:34.42ms step:122/1555 train_time:4203ms step_avg:34.45ms step:123/1555 train_time:4234ms step_avg:34.42ms step:124/1555 train_time:4271ms step_avg:34.45ms step:125/1555 train_time:4302ms step_avg:34.42ms step:126/1555 train_time:4340ms step_avg:34.45ms step:127/1555 train_time:4371ms step_avg:34.42ms step:128/1555 train_time:4409ms step_avg:34.45ms step:129/1555 train_time:4440ms step_avg:34.42ms step:130/1555 train_time:4479ms step_avg:34.45ms step:131/1555 train_time:4509ms step_avg:34.42ms step:132/1555 train_time:4546ms step_avg:34.44ms step:133/1555 train_time:4577ms step_avg:34.42ms step:134/1555 train_time:4615ms step_avg:34.44ms step:135/1555 train_time:4646ms step_avg:34.41ms step:136/1555 train_time:4684ms step_avg:34.44ms step:137/1555 train_time:4715ms step_avg:34.41ms step:138/1555 train_time:4752ms step_avg:34.44ms step:139/1555 train_time:4783ms step_avg:34.41ms step:140/1555 train_time:4821ms step_avg:34.44ms step:141/1555 train_time:4852ms step_avg:34.41ms step:142/1555 train_time:4889ms step_avg:34.43ms step:143/1555 train_time:4920ms step_avg:34.41ms step:144/1555 train_time:4958ms step_avg:34.43ms step:145/1555 train_time:4988ms step_avg:34.40ms step:146/1555 train_time:5026ms step_avg:34.42ms step:147/1555 train_time:5057ms step_avg:34.40ms step:148/1555 train_time:5095ms step_avg:34.43ms step:149/1555 train_time:5126ms step_avg:34.40ms step:150/1555 train_time:5163ms step_avg:34.42ms step:151/1555 train_time:5194ms step_avg:34.40ms step:152/1555 train_time:5231ms step_avg:34.42ms step:153/1555 train_time:5262ms step_avg:34.40ms step:154/1555 train_time:5300ms step_avg:34.42ms step:155/1555 train_time:5331ms step_avg:34.39ms step:156/1555 train_time:5368ms step_avg:34.41ms step:157/1555 train_time:5400ms step_avg:34.39ms step:158/1555 train_time:5437ms step_avg:34.41ms step:159/1555 train_time:5468ms step_avg:34.39ms step:160/1555 train_time:5506ms step_avg:34.41ms step:161/1555 train_time:5537ms step_avg:34.39ms step:162/1555 train_time:5575ms step_avg:34.41ms step:163/1555 train_time:5606ms step_avg:34.39ms step:164/1555 train_time:5643ms step_avg:34.41ms step:165/1555 train_time:5675ms step_avg:34.39ms step:166/1555 train_time:5712ms step_avg:34.41ms step:167/1555 train_time:5744ms step_avg:34.39ms step:168/1555 train_time:5782ms step_avg:34.42ms step:169/1555 train_time:5814ms step_avg:34.40ms step:170/1555 train_time:5851ms step_avg:34.42ms step:171/1555 train_time:5883ms step_avg:34.40ms step:172/1555 train_time:5921ms step_avg:34.42ms step:173/1555 train_time:5951ms step_avg:34.40ms step:174/1555 train_time:5989ms step_avg:34.42ms step:175/1555 train_time:6021ms step_avg:34.40ms step:176/1555 train_time:6058ms step_avg:34.42ms step:177/1555 train_time:6089ms step_avg:34.40ms step:178/1555 train_time:6127ms step_avg:34.42ms step:179/1555 train_time:6157ms step_avg:34.40ms step:180/1555 train_time:6195ms step_avg:34.42ms step:181/1555 train_time:6226ms step_avg:34.40ms step:182/1555 train_time:6264ms step_avg:34.42ms step:183/1555 train_time:6294ms step_avg:34.39ms step:184/1555 train_time:6331ms step_avg:34.41ms step:185/1555 train_time:6362ms step_avg:34.39ms step:186/1555 train_time:6399ms step_avg:34.41ms step:187/1555 train_time:6430ms step_avg:34.39ms step:188/1555 train_time:6467ms step_avg:34.40ms step:189/1555 train_time:6498ms step_avg:34.38ms step:190/1555 train_time:6535ms step_avg:34.40ms step:191/1555 train_time:6567ms step_avg:34.38ms step:192/1555 train_time:6604ms step_avg:34.40ms step:193/1555 train_time:6635ms step_avg:34.38ms step:194/1555 train_time:6673ms step_avg:34.40ms step:195/1555 train_time:6703ms step_avg:34.38ms step:196/1555 train_time:6741ms step_avg:34.39ms step:197/1555 train_time:6772ms step_avg:34.38ms step:198/1555 train_time:6809ms step_avg:34.39ms step:199/1555 train_time:6841ms step_avg:34.37ms step:200/1555 train_time:6878ms step_avg:34.39ms step:201/1555 train_time:6909ms step_avg:34.37ms step:202/1555 train_time:6946ms step_avg:34.39ms step:203/1555 train_time:6977ms step_avg:34.37ms step:204/1555 train_time:7014ms step_avg:34.38ms step:205/1555 train_time:7045ms step_avg:34.37ms step:206/1555 train_time:7083ms step_avg:34.38ms step:207/1555 train_time:7114ms step_avg:34.37ms step:208/1555 train_time:7152ms step_avg:34.38ms step:209/1555 train_time:7182ms step_avg:34.37ms step:210/1555 train_time:7220ms step_avg:34.38ms step:211/1555 train_time:7251ms step_avg:34.36ms step:212/1555 train_time:7288ms step_avg:34.38ms step:213/1555 train_time:7319ms step_avg:34.36ms step:214/1555 train_time:7357ms step_avg:34.38ms step:215/1555 train_time:7388ms step_avg:34.36ms step:216/1555 train_time:7425ms step_avg:34.38ms step:217/1555 train_time:7457ms step_avg:34.36ms step:218/1555 train_time:7494ms step_avg:34.38ms step:219/1555 train_time:7525ms step_avg:34.36ms step:220/1555 train_time:7562ms step_avg:34.37ms step:221/1555 train_time:7593ms step_avg:34.36ms step:222/1555 train_time:7631ms step_avg:34.37ms step:223/1555 train_time:7661ms step_avg:34.36ms step:224/1555 train_time:7700ms step_avg:34.37ms step:225/1555 train_time:7731ms step_avg:34.36ms step:226/1555 train_time:7768ms step_avg:34.37ms step:227/1555 train_time:7799ms step_avg:34.36ms step:228/1555 train_time:7837ms step_avg:34.37ms step:229/1555 train_time:7867ms step_avg:34.36ms step:230/1555 train_time:7905ms step_avg:34.37ms step:231/1555 train_time:7935ms step_avg:34.35ms step:232/1555 train_time:7973ms step_avg:34.37ms step:233/1555 train_time:8004ms step_avg:34.35ms step:234/1555 train_time:8042ms step_avg:34.37ms step:235/1555 train_time:8073ms step_avg:34.35ms step:236/1555 train_time:8110ms step_avg:34.36ms step:237/1555 train_time:8142ms step_avg:34.35ms step:238/1555 train_time:8179ms step_avg:34.37ms step:239/1555 train_time:8210ms step_avg:34.35ms step:240/1555 train_time:8248ms step_avg:34.37ms step:241/1555 train_time:8278ms step_avg:34.35ms step:242/1555 train_time:8316ms step_avg:34.36ms step:243/1555 train_time:8348ms step_avg:34.35ms step:244/1555 train_time:8385ms step_avg:34.37ms step:245/1555 train_time:8417ms step_avg:34.35ms step:246/1555 train_time:8454ms step_avg:34.37ms step:247/1555 train_time:8485ms step_avg:34.35ms step:248/1555 train_time:8523ms step_avg:34.37ms step:249/1555 train_time:8553ms step_avg:34.35ms step:250/1555 train_time:8591ms step_avg:34.36ms step:250/1555 val_loss:4.5532 train_time:8641ms step_avg:34.56ms step:251/1555 train_time:8662ms step_avg:34.51ms step:252/1555 train_time:8686ms step_avg:34.47ms step:253/1555 train_time:8707ms step_avg:34.41ms step:254/1555 train_time:8731ms step_avg:34.37ms step:255/1555 train_time:8763ms step_avg:34.37ms step:256/1555 train_time:8802ms step_avg:34.38ms step:257/1555 train_time:8835ms step_avg:34.38ms step:258/1555 train_time:8874ms step_avg:34.40ms step:259/1555 train_time:8907ms step_avg:34.39ms step:260/1555 train_time:8944ms step_avg:34.40ms step:261/1555 train_time:8974ms step_avg:34.38ms step:262/1555 train_time:9012ms step_avg:34.40ms step:263/1555 train_time:9042ms step_avg:34.38ms step:264/1555 train_time:9080ms step_avg:34.39ms step:265/1555 train_time:9110ms step_avg:34.38ms step:266/1555 train_time:9148ms step_avg:34.39ms step:267/1555 train_time:9179ms step_avg:34.38ms step:268/1555 train_time:9217ms step_avg:34.39ms step:269/1555 train_time:9247ms step_avg:34.38ms step:270/1555 train_time:9285ms step_avg:34.39ms step:271/1555 train_time:9315ms step_avg:34.37ms step:272/1555 train_time:9353ms step_avg:34.38ms step:273/1555 train_time:9383ms step_avg:34.37ms step:274/1555 train_time:9420ms step_avg:34.38ms step:275/1555 train_time:9451ms step_avg:34.37ms step:276/1555 train_time:9489ms step_avg:34.38ms step:277/1555 train_time:9520ms step_avg:34.37ms step:278/1555 train_time:9558ms step_avg:34.38ms step:279/1555 train_time:9593ms step_avg:34.39ms step:280/1555 train_time:9625ms step_avg:34.38ms step:281/1555 train_time:9656ms step_avg:34.36ms step:282/1555 train_time:9694ms step_avg:34.38ms step:283/1555 train_time:9725ms step_avg:34.36ms step:284/1555 train_time:9762ms step_avg:34.37ms step:285/1555 train_time:9794ms step_avg:34.36ms step:286/1555 train_time:9832ms step_avg:34.38ms step:287/1555 train_time:9863ms step_avg:34.36ms step:288/1555 train_time:9900ms step_avg:34.38ms step:289/1555 train_time:9931ms step_avg:34.36ms step:290/1555 train_time:9969ms step_avg:34.37ms step:291/1555 train_time:9999ms step_avg:34.36ms step:292/1555 train_time:10037ms step_avg:34.37ms step:293/1555 train_time:10068ms step_avg:34.36ms step:294/1555 train_time:10105ms step_avg:34.37ms step:295/1555 train_time:10136ms step_avg:34.36ms step:296/1555 train_time:10173ms step_avg:34.37ms step:297/1555 train_time:10204ms step_avg:34.36ms step:298/1555 train_time:10242ms step_avg:34.37ms step:299/1555 train_time:10273ms step_avg:34.36ms step:300/1555 train_time:10311ms step_avg:34.37ms step:301/1555 train_time:10341ms step_avg:34.36ms step:302/1555 train_time:10378ms step_avg:34.37ms step:303/1555 train_time:10409ms step_avg:34.35ms step:304/1555 train_time:10447ms step_avg:34.36ms step:305/1555 train_time:10477ms step_avg:34.35ms step:306/1555 train_time:10515ms step_avg:34.36ms step:307/1555 train_time:10546ms step_avg:34.35ms step:308/1555 train_time:10583ms step_avg:34.36ms step:309/1555 train_time:10613ms step_avg:34.35ms step:310/1555 train_time:10651ms step_avg:34.36ms step:311/1555 train_time:10682ms step_avg:34.35ms step:312/1555 train_time:10719ms step_avg:34.35ms step:313/1555 train_time:10750ms step_avg:34.34ms step:314/1555 train_time:10787ms step_avg:34.35ms step:315/1555 train_time:10818ms step_avg:34.34ms step:316/1555 train_time:10856ms step_avg:34.35ms step:317/1555 train_time:10886ms step_avg:34.34ms step:318/1555 train_time:10924ms step_avg:34.35ms step:319/1555 train_time:10955ms step_avg:34.34ms step:320/1555 train_time:10993ms step_avg:34.35ms step:321/1555 train_time:11024ms step_avg:34.34ms step:322/1555 train_time:11062ms step_avg:34.35ms step:323/1555 train_time:11092ms step_avg:34.34ms step:324/1555 train_time:11130ms step_avg:34.35ms step:325/1555 train_time:11161ms step_avg:34.34ms step:326/1555 train_time:11199ms step_avg:34.35ms step:327/1555 train_time:11230ms step_avg:34.34ms step:328/1555 train_time:11267ms step_avg:34.35ms step:329/1555 train_time:11298ms step_avg:34.34ms step:330/1555 train_time:11336ms step_avg:34.35ms step:331/1555 train_time:11366ms step_avg:34.34ms step:332/1555 train_time:11404ms step_avg:34.35ms step:333/1555 train_time:11434ms step_avg:34.34ms step:334/1555 train_time:11472ms step_avg:34.35ms step:335/1555 train_time:11502ms step_avg:34.34ms step:336/1555 train_time:11540ms step_avg:34.35ms step:337/1555 train_time:11571ms step_avg:34.34ms step:338/1555 train_time:11608ms step_avg:34.34ms step:339/1555 train_time:11639ms step_avg:34.33ms step:340/1555 train_time:11677ms step_avg:34.34ms step:341/1555 train_time:11708ms step_avg:34.33ms step:342/1555 train_time:11745ms step_avg:34.34ms step:343/1555 train_time:11776ms step_avg:34.33ms step:344/1555 train_time:11814ms step_avg:34.34ms step:345/1555 train_time:11845ms step_avg:34.33ms step:346/1555 train_time:11883ms step_avg:34.34ms step:347/1555 train_time:11913ms step_avg:34.33ms step:348/1555 train_time:11952ms step_avg:34.34ms step:349/1555 train_time:11983ms step_avg:34.33ms step:350/1555 train_time:12020ms step_avg:34.34ms step:351/1555 train_time:12051ms step_avg:34.33ms step:352/1555 train_time:12089ms step_avg:34.34ms step:353/1555 train_time:12119ms step_avg:34.33ms step:354/1555 train_time:12157ms step_avg:34.34ms step:355/1555 train_time:12188ms step_avg:34.33ms step:356/1555 train_time:12225ms step_avg:34.34ms step:357/1555 train_time:12256ms step_avg:34.33ms step:358/1555 train_time:12294ms step_avg:34.34ms step:359/1555 train_time:12325ms step_avg:34.33ms step:360/1555 train_time:12362ms step_avg:34.34ms step:361/1555 train_time:12393ms step_avg:34.33ms step:362/1555 train_time:12431ms step_avg:34.34ms step:363/1555 train_time:12462ms step_avg:34.33ms step:364/1555 train_time:12500ms step_avg:34.34ms step:365/1555 train_time:12530ms step_avg:34.33ms step:366/1555 train_time:12568ms step_avg:34.34ms step:367/1555 train_time:12598ms step_avg:34.33ms step:368/1555 train_time:12636ms step_avg:34.34ms step:369/1555 train_time:12667ms step_avg:34.33ms step:370/1555 train_time:12704ms step_avg:34.34ms step:371/1555 train_time:12735ms step_avg:34.33ms step:372/1555 train_time:12773ms step_avg:34.34ms step:373/1555 train_time:12804ms step_avg:34.33ms step:374/1555 train_time:12841ms step_avg:34.33ms step:375/1555 train_time:12872ms step_avg:34.33ms step:376/1555 train_time:12909ms step_avg:34.33ms step:377/1555 train_time:12940ms step_avg:34.32ms step:378/1555 train_time:12977ms step_avg:34.33ms step:379/1555 train_time:13008ms step_avg:34.32ms step:380/1555 train_time:13046ms step_avg:34.33ms step:381/1555 train_time:13077ms step_avg:34.32ms step:382/1555 train_time:13115ms step_avg:34.33ms step:383/1555 train_time:13146ms step_avg:34.32ms step:384/1555 train_time:13183ms step_avg:34.33ms step:385/1555 train_time:13214ms step_avg:34.32ms step:386/1555 train_time:13252ms step_avg:34.33ms step:387/1555 train_time:13283ms step_avg:34.32ms step:388/1555 train_time:13321ms step_avg:34.33ms step:389/1555 train_time:13351ms step_avg:34.32ms step:390/1555 train_time:13388ms step_avg:34.33ms step:391/1555 train_time:13419ms step_avg:34.32ms step:392/1555 train_time:13457ms step_avg:34.33ms step:393/1555 train_time:13488ms step_avg:34.32ms step:394/1555 train_time:13525ms step_avg:34.33ms step:395/1555 train_time:13556ms step_avg:34.32ms step:396/1555 train_time:13593ms step_avg:34.33ms step:397/1555 train_time:13624ms step_avg:34.32ms step:398/1555 train_time:13662ms step_avg:34.33ms step:399/1555 train_time:13693ms step_avg:34.32ms step:400/1555 train_time:13731ms step_avg:34.33ms step:401/1555 train_time:13761ms step_avg:34.32ms step:402/1555 train_time:13799ms step_avg:34.33ms step:403/1555 train_time:13829ms step_avg:34.32ms step:404/1555 train_time:13867ms step_avg:34.32ms step:405/1555 train_time:13898ms step_avg:34.32ms step:406/1555 train_time:13935ms step_avg:34.32ms step:407/1555 train_time:13966ms step_avg:34.31ms step:408/1555 train_time:14003ms step_avg:34.32ms step:409/1555 train_time:14034ms step_avg:34.31ms step:410/1555 train_time:14071ms step_avg:34.32ms step:411/1555 train_time:14102ms step_avg:34.31ms step:412/1555 train_time:14139ms step_avg:34.32ms step:413/1555 train_time:14170ms step_avg:34.31ms step:414/1555 train_time:14207ms step_avg:34.32ms step:415/1555 train_time:14238ms step_avg:34.31ms step:416/1555 train_time:14275ms step_avg:34.32ms step:417/1555 train_time:14306ms step_avg:34.31ms step:418/1555 train_time:14343ms step_avg:34.31ms step:419/1555 train_time:14374ms step_avg:34.30ms step:420/1555 train_time:14411ms step_avg:34.31ms step:421/1555 train_time:14442ms step_avg:34.30ms step:422/1555 train_time:14480ms step_avg:34.31ms step:423/1555 train_time:14511ms step_avg:34.30ms step:424/1555 train_time:14548ms step_avg:34.31ms step:425/1555 train_time:14579ms step_avg:34.30ms step:426/1555 train_time:14617ms step_avg:34.31ms step:427/1555 train_time:14648ms step_avg:34.30ms step:428/1555 train_time:14685ms step_avg:34.31ms step:429/1555 train_time:14716ms step_avg:34.30ms step:430/1555 train_time:14753ms step_avg:34.31ms step:431/1555 train_time:14784ms step_avg:34.30ms step:432/1555 train_time:14822ms step_avg:34.31ms step:433/1555 train_time:14852ms step_avg:34.30ms step:434/1555 train_time:14890ms step_avg:34.31ms step:435/1555 train_time:14921ms step_avg:34.30ms step:436/1555 train_time:14958ms step_avg:34.31ms step:437/1555 train_time:14989ms step_avg:34.30ms step:438/1555 train_time:15027ms step_avg:34.31ms step:439/1555 train_time:15058ms step_avg:34.30ms step:440/1555 train_time:15095ms step_avg:34.31ms step:441/1555 train_time:15125ms step_avg:34.30ms step:442/1555 train_time:15163ms step_avg:34.30ms step:443/1555 train_time:15193ms step_avg:34.30ms step:444/1555 train_time:15231ms step_avg:34.30ms step:445/1555 train_time:15262ms step_avg:34.30ms step:446/1555 train_time:15299ms step_avg:34.30ms step:447/1555 train_time:15330ms step_avg:34.29ms step:448/1555 train_time:15367ms step_avg:34.30ms step:449/1555 train_time:15398ms step_avg:34.29ms step:450/1555 train_time:15436ms step_avg:34.30ms step:451/1555 train_time:15467ms step_avg:34.29ms step:452/1555 train_time:15505ms step_avg:34.30ms step:453/1555 train_time:15535ms step_avg:34.29ms step:454/1555 train_time:15574ms step_avg:34.30ms step:455/1555 train_time:15604ms step_avg:34.30ms step:456/1555 train_time:15642ms step_avg:34.30ms step:457/1555 train_time:15672ms step_avg:34.29ms step:458/1555 train_time:15710ms step_avg:34.30ms step:459/1555 train_time:15741ms step_avg:34.29ms step:460/1555 train_time:15779ms step_avg:34.30ms step:461/1555 train_time:15809ms step_avg:34.29ms step:462/1555 train_time:15847ms step_avg:34.30ms step:463/1555 train_time:15878ms step_avg:34.29ms step:464/1555 train_time:15916ms step_avg:34.30ms step:465/1555 train_time:15947ms step_avg:34.29ms step:466/1555 train_time:15984ms step_avg:34.30ms step:467/1555 train_time:16016ms step_avg:34.30ms step:468/1555 train_time:16054ms step_avg:34.30ms step:469/1555 train_time:16085ms step_avg:34.30ms step:470/1555 train_time:16122ms step_avg:34.30ms step:471/1555 train_time:16154ms step_avg:34.30ms step:472/1555 train_time:16192ms step_avg:34.30ms step:473/1555 train_time:16222ms step_avg:34.30ms step:474/1555 train_time:16260ms step_avg:34.30ms step:475/1555 train_time:16291ms step_avg:34.30ms step:476/1555 train_time:16328ms step_avg:34.30ms step:477/1555 train_time:16359ms step_avg:34.29ms step:478/1555 train_time:16396ms step_avg:34.30ms step:479/1555 train_time:16427ms step_avg:34.29ms step:480/1555 train_time:16464ms step_avg:34.30ms step:481/1555 train_time:16495ms step_avg:34.29ms step:482/1555 train_time:16534ms step_avg:34.30ms step:483/1555 train_time:16564ms step_avg:34.29ms step:484/1555 train_time:16601ms step_avg:34.30ms step:485/1555 train_time:16632ms step_avg:34.29ms step:486/1555 train_time:16670ms step_avg:34.30ms step:487/1555 train_time:16701ms step_avg:34.29ms step:488/1555 train_time:16739ms step_avg:34.30ms step:489/1555 train_time:16769ms step_avg:34.29ms step:490/1555 train_time:16807ms step_avg:34.30ms step:491/1555 train_time:16837ms step_avg:34.29ms step:492/1555 train_time:16875ms step_avg:34.30ms step:493/1555 train_time:16905ms step_avg:34.29ms step:494/1555 train_time:16943ms step_avg:34.30ms step:495/1555 train_time:16974ms step_avg:34.29ms step:496/1555 train_time:17012ms step_avg:34.30ms step:497/1555 train_time:17042ms step_avg:34.29ms step:498/1555 train_time:17080ms step_avg:34.30ms step:499/1555 train_time:17111ms step_avg:34.29ms step:500/1555 train_time:17148ms step_avg:34.30ms step:500/1555 val_loss:4.2277 train_time:17197ms step_avg:34.39ms step:501/1555 train_time:17218ms step_avg:34.37ms step:502/1555 train_time:17238ms step_avg:34.34ms step:503/1555 train_time:17257ms step_avg:34.31ms step:504/1555 train_time:17287ms step_avg:34.30ms step:505/1555 train_time:17319ms step_avg:34.30ms step:506/1555 train_time:17364ms step_avg:34.32ms step:507/1555 train_time:17416ms step_avg:34.35ms step:508/1555 train_time:17480ms step_avg:34.41ms step:509/1555 train_time:17538ms step_avg:34.46ms step:510/1555 train_time:17603ms step_avg:34.52ms step:511/1555 train_time:17661ms step_avg:34.56ms step:512/1555 train_time:17724ms step_avg:34.62ms step:513/1555 train_time:17781ms step_avg:34.66ms step:514/1555 train_time:17845ms step_avg:34.72ms step:515/1555 train_time:17902ms step_avg:34.76ms step:516/1555 train_time:17966ms step_avg:34.82ms step:517/1555 train_time:18023ms step_avg:34.86ms step:518/1555 train_time:18086ms step_avg:34.92ms step:519/1555 train_time:18145ms step_avg:34.96ms step:520/1555 train_time:18211ms step_avg:35.02ms step:521/1555 train_time:18270ms step_avg:35.07ms step:522/1555 train_time:18337ms step_avg:35.13ms step:523/1555 train_time:18394ms step_avg:35.17ms step:524/1555 train_time:18458ms step_avg:35.23ms step:525/1555 train_time:18516ms step_avg:35.27ms step:526/1555 train_time:18580ms step_avg:35.32ms step:527/1555 train_time:18638ms step_avg:35.37ms step:528/1555 train_time:18702ms step_avg:35.42ms step:529/1555 train_time:18760ms step_avg:35.46ms step:530/1555 train_time:18823ms step_avg:35.52ms step:531/1555 train_time:18881ms step_avg:35.56ms step:532/1555 train_time:18945ms step_avg:35.61ms step:533/1555 train_time:19001ms step_avg:35.65ms step:534/1555 train_time:19065ms step_avg:35.70ms step:535/1555 train_time:19123ms step_avg:35.74ms step:536/1555 train_time:19188ms step_avg:35.80ms step:537/1555 train_time:19247ms step_avg:35.84ms step:538/1555 train_time:19312ms step_avg:35.90ms step:539/1555 train_time:19372ms step_avg:35.94ms step:540/1555 train_time:19435ms step_avg:35.99ms step:541/1555 train_time:19493ms step_avg:36.03ms step:542/1555 train_time:19556ms step_avg:36.08ms step:543/1555 train_time:19615ms step_avg:36.12ms step:544/1555 train_time:19679ms step_avg:36.17ms step:545/1555 train_time:19737ms step_avg:36.21ms step:546/1555 train_time:19801ms step_avg:36.27ms step:547/1555 train_time:19858ms step_avg:36.30ms step:548/1555 train_time:19922ms step_avg:36.35ms step:549/1555 train_time:19979ms step_avg:36.39ms step:550/1555 train_time:20043ms step_avg:36.44ms step:551/1555 train_time:20100ms step_avg:36.48ms step:552/1555 train_time:20165ms step_avg:36.53ms step:553/1555 train_time:20223ms step_avg:36.57ms step:554/1555 train_time:20288ms step_avg:36.62ms step:555/1555 train_time:20346ms step_avg:36.66ms step:556/1555 train_time:20412ms step_avg:36.71ms step:557/1555 train_time:20471ms step_avg:36.75ms step:558/1555 train_time:20534ms step_avg:36.80ms step:559/1555 train_time:20591ms step_avg:36.84ms step:560/1555 train_time:20655ms step_avg:36.88ms step:561/1555 train_time:20712ms step_avg:36.92ms step:562/1555 train_time:20776ms step_avg:36.97ms step:563/1555 train_time:20834ms step_avg:37.01ms step:564/1555 train_time:20898ms step_avg:37.05ms step:565/1555 train_time:20956ms step_avg:37.09ms step:566/1555 train_time:21020ms step_avg:37.14ms step:567/1555 train_time:21078ms step_avg:37.17ms step:568/1555 train_time:21142ms step_avg:37.22ms step:569/1555 train_time:21200ms step_avg:37.26ms step:570/1555 train_time:21265ms step_avg:37.31ms step:571/1555 train_time:21322ms step_avg:37.34ms step:572/1555 train_time:21387ms step_avg:37.39ms step:573/1555 train_time:21445ms step_avg:37.43ms step:574/1555 train_time:21510ms step_avg:37.47ms step:575/1555 train_time:21568ms step_avg:37.51ms step:576/1555 train_time:21632ms step_avg:37.56ms step:577/1555 train_time:21689ms step_avg:37.59ms step:578/1555 train_time:21753ms step_avg:37.63ms step:579/1555 train_time:21810ms step_avg:37.67ms step:580/1555 train_time:21876ms step_avg:37.72ms step:581/1555 train_time:21932ms step_avg:37.75ms step:582/1555 train_time:21997ms step_avg:37.79ms step:583/1555 train_time:22055ms step_avg:37.83ms step:584/1555 train_time:22119ms step_avg:37.87ms step:585/1555 train_time:22176ms step_avg:37.91ms step:586/1555 train_time:22241ms step_avg:37.95ms step:587/1555 train_time:22299ms step_avg:37.99ms step:588/1555 train_time:22365ms step_avg:38.03ms step:589/1555 train_time:22422ms step_avg:38.07ms step:590/1555 train_time:22486ms step_avg:38.11ms step:591/1555 train_time:22544ms step_avg:38.15ms step:592/1555 train_time:22609ms step_avg:38.19ms step:593/1555 train_time:22666ms step_avg:38.22ms step:594/1555 train_time:22730ms step_avg:38.27ms step:595/1555 train_time:22788ms step_avg:38.30ms step:596/1555 train_time:22852ms step_avg:38.34ms step:597/1555 train_time:22910ms step_avg:38.37ms step:598/1555 train_time:22974ms step_avg:38.42ms step:599/1555 train_time:23031ms step_avg:38.45ms step:600/1555 train_time:23094ms step_avg:38.49ms step:601/1555 train_time:23153ms step_avg:38.52ms step:602/1555 train_time:23216ms step_avg:38.57ms step:603/1555 train_time:23275ms step_avg:38.60ms step:604/1555 train_time:23340ms step_avg:38.64ms step:605/1555 train_time:23397ms step_avg:38.67ms step:606/1555 train_time:23462ms step_avg:38.72ms step:607/1555 train_time:23519ms step_avg:38.75ms step:608/1555 train_time:23584ms step_avg:38.79ms step:609/1555 train_time:23642ms step_avg:38.82ms step:610/1555 train_time:23707ms step_avg:38.86ms step:611/1555 train_time:23764ms step_avg:38.89ms step:612/1555 train_time:23829ms step_avg:38.94ms step:613/1555 train_time:23887ms step_avg:38.97ms step:614/1555 train_time:23951ms step_avg:39.01ms step:615/1555 train_time:24008ms step_avg:39.04ms step:616/1555 train_time:24074ms step_avg:39.08ms step:617/1555 train_time:24131ms step_avg:39.11ms step:618/1555 train_time:24195ms step_avg:39.15ms step:619/1555 train_time:24253ms step_avg:39.18ms step:620/1555 train_time:24317ms step_avg:39.22ms step:621/1555 train_time:24374ms step_avg:39.25ms step:622/1555 train_time:24438ms step_avg:39.29ms step:623/1555 train_time:24496ms step_avg:39.32ms step:624/1555 train_time:24561ms step_avg:39.36ms step:625/1555 train_time:24619ms step_avg:39.39ms step:626/1555 train_time:24684ms step_avg:39.43ms step:627/1555 train_time:24742ms step_avg:39.46ms step:628/1555 train_time:24806ms step_avg:39.50ms step:629/1555 train_time:24864ms step_avg:39.53ms step:630/1555 train_time:24929ms step_avg:39.57ms step:631/1555 train_time:24987ms step_avg:39.60ms step:632/1555 train_time:25051ms step_avg:39.64ms step:633/1555 train_time:25109ms step_avg:39.67ms step:634/1555 train_time:25173ms step_avg:39.70ms step:635/1555 train_time:25230ms step_avg:39.73ms step:636/1555 train_time:25293ms step_avg:39.77ms step:637/1555 train_time:25351ms step_avg:39.80ms step:638/1555 train_time:25415ms step_avg:39.84ms step:639/1555 train_time:25472ms step_avg:39.86ms step:640/1555 train_time:25537ms step_avg:39.90ms step:641/1555 train_time:25596ms step_avg:39.93ms step:642/1555 train_time:25661ms step_avg:39.97ms step:643/1555 train_time:25719ms step_avg:40.00ms step:644/1555 train_time:25783ms step_avg:40.04ms step:645/1555 train_time:25842ms step_avg:40.06ms step:646/1555 train_time:25907ms step_avg:40.10ms step:647/1555 train_time:25965ms step_avg:40.13ms step:648/1555 train_time:26029ms step_avg:40.17ms step:649/1555 train_time:26087ms step_avg:40.20ms step:650/1555 train_time:26151ms step_avg:40.23ms step:651/1555 train_time:26210ms step_avg:40.26ms step:652/1555 train_time:26274ms step_avg:40.30ms step:653/1555 train_time:26331ms step_avg:40.32ms step:654/1555 train_time:26394ms step_avg:40.36ms step:655/1555 train_time:26452ms step_avg:40.39ms step:656/1555 train_time:26516ms step_avg:40.42ms step:657/1555 train_time:26572ms step_avg:40.45ms step:658/1555 train_time:26637ms step_avg:40.48ms step:659/1555 train_time:26694ms step_avg:40.51ms step:660/1555 train_time:26760ms step_avg:40.54ms step:661/1555 train_time:26817ms step_avg:40.57ms step:662/1555 train_time:26882ms step_avg:40.61ms step:663/1555 train_time:26940ms step_avg:40.63ms step:664/1555 train_time:27005ms step_avg:40.67ms step:665/1555 train_time:27063ms step_avg:40.70ms step:666/1555 train_time:27128ms step_avg:40.73ms step:667/1555 train_time:27186ms step_avg:40.76ms step:668/1555 train_time:27250ms step_avg:40.79ms step:669/1555 train_time:27308ms step_avg:40.82ms step:670/1555 train_time:27372ms step_avg:40.85ms step:671/1555 train_time:27431ms step_avg:40.88ms step:672/1555 train_time:27494ms step_avg:40.91ms step:673/1555 train_time:27551ms step_avg:40.94ms step:674/1555 train_time:27615ms step_avg:40.97ms step:675/1555 train_time:27674ms step_avg:41.00ms step:676/1555 train_time:27738ms step_avg:41.03ms step:677/1555 train_time:27796ms step_avg:41.06ms step:678/1555 train_time:27861ms step_avg:41.09ms step:679/1555 train_time:27918ms step_avg:41.12ms step:680/1555 train_time:27983ms step_avg:41.15ms step:681/1555 train_time:28042ms step_avg:41.18ms step:682/1555 train_time:28106ms step_avg:41.21ms step:683/1555 train_time:28163ms step_avg:41.23ms step:684/1555 train_time:28228ms step_avg:41.27ms step:685/1555 train_time:28286ms step_avg:41.29ms step:686/1555 train_time:28351ms step_avg:41.33ms step:687/1555 train_time:28408ms step_avg:41.35ms step:688/1555 train_time:28473ms step_avg:41.38ms step:689/1555 train_time:28530ms step_avg:41.41ms step:690/1555 train_time:28594ms step_avg:41.44ms step:691/1555 train_time:28652ms step_avg:41.46ms step:692/1555 train_time:28716ms step_avg:41.50ms step:693/1555 train_time:28774ms step_avg:41.52ms step:694/1555 train_time:28838ms step_avg:41.55ms step:695/1555 train_time:28895ms step_avg:41.58ms step:696/1555 train_time:28960ms step_avg:41.61ms step:697/1555 train_time:29018ms step_avg:41.63ms step:698/1555 train_time:29083ms step_avg:41.67ms step:699/1555 train_time:29141ms step_avg:41.69ms step:700/1555 train_time:29205ms step_avg:41.72ms step:701/1555 train_time:29263ms step_avg:41.75ms step:702/1555 train_time:29327ms step_avg:41.78ms step:703/1555 train_time:29385ms step_avg:41.80ms step:704/1555 train_time:29450ms step_avg:41.83ms step:705/1555 train_time:29507ms step_avg:41.85ms step:706/1555 train_time:29571ms step_avg:41.89ms step:707/1555 train_time:29629ms step_avg:41.91ms step:708/1555 train_time:29694ms step_avg:41.94ms step:709/1555 train_time:29752ms step_avg:41.96ms step:710/1555 train_time:29816ms step_avg:41.99ms step:711/1555 train_time:29873ms step_avg:42.02ms step:712/1555 train_time:29937ms step_avg:42.05ms step:713/1555 train_time:29994ms step_avg:42.07ms step:714/1555 train_time:30058ms step_avg:42.10ms step:715/1555 train_time:30116ms step_avg:42.12ms step:716/1555 train_time:30182ms step_avg:42.15ms step:717/1555 train_time:30241ms step_avg:42.18ms step:718/1555 train_time:30305ms step_avg:42.21ms step:719/1555 train_time:30364ms step_avg:42.23ms step:720/1555 train_time:30428ms step_avg:42.26ms step:721/1555 train_time:30485ms step_avg:42.28ms step:722/1555 train_time:30549ms step_avg:42.31ms step:723/1555 train_time:30607ms step_avg:42.33ms step:724/1555 train_time:30671ms step_avg:42.36ms step:725/1555 train_time:30729ms step_avg:42.38ms step:726/1555 train_time:30792ms step_avg:42.41ms step:727/1555 train_time:30850ms step_avg:42.43ms step:728/1555 train_time:30913ms step_avg:42.46ms step:729/1555 train_time:30971ms step_avg:42.48ms step:730/1555 train_time:31035ms step_avg:42.51ms step:731/1555 train_time:31093ms step_avg:42.54ms step:732/1555 train_time:31158ms step_avg:42.57ms step:733/1555 train_time:31216ms step_avg:42.59ms step:734/1555 train_time:31281ms step_avg:42.62ms step:735/1555 train_time:31339ms step_avg:42.64ms step:736/1555 train_time:31404ms step_avg:42.67ms step:737/1555 train_time:31462ms step_avg:42.69ms step:738/1555 train_time:31526ms step_avg:42.72ms step:739/1555 train_time:31584ms step_avg:42.74ms step:740/1555 train_time:31648ms step_avg:42.77ms step:741/1555 train_time:31705ms step_avg:42.79ms step:742/1555 train_time:31770ms step_avg:42.82ms step:743/1555 train_time:31829ms step_avg:42.84ms step:744/1555 train_time:31892ms step_avg:42.87ms step:745/1555 train_time:31951ms step_avg:42.89ms step:746/1555 train_time:32014ms step_avg:42.91ms step:747/1555 train_time:32072ms step_avg:42.93ms step:748/1555 train_time:32137ms step_avg:42.96ms step:749/1555 train_time:32195ms step_avg:42.98ms step:750/1555 train_time:32258ms step_avg:43.01ms step:750/1555 val_loss:3.8724 train_time:32342ms step_avg:43.12ms step:751/1555 train_time:32363ms step_avg:43.09ms step:752/1555 train_time:32392ms step_avg:43.07ms step:753/1555 train_time:32443ms step_avg:43.09ms step:754/1555 train_time:32512ms step_avg:43.12ms step:755/1555 train_time:32569ms step_avg:43.14ms step:756/1555 train_time:32633ms step_avg:43.17ms step:757/1555 train_time:32691ms step_avg:43.19ms step:758/1555 train_time:32754ms step_avg:43.21ms step:759/1555 train_time:32811ms step_avg:43.23ms step:760/1555 train_time:32874ms step_avg:43.25ms step:761/1555 train_time:32930ms step_avg:43.27ms step:762/1555 train_time:32994ms step_avg:43.30ms step:763/1555 train_time:33050ms step_avg:43.32ms step:764/1555 train_time:33113ms step_avg:43.34ms step:765/1555 train_time:33170ms step_avg:43.36ms step:766/1555 train_time:33233ms step_avg:43.39ms step:767/1555 train_time:33291ms step_avg:43.40ms step:768/1555 train_time:33356ms step_avg:43.43ms step:769/1555 train_time:33416ms step_avg:43.45ms step:770/1555 train_time:33482ms step_avg:43.48ms step:771/1555 train_time:33540ms step_avg:43.50ms step:772/1555 train_time:33605ms step_avg:43.53ms step:773/1555 train_time:33664ms step_avg:43.55ms step:774/1555 train_time:33728ms step_avg:43.58ms step:775/1555 train_time:33785ms step_avg:43.59ms step:776/1555 train_time:33849ms step_avg:43.62ms step:777/1555 train_time:33906ms step_avg:43.64ms step:778/1555 train_time:33970ms step_avg:43.66ms step:779/1555 train_time:34027ms step_avg:43.68ms step:780/1555 train_time:34090ms step_avg:43.71ms step:781/1555 train_time:34147ms step_avg:43.72ms step:782/1555 train_time:34211ms step_avg:43.75ms step:783/1555 train_time:34269ms step_avg:43.77ms step:784/1555 train_time:34333ms step_avg:43.79ms step:785/1555 train_time:34393ms step_avg:43.81ms step:786/1555 train_time:34457ms step_avg:43.84ms step:787/1555 train_time:34515ms step_avg:43.86ms step:788/1555 train_time:34580ms step_avg:43.88ms step:789/1555 train_time:34637ms step_avg:43.90ms step:790/1555 train_time:34702ms step_avg:43.93ms step:791/1555 train_time:34759ms step_avg:43.94ms step:792/1555 train_time:34824ms step_avg:43.97ms step:793/1555 train_time:34882ms step_avg:43.99ms step:794/1555 train_time:34945ms step_avg:44.01ms step:795/1555 train_time:35003ms step_avg:44.03ms step:796/1555 train_time:35066ms step_avg:44.05ms step:797/1555 train_time:35124ms step_avg:44.07ms step:798/1555 train_time:35189ms step_avg:44.10ms step:799/1555 train_time:35247ms step_avg:44.11ms step:800/1555 train_time:35311ms step_avg:44.14ms step:801/1555 train_time:35371ms step_avg:44.16ms step:802/1555 train_time:35434ms step_avg:44.18ms step:803/1555 train_time:35494ms step_avg:44.20ms step:804/1555 train_time:35557ms step_avg:44.22ms step:805/1555 train_time:35614ms step_avg:44.24ms step:806/1555 train_time:35678ms step_avg:44.27ms step:807/1555 train_time:35735ms step_avg:44.28ms step:808/1555 train_time:35800ms step_avg:44.31ms step:809/1555 train_time:35857ms step_avg:44.32ms step:810/1555 train_time:35921ms step_avg:44.35ms step:811/1555 train_time:35978ms step_avg:44.36ms step:812/1555 train_time:36043ms step_avg:44.39ms step:813/1555 train_time:36100ms step_avg:44.40ms step:814/1555 train_time:36166ms step_avg:44.43ms step:815/1555 train_time:36224ms step_avg:44.45ms step:816/1555 train_time:36289ms step_avg:44.47ms step:817/1555 train_time:36347ms step_avg:44.49ms step:818/1555 train_time:36411ms step_avg:44.51ms step:819/1555 train_time:36469ms step_avg:44.53ms step:820/1555 train_time:36534ms step_avg:44.55ms step:821/1555 train_time:36593ms step_avg:44.57ms step:822/1555 train_time:36656ms step_avg:44.59ms step:823/1555 train_time:36713ms step_avg:44.61ms step:824/1555 train_time:36776ms step_avg:44.63ms step:825/1555 train_time:36834ms step_avg:44.65ms step:826/1555 train_time:36898ms step_avg:44.67ms step:827/1555 train_time:36955ms step_avg:44.69ms step:828/1555 train_time:37020ms step_avg:44.71ms step:829/1555 train_time:37077ms step_avg:44.73ms step:830/1555 train_time:37142ms step_avg:44.75ms step:831/1555 train_time:37200ms step_avg:44.77ms step:832/1555 train_time:37264ms step_avg:44.79ms step:833/1555 train_time:37323ms step_avg:44.81ms step:834/1555 train_time:37388ms step_avg:44.83ms step:835/1555 train_time:37445ms step_avg:44.84ms step:836/1555 train_time:37511ms step_avg:44.87ms step:837/1555 train_time:37569ms step_avg:44.89ms step:838/1555 train_time:37634ms step_avg:44.91ms step:839/1555 train_time:37691ms step_avg:44.92ms step:840/1555 train_time:37754ms step_avg:44.95ms step:841/1555 train_time:37812ms step_avg:44.96ms step:842/1555 train_time:37875ms step_avg:44.98ms step:843/1555 train_time:37933ms step_avg:45.00ms step:844/1555 train_time:37997ms step_avg:45.02ms step:845/1555 train_time:38054ms step_avg:45.03ms step:846/1555 train_time:38119ms step_avg:45.06ms step:847/1555 train_time:38175ms step_avg:45.07ms step:848/1555 train_time:38240ms step_avg:45.09ms step:849/1555 train_time:38298ms step_avg:45.11ms step:850/1555 train_time:38363ms step_avg:45.13ms step:851/1555 train_time:38421ms step_avg:45.15ms step:852/1555 train_time:38486ms step_avg:45.17ms step:853/1555 train_time:38544ms step_avg:45.19ms step:854/1555 train_time:38609ms step_avg:45.21ms step:855/1555 train_time:38667ms step_avg:45.22ms step:856/1555 train_time:38731ms step_avg:45.25ms step:857/1555 train_time:38788ms step_avg:45.26ms step:858/1555 train_time:38852ms step_avg:45.28ms step:859/1555 train_time:38911ms step_avg:45.30ms step:860/1555 train_time:38974ms step_avg:45.32ms step:861/1555 train_time:39032ms step_avg:45.33ms step:862/1555 train_time:39096ms step_avg:45.35ms step:863/1555 train_time:39153ms step_avg:45.37ms step:864/1555 train_time:39217ms step_avg:45.39ms step:865/1555 train_time:39274ms step_avg:45.40ms step:866/1555 train_time:39338ms step_avg:45.42ms step:867/1555 train_time:39396ms step_avg:45.44ms step:868/1555 train_time:39460ms step_avg:45.46ms step:869/1555 train_time:39518ms step_avg:45.48ms step:870/1555 train_time:39583ms step_avg:45.50ms step:871/1555 train_time:39641ms step_avg:45.51ms step:872/1555 train_time:39706ms step_avg:45.53ms step:873/1555 train_time:39764ms step_avg:45.55ms step:874/1555 train_time:39828ms step_avg:45.57ms step:875/1555 train_time:39886ms step_avg:45.58ms step:876/1555 train_time:39951ms step_avg:45.61ms step:877/1555 train_time:40008ms step_avg:45.62ms step:878/1555 train_time:40073ms step_avg:45.64ms step:879/1555 train_time:40131ms step_avg:45.65ms step:880/1555 train_time:40195ms step_avg:45.68ms step:881/1555 train_time:40252ms step_avg:45.69ms step:882/1555 train_time:40316ms step_avg:45.71ms step:883/1555 train_time:40374ms step_avg:45.72ms step:884/1555 train_time:40438ms step_avg:45.74ms step:885/1555 train_time:40496ms step_avg:45.76ms step:886/1555 train_time:40560ms step_avg:45.78ms step:887/1555 train_time:40618ms step_avg:45.79ms step:888/1555 train_time:40683ms step_avg:45.81ms step:889/1555 train_time:40740ms step_avg:45.83ms step:890/1555 train_time:40805ms step_avg:45.85ms step:891/1555 train_time:40863ms step_avg:45.86ms step:892/1555 train_time:40927ms step_avg:45.88ms step:893/1555 train_time:40985ms step_avg:45.90ms step:894/1555 train_time:41050ms step_avg:45.92ms step:895/1555 train_time:41107ms step_avg:45.93ms step:896/1555 train_time:41172ms step_avg:45.95ms step:897/1555 train_time:41231ms step_avg:45.97ms step:898/1555 train_time:41295ms step_avg:45.99ms step:899/1555 train_time:41353ms step_avg:46.00ms step:900/1555 train_time:41416ms step_avg:46.02ms step:901/1555 train_time:41475ms step_avg:46.03ms step:902/1555 train_time:41537ms step_avg:46.05ms step:903/1555 train_time:41595ms step_avg:46.06ms step:904/1555 train_time:41659ms step_avg:46.08ms step:905/1555 train_time:41716ms step_avg:46.10ms step:906/1555 train_time:41781ms step_avg:46.12ms step:907/1555 train_time:41839ms step_avg:46.13ms step:908/1555 train_time:41903ms step_avg:46.15ms step:909/1555 train_time:41961ms step_avg:46.16ms step:910/1555 train_time:42026ms step_avg:46.18ms step:911/1555 train_time:42084ms step_avg:46.20ms step:912/1555 train_time:42148ms step_avg:46.22ms step:913/1555 train_time:42207ms step_avg:46.23ms step:914/1555 train_time:42271ms step_avg:46.25ms step:915/1555 train_time:42329ms step_avg:46.26ms step:916/1555 train_time:42394ms step_avg:46.28ms step:917/1555 train_time:42451ms step_avg:46.29ms step:918/1555 train_time:42514ms step_avg:46.31ms step:919/1555 train_time:42572ms step_avg:46.32ms step:920/1555 train_time:42636ms step_avg:46.34ms step:921/1555 train_time:42694ms step_avg:46.36ms step:922/1555 train_time:42758ms step_avg:46.37ms step:923/1555 train_time:42816ms step_avg:46.39ms step:924/1555 train_time:42880ms step_avg:46.41ms step:925/1555 train_time:42937ms step_avg:46.42ms step:926/1555 train_time:43002ms step_avg:46.44ms step:927/1555 train_time:43060ms step_avg:46.45ms step:928/1555 train_time:43125ms step_avg:46.47ms step:929/1555 train_time:43183ms step_avg:46.48ms step:930/1555 train_time:43248ms step_avg:46.50ms step:931/1555 train_time:43306ms step_avg:46.52ms step:932/1555 train_time:43370ms step_avg:46.53ms step:933/1555 train_time:43429ms step_avg:46.55ms step:934/1555 train_time:43494ms step_avg:46.57ms step:935/1555 train_time:43551ms step_avg:46.58ms step:936/1555 train_time:43614ms step_avg:46.60ms step:937/1555 train_time:43672ms step_avg:46.61ms step:938/1555 train_time:43736ms step_avg:46.63ms step:939/1555 train_time:43794ms step_avg:46.64ms step:940/1555 train_time:43857ms step_avg:46.66ms step:941/1555 train_time:43916ms step_avg:46.67ms step:942/1555 train_time:43980ms step_avg:46.69ms step:943/1555 train_time:44038ms step_avg:46.70ms step:944/1555 train_time:44103ms step_avg:46.72ms step:945/1555 train_time:44161ms step_avg:46.73ms step:946/1555 train_time:44225ms step_avg:46.75ms step:947/1555 train_time:44284ms step_avg:46.76ms step:948/1555 train_time:44349ms step_avg:46.78ms step:949/1555 train_time:44407ms step_avg:46.79ms step:950/1555 train_time:44470ms step_avg:46.81ms step:951/1555 train_time:44528ms step_avg:46.82ms step:952/1555 train_time:44592ms step_avg:46.84ms step:953/1555 train_time:44650ms step_avg:46.85ms step:954/1555 train_time:44714ms step_avg:46.87ms step:955/1555 train_time:44771ms step_avg:46.88ms step:956/1555 train_time:44835ms step_avg:46.90ms step:957/1555 train_time:44894ms step_avg:46.91ms step:958/1555 train_time:44957ms step_avg:46.93ms step:959/1555 train_time:45015ms step_avg:46.94ms step:960/1555 train_time:45080ms step_avg:46.96ms step:961/1555 train_time:45138ms step_avg:46.97ms step:962/1555 train_time:45203ms step_avg:46.99ms step:963/1555 train_time:45261ms step_avg:47.00ms step:964/1555 train_time:45326ms step_avg:47.02ms step:965/1555 train_time:45384ms step_avg:47.03ms step:966/1555 train_time:45448ms step_avg:47.05ms step:967/1555 train_time:45506ms step_avg:47.06ms step:968/1555 train_time:45571ms step_avg:47.08ms step:969/1555 train_time:45629ms step_avg:47.09ms step:970/1555 train_time:45694ms step_avg:47.11ms step:971/1555 train_time:45751ms step_avg:47.12ms step:972/1555 train_time:45814ms step_avg:47.13ms step:973/1555 train_time:45873ms step_avg:47.15ms step:974/1555 train_time:45937ms step_avg:47.16ms step:975/1555 train_time:45994ms step_avg:47.17ms step:976/1555 train_time:46058ms step_avg:47.19ms step:977/1555 train_time:46116ms step_avg:47.20ms step:978/1555 train_time:46180ms step_avg:47.22ms step:979/1555 train_time:46237ms step_avg:47.23ms step:980/1555 train_time:46303ms step_avg:47.25ms step:981/1555 train_time:46360ms step_avg:47.26ms step:982/1555 train_time:46424ms step_avg:47.28ms step:983/1555 train_time:46483ms step_avg:47.29ms step:984/1555 train_time:46548ms step_avg:47.30ms step:985/1555 train_time:46606ms step_avg:47.32ms step:986/1555 train_time:46670ms step_avg:47.33ms step:987/1555 train_time:46728ms step_avg:47.34ms step:988/1555 train_time:46792ms step_avg:47.36ms step:989/1555 train_time:46850ms step_avg:47.37ms step:990/1555 train_time:46914ms step_avg:47.39ms step:991/1555 train_time:46973ms step_avg:47.40ms step:992/1555 train_time:47036ms step_avg:47.42ms step:993/1555 train_time:47094ms step_avg:47.43ms step:994/1555 train_time:47157ms step_avg:47.44ms step:995/1555 train_time:47215ms step_avg:47.45ms step:996/1555 train_time:47279ms step_avg:47.47ms step:997/1555 train_time:47336ms step_avg:47.48ms step:998/1555 train_time:47402ms step_avg:47.50ms step:999/1555 train_time:47460ms step_avg:47.51ms step:1000/1555 train_time:47525ms step_avg:47.52ms step:1000/1555 val_loss:3.5699 train_time:47608ms step_avg:47.61ms step:1001/1555 train_time:47628ms step_avg:47.58ms step:1002/1555 train_time:47650ms step_avg:47.55ms step:1003/1555 train_time:47705ms step_avg:47.56ms step:1004/1555 train_time:47775ms step_avg:47.58ms step:1005/1555 train_time:47834ms step_avg:47.60ms step:1006/1555 train_time:47899ms step_avg:47.61ms step:1007/1555 train_time:47956ms step_avg:47.62ms step:1008/1555 train_time:48020ms step_avg:47.64ms step:1009/1555 train_time:48077ms step_avg:47.65ms step:1010/1555 train_time:48141ms step_avg:47.66ms step:1011/1555 train_time:48201ms step_avg:47.68ms step:1012/1555 train_time:48285ms step_avg:47.71ms step:1013/1555 train_time:48368ms step_avg:47.75ms step:1014/1555 train_time:48458ms step_avg:47.79ms step:1015/1555 train_time:48541ms step_avg:47.82ms step:1016/1555 train_time:48632ms step_avg:47.87ms step:1017/1555 train_time:48718ms step_avg:47.90ms step:1018/1555 train_time:48811ms step_avg:47.95ms step:1019/1555 train_time:48898ms step_avg:47.99ms step:1020/1555 train_time:48988ms step_avg:48.03ms step:1021/1555 train_time:49072ms step_avg:48.06ms step:1022/1555 train_time:49162ms step_avg:48.10ms step:1023/1555 train_time:49244ms step_avg:48.14ms step:1024/1555 train_time:49334ms step_avg:48.18ms step:1025/1555 train_time:49416ms step_avg:48.21ms step:1026/1555 train_time:49505ms step_avg:48.25ms step:1027/1555 train_time:49590ms step_avg:48.29ms step:1028/1555 train_time:49681ms step_avg:48.33ms step:1029/1555 train_time:49767ms step_avg:48.36ms step:1030/1555 train_time:49860ms step_avg:48.41ms step:1031/1555 train_time:49944ms step_avg:48.44ms step:1032/1555 train_time:50034ms step_avg:48.48ms step:1033/1555 train_time:50118ms step_avg:48.52ms step:1034/1555 train_time:50207ms step_avg:48.56ms step:1035/1555 train_time:50291ms step_avg:48.59ms step:1036/1555 train_time:50380ms step_avg:48.63ms step:1037/1555 train_time:50463ms step_avg:48.66ms step:1038/1555 train_time:50553ms step_avg:48.70ms step:1039/1555 train_time:50638ms step_avg:48.74ms step:1040/1555 train_time:50727ms step_avg:48.78ms step:1041/1555 train_time:50812ms step_avg:48.81ms step:1042/1555 train_time:50904ms step_avg:48.85ms step:1043/1555 train_time:50987ms step_avg:48.89ms step:1044/1555 train_time:51079ms step_avg:48.93ms step:1045/1555 train_time:51162ms step_avg:48.96ms step:1046/1555 train_time:51252ms step_avg:49.00ms step:1047/1555 train_time:51336ms step_avg:49.03ms step:1048/1555 train_time:51424ms step_avg:49.07ms step:1049/1555 train_time:51508ms step_avg:49.10ms step:1050/1555 train_time:51599ms step_avg:49.14ms step:1051/1555 train_time:51682ms step_avg:49.17ms step:1052/1555 train_time:51773ms step_avg:49.21ms step:1053/1555 train_time:51858ms step_avg:49.25ms step:1054/1555 train_time:51947ms step_avg:49.29ms step:1055/1555 train_time:52032ms step_avg:49.32ms step:1056/1555 train_time:52122ms step_avg:49.36ms step:1057/1555 train_time:52205ms step_avg:49.39ms step:1058/1555 train_time:52296ms step_avg:49.43ms step:1059/1555 train_time:52379ms step_avg:49.46ms step:1060/1555 train_time:52469ms step_avg:49.50ms step:1061/1555 train_time:52554ms step_avg:49.53ms step:1062/1555 train_time:52643ms step_avg:49.57ms step:1063/1555 train_time:52728ms step_avg:49.60ms step:1064/1555 train_time:52818ms step_avg:49.64ms step:1065/1555 train_time:52902ms step_avg:49.67ms step:1066/1555 train_time:52993ms step_avg:49.71ms step:1067/1555 train_time:53077ms step_avg:49.74ms step:1068/1555 train_time:53168ms step_avg:49.78ms step:1069/1555 train_time:53252ms step_avg:49.81ms step:1070/1555 train_time:53342ms step_avg:49.85ms step:1071/1555 train_time:53425ms step_avg:49.88ms step:1072/1555 train_time:53516ms step_avg:49.92ms step:1073/1555 train_time:53599ms step_avg:49.95ms step:1074/1555 train_time:53689ms step_avg:49.99ms step:1075/1555 train_time:53774ms step_avg:50.02ms step:1076/1555 train_time:53864ms step_avg:50.06ms step:1077/1555 train_time:53948ms step_avg:50.09ms step:1078/1555 train_time:54040ms step_avg:50.13ms step:1079/1555 train_time:54122ms step_avg:50.16ms step:1080/1555 train_time:54212ms step_avg:50.20ms step:1081/1555 train_time:54298ms step_avg:50.23ms step:1082/1555 train_time:54386ms step_avg:50.26ms step:1083/1555 train_time:54472ms step_avg:50.30ms step:1084/1555 train_time:54562ms step_avg:50.33ms step:1085/1555 train_time:54646ms step_avg:50.36ms step:1086/1555 train_time:54737ms step_avg:50.40ms step:1087/1555 train_time:54820ms step_avg:50.43ms step:1088/1555 train_time:54910ms step_avg:50.47ms step:1089/1555 train_time:54994ms step_avg:50.50ms step:1090/1555 train_time:55083ms step_avg:50.54ms step:1091/1555 train_time:55168ms step_avg:50.57ms step:1092/1555 train_time:55258ms step_avg:50.60ms step:1093/1555 train_time:55341ms step_avg:50.63ms step:1094/1555 train_time:55431ms step_avg:50.67ms step:1095/1555 train_time:55515ms step_avg:50.70ms step:1096/1555 train_time:55604ms step_avg:50.73ms step:1097/1555 train_time:55688ms step_avg:50.76ms step:1098/1555 train_time:55778ms step_avg:50.80ms step:1099/1555 train_time:55863ms step_avg:50.83ms step:1100/1555 train_time:55953ms step_avg:50.87ms step:1101/1555 train_time:56038ms step_avg:50.90ms step:1102/1555 train_time:56127ms step_avg:50.93ms step:1103/1555 train_time:56212ms step_avg:50.96ms step:1104/1555 train_time:56301ms step_avg:51.00ms step:1105/1555 train_time:56385ms step_avg:51.03ms step:1106/1555 train_time:56475ms step_avg:51.06ms step:1107/1555 train_time:56559ms step_avg:51.09ms step:1108/1555 train_time:56649ms step_avg:51.13ms step:1109/1555 train_time:56733ms step_avg:51.16ms step:1110/1555 train_time:56823ms step_avg:51.19ms step:1111/1555 train_time:56908ms step_avg:51.22ms step:1112/1555 train_time:56999ms step_avg:51.26ms step:1113/1555 train_time:57083ms step_avg:51.29ms step:1114/1555 train_time:57173ms step_avg:51.32ms step:1115/1555 train_time:57258ms step_avg:51.35ms step:1116/1555 train_time:57346ms step_avg:51.38ms step:1117/1555 train_time:57430ms step_avg:51.41ms step:1118/1555 train_time:57521ms step_avg:51.45ms step:1119/1555 train_time:57605ms step_avg:51.48ms step:1120/1555 train_time:57695ms step_avg:51.51ms step:1121/1555 train_time:57779ms step_avg:51.54ms step:1122/1555 train_time:57868ms step_avg:51.58ms step:1123/1555 train_time:57952ms step_avg:51.60ms step:1124/1555 train_time:58042ms step_avg:51.64ms step:1125/1555 train_time:58127ms step_avg:51.67ms step:1126/1555 train_time:58217ms step_avg:51.70ms step:1127/1555 train_time:58301ms step_avg:51.73ms step:1128/1555 train_time:58390ms step_avg:51.76ms step:1129/1555 train_time:58474ms step_avg:51.79ms step:1130/1555 train_time:58564ms step_avg:51.83ms step:1131/1555 train_time:58649ms step_avg:51.86ms step:1132/1555 train_time:58739ms step_avg:51.89ms step:1133/1555 train_time:58822ms step_avg:51.92ms step:1134/1555 train_time:58913ms step_avg:51.95ms step:1135/1555 train_time:58997ms step_avg:51.98ms step:1136/1555 train_time:59087ms step_avg:52.01ms step:1137/1555 train_time:59171ms step_avg:52.04ms step:1138/1555 train_time:59261ms step_avg:52.07ms step:1139/1555 train_time:59345ms step_avg:52.10ms step:1140/1555 train_time:59436ms step_avg:52.14ms step:1141/1555 train_time:59519ms step_avg:52.16ms step:1142/1555 train_time:59608ms step_avg:52.20ms step:1143/1555 train_time:59693ms step_avg:52.22ms step:1144/1555 train_time:59783ms step_avg:52.26ms step:1145/1555 train_time:59868ms step_avg:52.29ms step:1146/1555 train_time:59960ms step_avg:52.32ms step:1147/1555 train_time:60043ms step_avg:52.35ms step:1148/1555 train_time:60132ms step_avg:52.38ms step:1149/1555 train_time:60216ms step_avg:52.41ms step:1150/1555 train_time:60306ms step_avg:52.44ms step:1151/1555 train_time:60391ms step_avg:52.47ms step:1152/1555 train_time:60481ms step_avg:52.50ms step:1153/1555 train_time:60565ms step_avg:52.53ms step:1154/1555 train_time:60655ms step_avg:52.56ms step:1155/1555 train_time:60739ms step_avg:52.59ms step:1156/1555 train_time:60829ms step_avg:52.62ms step:1157/1555 train_time:60913ms step_avg:52.65ms step:1158/1555 train_time:61003ms step_avg:52.68ms step:1159/1555 train_time:61088ms step_avg:52.71ms step:1160/1555 train_time:61179ms step_avg:52.74ms step:1161/1555 train_time:61262ms step_avg:52.77ms step:1162/1555 train_time:61352ms step_avg:52.80ms step:1163/1555 train_time:61438ms step_avg:52.83ms step:1164/1555 train_time:61526ms step_avg:52.86ms step:1165/1555 train_time:61611ms step_avg:52.88ms step:1166/1555 train_time:61701ms step_avg:52.92ms step:1167/1555 train_time:61785ms step_avg:52.94ms step:1168/1555 train_time:61876ms step_avg:52.98ms step:1169/1555 train_time:61960ms step_avg:53.00ms step:1170/1555 train_time:62049ms step_avg:53.03ms step:1171/1555 train_time:62134ms step_avg:53.06ms step:1172/1555 train_time:62224ms step_avg:53.09ms step:1173/1555 train_time:62309ms step_avg:53.12ms step:1174/1555 train_time:62399ms step_avg:53.15ms step:1175/1555 train_time:62483ms step_avg:53.18ms step:1176/1555 train_time:62574ms step_avg:53.21ms step:1177/1555 train_time:62658ms step_avg:53.24ms step:1178/1555 train_time:62747ms step_avg:53.27ms step:1179/1555 train_time:62831ms step_avg:53.29ms step:1180/1555 train_time:62920ms step_avg:53.32ms step:1181/1555 train_time:63004ms step_avg:53.35ms step:1182/1555 train_time:63095ms step_avg:53.38ms step:1183/1555 train_time:63178ms step_avg:53.40ms step:1184/1555 train_time:63267ms step_avg:53.44ms step:1185/1555 train_time:63352ms step_avg:53.46ms step:1186/1555 train_time:63442ms step_avg:53.49ms step:1187/1555 train_time:63526ms step_avg:53.52ms step:1188/1555 train_time:63617ms step_avg:53.55ms step:1189/1555 train_time:63700ms step_avg:53.57ms step:1190/1555 train_time:63790ms step_avg:53.61ms step:1191/1555 train_time:63874ms step_avg:53.63ms step:1192/1555 train_time:63963ms step_avg:53.66ms step:1193/1555 train_time:64047ms step_avg:53.69ms step:1194/1555 train_time:64138ms step_avg:53.72ms step:1195/1555 train_time:64221ms step_avg:53.74ms step:1196/1555 train_time:64311ms step_avg:53.77ms step:1197/1555 train_time:64395ms step_avg:53.80ms step:1198/1555 train_time:64485ms step_avg:53.83ms step:1199/1555 train_time:64570ms step_avg:53.85ms step:1200/1555 train_time:64660ms step_avg:53.88ms step:1201/1555 train_time:64744ms step_avg:53.91ms step:1202/1555 train_time:64834ms step_avg:53.94ms step:1203/1555 train_time:64919ms step_avg:53.96ms step:1204/1555 train_time:65008ms step_avg:53.99ms step:1205/1555 train_time:65092ms step_avg:54.02ms step:1206/1555 train_time:65182ms step_avg:54.05ms step:1207/1555 train_time:65265ms step_avg:54.07ms step:1208/1555 train_time:65356ms step_avg:54.10ms step:1209/1555 train_time:65440ms step_avg:54.13ms step:1210/1555 train_time:65529ms step_avg:54.16ms step:1211/1555 train_time:65614ms step_avg:54.18ms step:1212/1555 train_time:65704ms step_avg:54.21ms step:1213/1555 train_time:65789ms step_avg:54.24ms step:1214/1555 train_time:65880ms step_avg:54.27ms step:1215/1555 train_time:65964ms step_avg:54.29ms step:1216/1555 train_time:66054ms step_avg:54.32ms step:1217/1555 train_time:66138ms step_avg:54.34ms step:1218/1555 train_time:66226ms step_avg:54.37ms step:1219/1555 train_time:66311ms step_avg:54.40ms step:1220/1555 train_time:66401ms step_avg:54.43ms step:1221/1555 train_time:66484ms step_avg:54.45ms step:1222/1555 train_time:66575ms step_avg:54.48ms step:1223/1555 train_time:66660ms step_avg:54.51ms step:1224/1555 train_time:66749ms step_avg:54.53ms step:1225/1555 train_time:66834ms step_avg:54.56ms step:1226/1555 train_time:66923ms step_avg:54.59ms step:1227/1555 train_time:67008ms step_avg:54.61ms step:1228/1555 train_time:67098ms step_avg:54.64ms step:1229/1555 train_time:67182ms step_avg:54.66ms step:1230/1555 train_time:67272ms step_avg:54.69ms step:1231/1555 train_time:67356ms step_avg:54.72ms step:1232/1555 train_time:67446ms step_avg:54.75ms step:1233/1555 train_time:67531ms step_avg:54.77ms step:1234/1555 train_time:67621ms step_avg:54.80ms step:1235/1555 train_time:67704ms step_avg:54.82ms step:1236/1555 train_time:67794ms step_avg:54.85ms step:1237/1555 train_time:67879ms step_avg:54.87ms step:1238/1555 train_time:67968ms step_avg:54.90ms step:1239/1555 train_time:68053ms step_avg:54.93ms step:1240/1555 train_time:68143ms step_avg:54.95ms step:1241/1555 train_time:68227ms step_avg:54.98ms step:1242/1555 train_time:68317ms step_avg:55.01ms step:1243/1555 train_time:68400ms step_avg:55.03ms step:1244/1555 train_time:68489ms step_avg:55.06ms step:1245/1555 train_time:68574ms step_avg:55.08ms step:1246/1555 train_time:68664ms step_avg:55.11ms step:1247/1555 train_time:68749ms step_avg:55.13ms step:1248/1555 train_time:68839ms step_avg:55.16ms step:1249/1555 train_time:68922ms step_avg:55.18ms step:1250/1555 train_time:69013ms step_avg:55.21ms step:1250/1555 val_loss:3.3977 train_time:69127ms step_avg:55.30ms step:1251/1555 train_time:69149ms step_avg:55.28ms step:1252/1555 train_time:69188ms step_avg:55.26ms step:1253/1555 train_time:69275ms step_avg:55.29ms step:1254/1555 train_time:69371ms step_avg:55.32ms step:1255/1555 train_time:69455ms step_avg:55.34ms step:1256/1555 train_time:69546ms step_avg:55.37ms step:1257/1555 train_time:69629ms step_avg:55.39ms step:1258/1555 train_time:69717ms step_avg:55.42ms step:1259/1555 train_time:69801ms step_avg:55.44ms step:1260/1555 train_time:69890ms step_avg:55.47ms step:1261/1555 train_time:69972ms step_avg:55.49ms step:1262/1555 train_time:70062ms step_avg:55.52ms step:1263/1555 train_time:70149ms step_avg:55.54ms step:1264/1555 train_time:70240ms step_avg:55.57ms step:1265/1555 train_time:70327ms step_avg:55.59ms step:1266/1555 train_time:70417ms step_avg:55.62ms step:1267/1555 train_time:70502ms step_avg:55.65ms step:1268/1555 train_time:70592ms step_avg:55.67ms step:1269/1555 train_time:70674ms step_avg:55.69ms step:1270/1555 train_time:70765ms step_avg:55.72ms step:1271/1555 train_time:70848ms step_avg:55.74ms step:1272/1555 train_time:70936ms step_avg:55.77ms step:1273/1555 train_time:71020ms step_avg:55.79ms step:1274/1555 train_time:71110ms step_avg:55.82ms step:1275/1555 train_time:71194ms step_avg:55.84ms step:1276/1555 train_time:71286ms step_avg:55.87ms step:1277/1555 train_time:71371ms step_avg:55.89ms step:1278/1555 train_time:71461ms step_avg:55.92ms step:1279/1555 train_time:71546ms step_avg:55.94ms step:1280/1555 train_time:71636ms step_avg:55.97ms step:1281/1555 train_time:71721ms step_avg:55.99ms step:1282/1555 train_time:71810ms step_avg:56.01ms step:1283/1555 train_time:71893ms step_avg:56.03ms step:1284/1555 train_time:71982ms step_avg:56.06ms step:1285/1555 train_time:72065ms step_avg:56.08ms step:1286/1555 train_time:72156ms step_avg:56.11ms step:1287/1555 train_time:72241ms step_avg:56.13ms step:1288/1555 train_time:72332ms step_avg:56.16ms step:1289/1555 train_time:72416ms step_avg:56.18ms step:1290/1555 train_time:72507ms step_avg:56.21ms step:1291/1555 train_time:72591ms step_avg:56.23ms step:1292/1555 train_time:72679ms step_avg:56.25ms step:1293/1555 train_time:72763ms step_avg:56.27ms step:1294/1555 train_time:72852ms step_avg:56.30ms step:1295/1555 train_time:72935ms step_avg:56.32ms step:1296/1555 train_time:73025ms step_avg:56.35ms step:1297/1555 train_time:73109ms step_avg:56.37ms step:1298/1555 train_time:73200ms step_avg:56.39ms step:1299/1555 train_time:73285ms step_avg:56.42ms step:1300/1555 train_time:73376ms step_avg:56.44ms step:1301/1555 train_time:73460ms step_avg:56.46ms step:1302/1555 train_time:73552ms step_avg:56.49ms step:1303/1555 train_time:73634ms step_avg:56.51ms step:1304/1555 train_time:73725ms step_avg:56.54ms step:1305/1555 train_time:73810ms step_avg:56.56ms step:1306/1555 train_time:73898ms step_avg:56.58ms step:1307/1555 train_time:73982ms step_avg:56.60ms step:1308/1555 train_time:74072ms step_avg:56.63ms step:1309/1555 train_time:74156ms step_avg:56.65ms step:1310/1555 train_time:74247ms step_avg:56.68ms step:1311/1555 train_time:74331ms step_avg:56.70ms step:1312/1555 train_time:74420ms step_avg:56.72ms step:1313/1555 train_time:74505ms step_avg:56.74ms step:1314/1555 train_time:74596ms step_avg:56.77ms step:1315/1555 train_time:74680ms step_avg:56.79ms step:1316/1555 train_time:74771ms step_avg:56.82ms step:1317/1555 train_time:74854ms step_avg:56.84ms step:1318/1555 train_time:74943ms step_avg:56.86ms step:1319/1555 train_time:75026ms step_avg:56.88ms step:1320/1555 train_time:75116ms step_avg:56.91ms step:1321/1555 train_time:75201ms step_avg:56.93ms step:1322/1555 train_time:75292ms step_avg:56.95ms step:1323/1555 train_time:75376ms step_avg:56.97ms step:1324/1555 train_time:75466ms step_avg:57.00ms step:1325/1555 train_time:75551ms step_avg:57.02ms step:1326/1555 train_time:75642ms step_avg:57.05ms step:1327/1555 train_time:75725ms step_avg:57.07ms step:1328/1555 train_time:75815ms step_avg:57.09ms step:1329/1555 train_time:75898ms step_avg:57.11ms step:1330/1555 train_time:75988ms step_avg:57.13ms step:1331/1555 train_time:76072ms step_avg:57.15ms step:1332/1555 train_time:76161ms step_avg:57.18ms step:1333/1555 train_time:76246ms step_avg:57.20ms step:1334/1555 train_time:76335ms step_avg:57.22ms step:1335/1555 train_time:76420ms step_avg:57.24ms step:1336/1555 train_time:76510ms step_avg:57.27ms step:1337/1555 train_time:76594ms step_avg:57.29ms step:1338/1555 train_time:76684ms step_avg:57.31ms step:1339/1555 train_time:76768ms step_avg:57.33ms step:1340/1555 train_time:76856ms step_avg:57.36ms step:1341/1555 train_time:76940ms step_avg:57.38ms step:1342/1555 train_time:77031ms step_avg:57.40ms step:1343/1555 train_time:77114ms step_avg:57.42ms step:1344/1555 train_time:77204ms step_avg:57.44ms step:1345/1555 train_time:77289ms step_avg:57.46ms step:1346/1555 train_time:77378ms step_avg:57.49ms step:1347/1555 train_time:77462ms step_avg:57.51ms step:1348/1555 train_time:77553ms step_avg:57.53ms step:1349/1555 train_time:77636ms step_avg:57.55ms step:1350/1555 train_time:77726ms step_avg:57.58ms step:1351/1555 train_time:77810ms step_avg:57.59ms step:1352/1555 train_time:77899ms step_avg:57.62ms step:1353/1555 train_time:77984ms step_avg:57.64ms step:1354/1555 train_time:78075ms step_avg:57.66ms step:1355/1555 train_time:78158ms step_avg:57.68ms step:1356/1555 train_time:78249ms step_avg:57.71ms step:1357/1555 train_time:78333ms step_avg:57.73ms step:1358/1555 train_time:78423ms step_avg:57.75ms step:1359/1555 train_time:78507ms step_avg:57.77ms step:1360/1555 train_time:78599ms step_avg:57.79ms step:1361/1555 train_time:78684ms step_avg:57.81ms step:1362/1555 train_time:78775ms step_avg:57.84ms step:1363/1555 train_time:78858ms step_avg:57.86ms step:1364/1555 train_time:78950ms step_avg:57.88ms step:1365/1555 train_time:79032ms step_avg:57.90ms step:1366/1555 train_time:79123ms step_avg:57.92ms step:1367/1555 train_time:79207ms step_avg:57.94ms step:1368/1555 train_time:79297ms step_avg:57.97ms step:1369/1555 train_time:79381ms step_avg:57.98ms step:1370/1555 train_time:79473ms step_avg:58.01ms step:1371/1555 train_time:79557ms step_avg:58.03ms step:1372/1555 train_time:79647ms step_avg:58.05ms step:1373/1555 train_time:79731ms step_avg:58.07ms step:1374/1555 train_time:79820ms step_avg:58.09ms step:1375/1555 train_time:79905ms step_avg:58.11ms step:1376/1555 train_time:79994ms step_avg:58.14ms step:1377/1555 train_time:80078ms step_avg:58.15ms step:1378/1555 train_time:80169ms step_avg:58.18ms step:1379/1555 train_time:80252ms step_avg:58.20ms step:1380/1555 train_time:80342ms step_avg:58.22ms step:1381/1555 train_time:80426ms step_avg:58.24ms step:1382/1555 train_time:80515ms step_avg:58.26ms step:1383/1555 train_time:80600ms step_avg:58.28ms step:1384/1555 train_time:80691ms step_avg:58.30ms step:1385/1555 train_time:80774ms step_avg:58.32ms step:1386/1555 train_time:80864ms step_avg:58.34ms step:1387/1555 train_time:80948ms step_avg:58.36ms step:1388/1555 train_time:81037ms step_avg:58.38ms step:1389/1555 train_time:81120ms step_avg:58.40ms step:1390/1555 train_time:81211ms step_avg:58.43ms step:1391/1555 train_time:81295ms step_avg:58.44ms step:1392/1555 train_time:81386ms step_avg:58.47ms step:1393/1555 train_time:81470ms step_avg:58.49ms step:1394/1555 train_time:81559ms step_avg:58.51ms step:1395/1555 train_time:81644ms step_avg:58.53ms step:1396/1555 train_time:81733ms step_avg:58.55ms step:1397/1555 train_time:81818ms step_avg:58.57ms step:1398/1555 train_time:81909ms step_avg:58.59ms step:1399/1555 train_time:81992ms step_avg:58.61ms step:1400/1555 train_time:82081ms step_avg:58.63ms step:1401/1555 train_time:82165ms step_avg:58.65ms step:1402/1555 train_time:82256ms step_avg:58.67ms step:1403/1555 train_time:82340ms step_avg:58.69ms step:1404/1555 train_time:82431ms step_avg:58.71ms step:1405/1555 train_time:82514ms step_avg:58.73ms step:1406/1555 train_time:82603ms step_avg:58.75ms step:1407/1555 train_time:82687ms step_avg:58.77ms step:1408/1555 train_time:82777ms step_avg:58.79ms step:1409/1555 train_time:82861ms step_avg:58.81ms step:1410/1555 train_time:82953ms step_avg:58.83ms step:1411/1555 train_time:83036ms step_avg:58.85ms step:1412/1555 train_time:83126ms step_avg:58.87ms step:1413/1555 train_time:83210ms step_avg:58.89ms step:1414/1555 train_time:83300ms step_avg:58.91ms step:1415/1555 train_time:83385ms step_avg:58.93ms step:1416/1555 train_time:83477ms step_avg:58.95ms step:1417/1555 train_time:83559ms step_avg:58.97ms step:1418/1555 train_time:83650ms step_avg:58.99ms step:1419/1555 train_time:83733ms step_avg:59.01ms step:1420/1555 train_time:83824ms step_avg:59.03ms step:1421/1555 train_time:83907ms step_avg:59.05ms step:1422/1555 train_time:83997ms step_avg:59.07ms step:1423/1555 train_time:84080ms step_avg:59.09ms step:1424/1555 train_time:84172ms step_avg:59.11ms step:1425/1555 train_time:84255ms step_avg:59.13ms step:1426/1555 train_time:84345ms step_avg:59.15ms step:1427/1555 train_time:84430ms step_avg:59.17ms step:1428/1555 train_time:84519ms step_avg:59.19ms step:1429/1555 train_time:84603ms step_avg:59.20ms step:1430/1555 train_time:84693ms step_avg:59.23ms step:1431/1555 train_time:84777ms step_avg:59.24ms step:1432/1555 train_time:84870ms step_avg:59.27ms step:1433/1555 train_time:84953ms step_avg:59.28ms step:1434/1555 train_time:85043ms step_avg:59.31ms step:1435/1555 train_time:85127ms step_avg:59.32ms step:1436/1555 train_time:85216ms step_avg:59.34ms step:1437/1555 train_time:85300ms step_avg:59.36ms step:1438/1555 train_time:85392ms step_avg:59.38ms step:1439/1555 train_time:85475ms step_avg:59.40ms step:1440/1555 train_time:85565ms step_avg:59.42ms step:1441/1555 train_time:85649ms step_avg:59.44ms step:1442/1555 train_time:85739ms step_avg:59.46ms step:1443/1555 train_time:85824ms step_avg:59.48ms step:1444/1555 train_time:85914ms step_avg:59.50ms step:1445/1555 train_time:85998ms step_avg:59.51ms step:1446/1555 train_time:86090ms step_avg:59.54ms step:1447/1555 train_time:86174ms step_avg:59.55ms step:1448/1555 train_time:86263ms step_avg:59.57ms step:1449/1555 train_time:86349ms step_avg:59.59ms step:1450/1555 train_time:86437ms step_avg:59.61ms step:1451/1555 train_time:86522ms step_avg:59.63ms step:1452/1555 train_time:86612ms step_avg:59.65ms step:1453/1555 train_time:86696ms step_avg:59.67ms step:1454/1555 train_time:86787ms step_avg:59.69ms step:1455/1555 train_time:86870ms step_avg:59.70ms step:1456/1555 train_time:86960ms step_avg:59.73ms step:1457/1555 train_time:87046ms step_avg:59.74ms step:1458/1555 train_time:87134ms step_avg:59.76ms step:1459/1555 train_time:87219ms step_avg:59.78ms step:1460/1555 train_time:87310ms step_avg:59.80ms step:1461/1555 train_time:87393ms step_avg:59.82ms step:1462/1555 train_time:87483ms step_avg:59.84ms step:1463/1555 train_time:87567ms step_avg:59.85ms step:1464/1555 train_time:87657ms step_avg:59.88ms step:1465/1555 train_time:87740ms step_avg:59.89ms step:1466/1555 train_time:87831ms step_avg:59.91ms step:1467/1555 train_time:87914ms step_avg:59.93ms step:1468/1555 train_time:88005ms step_avg:59.95ms step:1469/1555 train_time:88090ms step_avg:59.97ms step:1470/1555 train_time:88179ms step_avg:59.99ms step:1471/1555 train_time:88264ms step_avg:60.00ms step:1472/1555 train_time:88354ms step_avg:60.02ms step:1473/1555 train_time:88439ms step_avg:60.04ms step:1474/1555 train_time:88528ms step_avg:60.06ms step:1475/1555 train_time:88611ms step_avg:60.08ms step:1476/1555 train_time:88700ms step_avg:60.09ms step:1477/1555 train_time:88784ms step_avg:60.11ms step:1478/1555 train_time:88875ms step_avg:60.13ms step:1479/1555 train_time:88958ms step_avg:60.15ms step:1480/1555 train_time:89050ms step_avg:60.17ms step:1481/1555 train_time:89133ms step_avg:60.18ms step:1482/1555 train_time:89224ms step_avg:60.20ms step:1483/1555 train_time:89309ms step_avg:60.22ms step:1484/1555 train_time:89397ms step_avg:60.24ms step:1485/1555 train_time:89481ms step_avg:60.26ms step:1486/1555 train_time:89571ms step_avg:60.28ms step:1487/1555 train_time:89655ms step_avg:60.29ms step:1488/1555 train_time:89746ms step_avg:60.31ms step:1489/1555 train_time:89829ms step_avg:60.33ms step:1490/1555 train_time:89919ms step_avg:60.35ms step:1491/1555 train_time:90004ms step_avg:60.36ms step:1492/1555 train_time:90094ms step_avg:60.38ms step:1493/1555 train_time:90178ms step_avg:60.40ms step:1494/1555 train_time:90269ms step_avg:60.42ms step:1495/1555 train_time:90353ms step_avg:60.44ms step:1496/1555 train_time:90442ms step_avg:60.46ms step:1497/1555 train_time:90526ms step_avg:60.47ms step:1498/1555 train_time:90615ms step_avg:60.49ms step:1499/1555 train_time:90699ms step_avg:60.51ms step:1500/1555 train_time:90791ms step_avg:60.53ms step:1500/1555 val_loss:3.2939 train_time:90905ms step_avg:60.60ms step:1501/1555 train_time:90928ms step_avg:60.58ms step:1502/1555 train_time:90967ms step_avg:60.56ms step:1503/1555 train_time:91051ms step_avg:60.58ms step:1504/1555 train_time:91146ms step_avg:60.60ms step:1505/1555 train_time:91231ms step_avg:60.62ms step:1506/1555 train_time:91322ms step_avg:60.64ms step:1507/1555 train_time:91405ms step_avg:60.65ms step:1508/1555 train_time:91493ms step_avg:60.67ms step:1509/1555 train_time:91577ms step_avg:60.69ms step:1510/1555 train_time:91666ms step_avg:60.71ms step:1511/1555 train_time:91749ms step_avg:60.72ms step:1512/1555 train_time:91840ms step_avg:60.74ms step:1513/1555 train_time:91925ms step_avg:60.76ms step:1514/1555 train_time:92016ms step_avg:60.78ms step:1515/1555 train_time:92103ms step_avg:60.79ms step:1516/1555 train_time:92198ms step_avg:60.82ms step:1517/1555 train_time:92284ms step_avg:60.83ms step:1518/1555 train_time:92373ms step_avg:60.85ms step:1519/1555 train_time:92457ms step_avg:60.87ms step:1520/1555 train_time:92546ms step_avg:60.89ms step:1521/1555 train_time:92630ms step_avg:60.90ms step:1522/1555 train_time:92720ms step_avg:60.92ms step:1523/1555 train_time:92804ms step_avg:60.94ms step:1524/1555 train_time:92895ms step_avg:60.95ms step:1525/1555 train_time:92980ms step_avg:60.97ms step:1526/1555 train_time:93071ms step_avg:60.99ms step:1527/1555 train_time:93159ms step_avg:61.01ms step:1528/1555 train_time:93248ms step_avg:61.03ms step:1529/1555 train_time:93334ms step_avg:61.04ms step:1530/1555 train_time:93423ms step_avg:61.06ms step:1531/1555 train_time:93506ms step_avg:61.08ms step:1532/1555 train_time:93596ms step_avg:61.09ms step:1533/1555 train_time:93680ms step_avg:61.11ms step:1534/1555 train_time:93769ms step_avg:61.13ms step:1535/1555 train_time:93854ms step_avg:61.14ms step:1536/1555 train_time:93945ms step_avg:61.16ms step:1537/1555 train_time:94031ms step_avg:61.18ms step:1538/1555 train_time:94123ms step_avg:61.20ms step:1539/1555 train_time:94207ms step_avg:61.21ms step:1540/1555 train_time:94299ms step_avg:61.23ms step:1541/1555 train_time:94383ms step_avg:61.25ms step:1542/1555 train_time:94473ms step_avg:61.27ms step:1543/1555 train_time:94557ms step_avg:61.28ms step:1544/1555 train_time:94646ms step_avg:61.30ms step:1545/1555 train_time:94730ms step_avg:61.31ms step:1546/1555 train_time:94821ms step_avg:61.33ms step:1547/1555 train_time:94906ms step_avg:61.35ms step:1548/1555 train_time:94997ms step_avg:61.37ms step:1549/1555 train_time:95083ms step_avg:61.38ms step:1550/1555 train_time:95172ms step_avg:61.40ms step:1551/1555 train_time:95259ms step_avg:61.42ms step:1552/1555 train_time:95348ms step_avg:61.44ms step:1553/1555 train_time:95433ms step_avg:61.45ms step:1554/1555 train_time:95524ms step_avg:61.47ms step:1555/1555 train_time:95608ms step_avg:61.48ms step:1555/1555 val_loss:3.2777 train_time:95723ms step_avg:61.56ms peak memory allocated: 30746 MiB reserved: 46798 MiB