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:06:54 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 33C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 37C P0 123W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:71:00.0 Off | 0 | | N/A 39C P0 125W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:79:00.0 Off | 0 | | N/A 34C P0 124W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:7F:00.0 Off | 0 | | N/A 32C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:87:00.0 Off | 0 | | N/A 39C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:8D:00.0 Off | 0 | | N/A 37C P0 123W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:95:00.0 Off | 0 | | N/A 34C P0 118W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 15516 C /usr/local/bin/python 1510MiB | | 1 N/A N/A 15517 C /usr/local/bin/python 1510MiB | | 2 N/A N/A 15518 C /usr/local/bin/python 1510MiB | | 3 N/A N/A 15519 C /usr/local/bin/python 1510MiB | | 4 N/A N/A 15520 C /usr/local/bin/python 1510MiB | | 5 N/A N/A 15521 C /usr/local/bin/python 1510MiB | | 6 N/A N/A 15522 C /usr/local/bin/python 1510MiB | | 7 N/A N/A 15523 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.8306 train_time:0ms step_avg:0.03ms step:1/1555 train_time:94ms step_avg:93.96ms step:2/1555 train_time:121ms step_avg:60.66ms step:3/1555 train_time:140ms step_avg:46.72ms step:4/1555 train_time:159ms step_avg:39.65ms step:5/1555 train_time:183ms step_avg:36.60ms step:6/1555 train_time:220ms step_avg:36.68ms step:7/1555 train_time:251ms step_avg:35.85ms step:8/1555 train_time:288ms step_avg:36.05ms step:9/1555 train_time:319ms step_avg:35.45ms step:10/1555 train_time:356ms step_avg:35.65ms step:11/1555 train_time:388ms step_avg:35.25ms step:12/1555 train_time:425ms step_avg:35.43ms step:13/1555 train_time:456ms step_avg:35.10ms step:14/1555 train_time:494ms step_avg:35.26ms step:15/1555 train_time:525ms step_avg:34.97ms step:16/1555 train_time:562ms step_avg:35.15ms step:17/1555 train_time:593ms step_avg:34.91ms step:18/1555 train_time:631ms step_avg:35.04ms step:19/1555 train_time:662ms step_avg:34.83ms step:20/1555 train_time:700ms step_avg:34.99ms step:21/1555 train_time:731ms step_avg:34.79ms step:22/1555 train_time:768ms step_avg:34.91ms step:23/1555 train_time:799ms step_avg:34.75ms step:24/1555 train_time:837ms step_avg:34.86ms step:25/1555 train_time:868ms step_avg:34.71ms step:26/1555 train_time:905ms step_avg:34.82ms step:27/1555 train_time:937ms step_avg:34.70ms step:28/1555 train_time:974ms step_avg:34.79ms step:29/1555 train_time:1006ms step_avg:34.68ms step:30/1555 train_time:1045ms step_avg:34.82ms step:31/1555 train_time:1076ms step_avg:34.71ms step:32/1555 train_time:1113ms step_avg:34.79ms step:33/1555 train_time:1145ms step_avg:34.70ms step:34/1555 train_time:1183ms step_avg:34.80ms step:35/1555 train_time:1214ms step_avg:34.69ms step:36/1555 train_time:1252ms step_avg:34.77ms step:37/1555 train_time:1284ms step_avg:34.69ms step:38/1555 train_time:1321ms step_avg:34.77ms step:39/1555 train_time:1353ms step_avg:34.68ms step:40/1555 train_time:1390ms step_avg:34.75ms step:41/1555 train_time:1421ms step_avg:34.66ms step:42/1555 train_time:1458ms step_avg:34.72ms step:43/1555 train_time:1489ms step_avg:34.64ms step:44/1555 train_time:1528ms step_avg:34.72ms step:45/1555 train_time:1558ms step_avg:34.62ms step:46/1555 train_time:1595ms step_avg:34.68ms step:47/1555 train_time:1626ms step_avg:34.60ms step:48/1555 train_time:1664ms step_avg:34.66ms step:49/1555 train_time:1695ms step_avg:34.58ms step:50/1555 train_time:1732ms step_avg:34.64ms step:51/1555 train_time:1763ms step_avg:34.56ms step:52/1555 train_time:1800ms step_avg:34.62ms step:53/1555 train_time:1832ms step_avg:34.56ms step:54/1555 train_time:1869ms step_avg:34.61ms step:55/1555 train_time:1900ms step_avg:34.54ms step:56/1555 train_time:1938ms step_avg:34.60ms step:57/1555 train_time:1969ms step_avg:34.54ms step:58/1555 train_time:2007ms step_avg:34.60ms step:59/1555 train_time:2038ms step_avg:34.54ms step:60/1555 train_time:2075ms step_avg:34.59ms step:61/1555 train_time:2107ms step_avg:34.54ms step:62/1555 train_time:2145ms step_avg:34.59ms step:63/1555 train_time:2176ms step_avg:34.54ms step:64/1555 train_time:2214ms step_avg:34.59ms step:65/1555 train_time:2245ms step_avg:34.54ms step:66/1555 train_time:2283ms step_avg:34.60ms step:67/1555 train_time:2315ms step_avg:34.55ms step:68/1555 train_time:2352ms step_avg:34.59ms step:69/1555 train_time:2384ms step_avg:34.55ms step:70/1555 train_time:2422ms step_avg:34.60ms step:71/1555 train_time:2453ms step_avg:34.55ms step:72/1555 train_time:2490ms step_avg:34.59ms step:73/1555 train_time:2522ms step_avg:34.55ms step:74/1555 train_time:2560ms step_avg:34.59ms step:75/1555 train_time:2591ms step_avg:34.55ms step:76/1555 train_time:2629ms step_avg:34.59ms step:77/1555 train_time:2660ms step_avg:34.54ms step:78/1555 train_time:2697ms step_avg:34.58ms step:79/1555 train_time:2728ms step_avg:34.54ms step:80/1555 train_time:2766ms step_avg:34.57ms step:81/1555 train_time:2797ms step_avg:34.53ms step:82/1555 train_time:2834ms step_avg:34.56ms step:83/1555 train_time:2865ms step_avg:34.52ms step:84/1555 train_time:2903ms step_avg:34.56ms step:85/1555 train_time:2935ms step_avg:34.52ms step:86/1555 train_time:2972ms step_avg:34.56ms step:87/1555 train_time:3004ms step_avg:34.53ms step:88/1555 train_time:3042ms step_avg:34.57ms step:89/1555 train_time:3073ms step_avg:34.53ms step:90/1555 train_time:3110ms step_avg:34.56ms step:91/1555 train_time:3141ms step_avg:34.52ms step:92/1555 train_time:3179ms step_avg:34.56ms step:93/1555 train_time:3210ms step_avg:34.52ms step:94/1555 train_time:3247ms step_avg:34.55ms step:95/1555 train_time:3278ms step_avg:34.51ms step:96/1555 train_time:3315ms step_avg:34.53ms step:97/1555 train_time:3347ms step_avg:34.50ms step:98/1555 train_time:3384ms step_avg:34.53ms step:99/1555 train_time:3415ms step_avg:34.50ms step:100/1555 train_time:3452ms step_avg:34.52ms step:101/1555 train_time:3483ms step_avg:34.49ms step:102/1555 train_time:3521ms step_avg:34.52ms step:103/1555 train_time:3552ms step_avg:34.48ms step:104/1555 train_time:3589ms step_avg:34.51ms step:105/1555 train_time:3620ms step_avg:34.48ms step:106/1555 train_time:3658ms step_avg:34.51ms step:107/1555 train_time:3689ms step_avg:34.47ms step:108/1555 train_time:3726ms step_avg:34.50ms step:109/1555 train_time:3757ms step_avg:34.47ms step:110/1555 train_time:3794ms step_avg:34.49ms step:111/1555 train_time:3826ms step_avg:34.46ms step:112/1555 train_time:3863ms step_avg:34.49ms step:113/1555 train_time:3894ms step_avg:34.46ms step:114/1555 train_time:3931ms step_avg:34.48ms step:115/1555 train_time:3962ms step_avg:34.46ms step:116/1555 train_time:4000ms step_avg:34.48ms step:117/1555 train_time:4031ms step_avg:34.45ms step:118/1555 train_time:4068ms step_avg:34.48ms step:119/1555 train_time:4099ms step_avg:34.45ms step:120/1555 train_time:4137ms step_avg:34.47ms step:121/1555 train_time:4168ms step_avg:34.45ms step:122/1555 train_time:4206ms step_avg:34.48ms step:123/1555 train_time:4236ms step_avg:34.44ms step:124/1555 train_time:4274ms step_avg:34.47ms step:125/1555 train_time:4305ms step_avg:34.44ms step:126/1555 train_time:4343ms step_avg:34.47ms step:127/1555 train_time:4374ms step_avg:34.44ms step:128/1555 train_time:4411ms step_avg:34.46ms step:129/1555 train_time:4443ms step_avg:34.44ms step:130/1555 train_time:4480ms step_avg:34.46ms step:131/1555 train_time:4511ms step_avg:34.44ms step:132/1555 train_time:4549ms step_avg:34.46ms step:133/1555 train_time:4581ms step_avg:34.44ms step:134/1555 train_time:4619ms step_avg:34.47ms step:135/1555 train_time:4650ms step_avg:34.44ms step:136/1555 train_time:4688ms step_avg:34.47ms step:137/1555 train_time:4719ms step_avg:34.44ms step:138/1555 train_time:4756ms step_avg:34.46ms step:139/1555 train_time:4787ms step_avg:34.44ms step:140/1555 train_time:4824ms step_avg:34.46ms step:141/1555 train_time:4855ms step_avg:34.43ms step:142/1555 train_time:4893ms step_avg:34.45ms step:143/1555 train_time:4924ms step_avg:34.43ms step:144/1555 train_time:4961ms step_avg:34.45ms step:145/1555 train_time:4993ms step_avg:34.43ms step:146/1555 train_time:5030ms step_avg:34.45ms step:147/1555 train_time:5061ms step_avg:34.43ms step:148/1555 train_time:5099ms step_avg:34.45ms step:149/1555 train_time:5130ms step_avg:34.43ms step:150/1555 train_time:5167ms step_avg:34.45ms step:151/1555 train_time:5198ms step_avg:34.43ms step:152/1555 train_time:5236ms step_avg:34.45ms step:153/1555 train_time:5267ms step_avg:34.42ms step:154/1555 train_time:5304ms step_avg:34.44ms step:155/1555 train_time:5336ms step_avg:34.43ms step:156/1555 train_time:5373ms step_avg:34.45ms step:157/1555 train_time:5405ms step_avg:34.43ms step:158/1555 train_time:5442ms step_avg:34.45ms step:159/1555 train_time:5474ms step_avg:34.43ms step:160/1555 train_time:5511ms step_avg:34.44ms step:161/1555 train_time:5543ms step_avg:34.43ms step:162/1555 train_time:5580ms step_avg:34.45ms step:163/1555 train_time:5612ms step_avg:34.43ms step:164/1555 train_time:5649ms step_avg:34.44ms step:165/1555 train_time:5680ms step_avg:34.42ms step:166/1555 train_time:5718ms step_avg:34.44ms step:167/1555 train_time:5749ms step_avg:34.42ms step:168/1555 train_time:5786ms step_avg:34.44ms step:169/1555 train_time:5817ms step_avg:34.42ms step:170/1555 train_time:5854ms step_avg:34.44ms step:171/1555 train_time:5885ms step_avg:34.42ms step:172/1555 train_time:5923ms step_avg:34.44ms step:173/1555 train_time:5954ms step_avg:34.41ms step:174/1555 train_time:5991ms step_avg:34.43ms step:175/1555 train_time:6022ms step_avg:34.41ms step:176/1555 train_time:6060ms step_avg:34.43ms step:177/1555 train_time:6091ms step_avg:34.41ms step:178/1555 train_time:6129ms step_avg:34.43ms step:179/1555 train_time:6160ms step_avg:34.41ms step:180/1555 train_time:6197ms step_avg:34.43ms step:181/1555 train_time:6230ms step_avg:34.42ms step:182/1555 train_time:6266ms step_avg:34.43ms step:183/1555 train_time:6297ms step_avg:34.41ms step:184/1555 train_time:6334ms step_avg:34.42ms step:185/1555 train_time:6366ms step_avg:34.41ms step:186/1555 train_time:6403ms step_avg:34.43ms step:187/1555 train_time:6434ms step_avg:34.41ms step:188/1555 train_time:6472ms step_avg:34.42ms step:189/1555 train_time:6503ms step_avg:34.41ms step:190/1555 train_time:6540ms step_avg:34.42ms step:191/1555 train_time:6571ms step_avg:34.40ms step:192/1555 train_time:6609ms step_avg:34.42ms step:193/1555 train_time:6640ms step_avg:34.41ms step:194/1555 train_time:6678ms step_avg:34.42ms step:195/1555 train_time:6709ms step_avg:34.41ms step:196/1555 train_time:6747ms step_avg:34.42ms step:197/1555 train_time:6778ms step_avg:34.41ms step:198/1555 train_time:6815ms step_avg:34.42ms step:199/1555 train_time:6846ms step_avg:34.40ms step:200/1555 train_time:6883ms step_avg:34.42ms step:201/1555 train_time:6914ms step_avg:34.40ms step:202/1555 train_time:6951ms step_avg:34.41ms step:203/1555 train_time:6982ms step_avg:34.40ms step:204/1555 train_time:7020ms step_avg:34.41ms step:205/1555 train_time:7051ms step_avg:34.39ms step:206/1555 train_time:7088ms step_avg:34.41ms step:207/1555 train_time:7119ms step_avg:34.39ms step:208/1555 train_time:7157ms step_avg:34.41ms step:209/1555 train_time:7188ms step_avg:34.39ms step:210/1555 train_time:7225ms step_avg:34.41ms step:211/1555 train_time:7256ms step_avg:34.39ms step:212/1555 train_time:7293ms step_avg:34.40ms step:213/1555 train_time:7325ms step_avg:34.39ms step:214/1555 train_time:7362ms step_avg:34.40ms step:215/1555 train_time:7393ms step_avg:34.39ms step:216/1555 train_time:7431ms step_avg:34.40ms step:217/1555 train_time:7462ms step_avg:34.39ms step:218/1555 train_time:7499ms step_avg:34.40ms step:219/1555 train_time:7530ms step_avg:34.39ms step:220/1555 train_time:7568ms step_avg:34.40ms step:221/1555 train_time:7599ms step_avg:34.38ms step:222/1555 train_time:7636ms step_avg:34.40ms step:223/1555 train_time:7668ms step_avg:34.38ms step:224/1555 train_time:7705ms step_avg:34.40ms step:225/1555 train_time:7736ms step_avg:34.38ms step:226/1555 train_time:7774ms step_avg:34.40ms step:227/1555 train_time:7805ms step_avg:34.38ms step:228/1555 train_time:7843ms step_avg:34.40ms step:229/1555 train_time:7874ms step_avg:34.38ms step:230/1555 train_time:7911ms step_avg:34.40ms step:231/1555 train_time:7942ms step_avg:34.38ms step:232/1555 train_time:7980ms step_avg:34.39ms step:233/1555 train_time:8011ms step_avg:34.38ms step:234/1555 train_time:8048ms step_avg:34.39ms step:235/1555 train_time:8079ms step_avg:34.38ms step:236/1555 train_time:8117ms step_avg:34.39ms step:237/1555 train_time:8148ms step_avg:34.38ms step:238/1555 train_time:8185ms step_avg:34.39ms step:239/1555 train_time:8216ms step_avg:34.38ms step:240/1555 train_time:8254ms step_avg:34.39ms step:241/1555 train_time:8286ms step_avg:34.38ms step:242/1555 train_time:8324ms step_avg:34.40ms step:243/1555 train_time:8355ms step_avg:34.38ms step:244/1555 train_time:8392ms step_avg:34.39ms step:245/1555 train_time:8423ms step_avg:34.38ms step:246/1555 train_time:8461ms step_avg:34.39ms step:247/1555 train_time:8492ms step_avg:34.38ms step:248/1555 train_time:8529ms step_avg:34.39ms step:249/1555 train_time:8560ms step_avg:34.38ms step:250/1555 train_time:8598ms step_avg:34.39ms step:250/1555 val_loss:4.5545 train_time:8648ms step_avg:34.59ms step:251/1555 train_time:8668ms step_avg:34.53ms step:252/1555 train_time:8687ms step_avg:34.47ms step:253/1555 train_time:8704ms step_avg:34.40ms step:254/1555 train_time:8737ms step_avg:34.40ms step:255/1555 train_time:8771ms step_avg:34.39ms step:256/1555 train_time:8809ms step_avg:34.41ms step:257/1555 train_time:8840ms step_avg:34.40ms step:258/1555 train_time:8878ms step_avg:34.41ms step:259/1555 train_time:8910ms step_avg:34.40ms step:260/1555 train_time:8948ms step_avg:34.41ms step:261/1555 train_time:8979ms step_avg:34.40ms step:262/1555 train_time:9016ms step_avg:34.41ms step:263/1555 train_time:9047ms step_avg:34.40ms step:264/1555 train_time:9084ms step_avg:34.41ms step:265/1555 train_time:9115ms step_avg:34.40ms step:266/1555 train_time:9153ms step_avg:34.41ms step:267/1555 train_time:9184ms step_avg:34.40ms step:268/1555 train_time:9221ms step_avg:34.41ms step:269/1555 train_time:9252ms step_avg:34.39ms step:270/1555 train_time:9289ms step_avg:34.40ms step:271/1555 train_time:9320ms step_avg:34.39ms step:272/1555 train_time:9357ms step_avg:34.40ms step:273/1555 train_time:9388ms step_avg:34.39ms step:274/1555 train_time:9425ms step_avg:34.40ms step:275/1555 train_time:9456ms step_avg:34.39ms step:276/1555 train_time:9493ms step_avg:34.40ms step:277/1555 train_time:9524ms step_avg:34.38ms step:278/1555 train_time:9561ms step_avg:34.39ms step:279/1555 train_time:9592ms step_avg:34.38ms step:280/1555 train_time:9630ms step_avg:34.39ms step:281/1555 train_time:9661ms step_avg:34.38ms step:282/1555 train_time:9698ms step_avg:34.39ms step:283/1555 train_time:9729ms step_avg:34.38ms step:284/1555 train_time:9767ms step_avg:34.39ms step:285/1555 train_time:9798ms step_avg:34.38ms step:286/1555 train_time:9836ms step_avg:34.39ms step:287/1555 train_time:9867ms step_avg:34.38ms step:288/1555 train_time:9905ms step_avg:34.39ms step:289/1555 train_time:9936ms step_avg:34.38ms step:290/1555 train_time:9974ms step_avg:34.39ms step:291/1555 train_time:10005ms step_avg:34.38ms step:292/1555 train_time:10042ms step_avg:34.39ms step:293/1555 train_time:10073ms step_avg:34.38ms step:294/1555 train_time:10111ms step_avg:34.39ms step:295/1555 train_time:10142ms step_avg:34.38ms step:296/1555 train_time:10180ms step_avg:34.39ms step:297/1555 train_time:10211ms step_avg:34.38ms step:298/1555 train_time:10248ms step_avg:34.39ms step:299/1555 train_time:10279ms step_avg:34.38ms step:300/1555 train_time:10317ms step_avg:34.39ms step:301/1555 train_time:10347ms step_avg:34.38ms step:302/1555 train_time:10385ms step_avg:34.39ms step:303/1555 train_time:10416ms step_avg:34.38ms step:304/1555 train_time:10454ms step_avg:34.39ms step:305/1555 train_time:10484ms step_avg:34.38ms step:306/1555 train_time:10522ms step_avg:34.38ms step:307/1555 train_time:10553ms step_avg:34.37ms step:308/1555 train_time:10591ms step_avg:34.39ms step:309/1555 train_time:10622ms step_avg:34.37ms step:310/1555 train_time:10659ms step_avg:34.38ms step:311/1555 train_time:10690ms step_avg:34.37ms step:312/1555 train_time:10727ms step_avg:34.38ms step:313/1555 train_time:10758ms step_avg:34.37ms step:314/1555 train_time:10796ms step_avg:34.38ms step:315/1555 train_time:10827ms step_avg:34.37ms step:316/1555 train_time:10864ms step_avg:34.38ms step:317/1555 train_time:10895ms step_avg:34.37ms step:318/1555 train_time:10933ms step_avg:34.38ms step:319/1555 train_time:10964ms step_avg:34.37ms step:320/1555 train_time:11001ms step_avg:34.38ms step:321/1555 train_time:11032ms step_avg:34.37ms step:322/1555 train_time:11070ms step_avg:34.38ms step:323/1555 train_time:11101ms step_avg:34.37ms step:324/1555 train_time:11139ms step_avg:34.38ms step:325/1555 train_time:11170ms step_avg:34.37ms step:326/1555 train_time:11208ms step_avg:34.38ms step:327/1555 train_time:11239ms step_avg:34.37ms step:328/1555 train_time:11277ms step_avg:34.38ms step:329/1555 train_time:11308ms step_avg:34.37ms step:330/1555 train_time:11345ms step_avg:34.38ms step:331/1555 train_time:11376ms step_avg:34.37ms step:332/1555 train_time:11414ms step_avg:34.38ms step:333/1555 train_time:11445ms step_avg:34.37ms step:334/1555 train_time:11482ms step_avg:34.38ms step:335/1555 train_time:11513ms step_avg:34.37ms step:336/1555 train_time:11551ms step_avg:34.38ms step:337/1555 train_time:11582ms step_avg:34.37ms step:338/1555 train_time:11619ms step_avg:34.38ms step:339/1555 train_time:11650ms step_avg:34.37ms step:340/1555 train_time:11687ms step_avg:34.37ms step:341/1555 train_time:11718ms step_avg:34.36ms step:342/1555 train_time:11756ms step_avg:34.37ms step:343/1555 train_time:11786ms step_avg:34.36ms step:344/1555 train_time:11823ms step_avg:34.37ms step:345/1555 train_time:11854ms step_avg:34.36ms step:346/1555 train_time:11892ms step_avg:34.37ms step:347/1555 train_time:11923ms step_avg:34.36ms step:348/1555 train_time:11960ms step_avg:34.37ms step:349/1555 train_time:11992ms step_avg:34.36ms step:350/1555 train_time:12029ms step_avg:34.37ms step:351/1555 train_time:12060ms step_avg:34.36ms step:352/1555 train_time:12098ms step_avg:34.37ms step:353/1555 train_time:12129ms step_avg:34.36ms step:354/1555 train_time:12166ms step_avg:34.37ms step:355/1555 train_time:12197ms step_avg:34.36ms step:356/1555 train_time:12235ms step_avg:34.37ms step:357/1555 train_time:12266ms step_avg:34.36ms step:358/1555 train_time:12304ms step_avg:34.37ms step:359/1555 train_time:12334ms step_avg:34.36ms step:360/1555 train_time:12372ms step_avg:34.37ms step:361/1555 train_time:12403ms step_avg:34.36ms step:362/1555 train_time:12440ms step_avg:34.37ms step:363/1555 train_time:12472ms step_avg:34.36ms step:364/1555 train_time:12509ms step_avg:34.37ms step:365/1555 train_time:12540ms step_avg:34.36ms step:366/1555 train_time:12578ms step_avg:34.37ms step:367/1555 train_time:12609ms step_avg:34.36ms step:368/1555 train_time:12646ms step_avg:34.36ms step:369/1555 train_time:12677ms step_avg:34.36ms step:370/1555 train_time:12715ms step_avg:34.36ms step:371/1555 train_time:12746ms step_avg:34.36ms step:372/1555 train_time:12783ms step_avg:34.36ms step:373/1555 train_time:12815ms step_avg:34.36ms step:374/1555 train_time:12853ms step_avg:34.37ms step:375/1555 train_time:12884ms step_avg:34.36ms step:376/1555 train_time:12921ms step_avg:34.36ms step:377/1555 train_time:12952ms step_avg:34.36ms step:378/1555 train_time:12989ms step_avg:34.36ms step:379/1555 train_time:13020ms step_avg:34.35ms step:380/1555 train_time:13058ms step_avg:34.36ms step:381/1555 train_time:13089ms step_avg:34.35ms step:382/1555 train_time:13126ms step_avg:34.36ms step:383/1555 train_time:13157ms step_avg:34.35ms step:384/1555 train_time:13195ms step_avg:34.36ms step:385/1555 train_time:13226ms step_avg:34.35ms step:386/1555 train_time:13263ms step_avg:34.36ms step:387/1555 train_time:13294ms step_avg:34.35ms step:388/1555 train_time:13332ms step_avg:34.36ms step:389/1555 train_time:13363ms step_avg:34.35ms step:390/1555 train_time:13401ms step_avg:34.36ms step:391/1555 train_time:13432ms step_avg:34.35ms step:392/1555 train_time:13470ms step_avg:34.36ms step:393/1555 train_time:13501ms step_avg:34.35ms step:394/1555 train_time:13538ms step_avg:34.36ms step:395/1555 train_time:13569ms step_avg:34.35ms step:396/1555 train_time:13607ms step_avg:34.36ms step:397/1555 train_time:13638ms step_avg:34.35ms step:398/1555 train_time:13675ms step_avg:34.36ms step:399/1555 train_time:13706ms step_avg:34.35ms step:400/1555 train_time:13743ms step_avg:34.36ms step:401/1555 train_time:13774ms step_avg:34.35ms step:402/1555 train_time:13812ms step_avg:34.36ms step:403/1555 train_time:13843ms step_avg:34.35ms step:404/1555 train_time:13880ms step_avg:34.36ms step:405/1555 train_time:13912ms step_avg:34.35ms step:406/1555 train_time:13949ms step_avg:34.36ms step:407/1555 train_time:13980ms step_avg:34.35ms step:408/1555 train_time:14017ms step_avg:34.36ms step:409/1555 train_time:14048ms step_avg:34.35ms step:410/1555 train_time:14085ms step_avg:34.35ms step:411/1555 train_time:14116ms step_avg:34.35ms step:412/1555 train_time:14154ms step_avg:34.35ms step:413/1555 train_time:14185ms step_avg:34.35ms step:414/1555 train_time:14222ms step_avg:34.35ms step:415/1555 train_time:14254ms step_avg:34.35ms step:416/1555 train_time:14291ms step_avg:34.35ms step:417/1555 train_time:14323ms step_avg:34.35ms step:418/1555 train_time:14360ms step_avg:34.35ms step:419/1555 train_time:14391ms step_avg:34.35ms step:420/1555 train_time:14429ms step_avg:34.35ms step:421/1555 train_time:14460ms step_avg:34.35ms step:422/1555 train_time:14497ms step_avg:34.35ms step:423/1555 train_time:14528ms step_avg:34.35ms step:424/1555 train_time:14565ms step_avg:34.35ms step:425/1555 train_time:14596ms step_avg:34.34ms step:426/1555 train_time:14634ms step_avg:34.35ms step:427/1555 train_time:14664ms step_avg:34.34ms step:428/1555 train_time:14701ms step_avg:34.35ms step:429/1555 train_time:14733ms step_avg:34.34ms step:430/1555 train_time:14770ms step_avg:34.35ms step:431/1555 train_time:14801ms step_avg:34.34ms step:432/1555 train_time:14838ms step_avg:34.35ms step:433/1555 train_time:14870ms step_avg:34.34ms step:434/1555 train_time:14907ms step_avg:34.35ms step:435/1555 train_time:14939ms step_avg:34.34ms step:436/1555 train_time:14977ms step_avg:34.35ms step:437/1555 train_time:15007ms step_avg:34.34ms step:438/1555 train_time:15045ms step_avg:34.35ms step:439/1555 train_time:15076ms step_avg:34.34ms step:440/1555 train_time:15113ms step_avg:34.35ms step:441/1555 train_time:15145ms step_avg:34.34ms step:442/1555 train_time:15182ms step_avg:34.35ms step:443/1555 train_time:15213ms step_avg:34.34ms step:444/1555 train_time:15251ms step_avg:34.35ms step:445/1555 train_time:15282ms step_avg:34.34ms step:446/1555 train_time:15319ms step_avg:34.35ms step:447/1555 train_time:15350ms step_avg:34.34ms step:448/1555 train_time:15388ms step_avg:34.35ms step:449/1555 train_time:15419ms step_avg:34.34ms step:450/1555 train_time:15457ms step_avg:34.35ms step:451/1555 train_time:15488ms step_avg:34.34ms step:452/1555 train_time:15525ms step_avg:34.35ms step:453/1555 train_time:15556ms step_avg:34.34ms step:454/1555 train_time:15593ms step_avg:34.35ms step:455/1555 train_time:15624ms step_avg:34.34ms step:456/1555 train_time:15662ms step_avg:34.35ms step:457/1555 train_time:15692ms step_avg:34.34ms step:458/1555 train_time:15729ms step_avg:34.34ms step:459/1555 train_time:15760ms step_avg:34.34ms step:460/1555 train_time:15797ms step_avg:34.34ms step:461/1555 train_time:15828ms step_avg:34.33ms step:462/1555 train_time:15866ms step_avg:34.34ms step:463/1555 train_time:15897ms step_avg:34.33ms step:464/1555 train_time:15934ms step_avg:34.34ms step:465/1555 train_time:15965ms step_avg:34.33ms step:466/1555 train_time:16002ms step_avg:34.34ms step:467/1555 train_time:16033ms step_avg:34.33ms step:468/1555 train_time:16071ms step_avg:34.34ms step:469/1555 train_time:16102ms step_avg:34.33ms step:470/1555 train_time:16139ms step_avg:34.34ms step:471/1555 train_time:16170ms step_avg:34.33ms step:472/1555 train_time:16207ms step_avg:34.34ms step:473/1555 train_time:16238ms step_avg:34.33ms step:474/1555 train_time:16276ms step_avg:34.34ms step:475/1555 train_time:16307ms step_avg:34.33ms step:476/1555 train_time:16344ms step_avg:34.34ms step:477/1555 train_time:16375ms step_avg:34.33ms step:478/1555 train_time:16414ms step_avg:34.34ms step:479/1555 train_time:16444ms step_avg:34.33ms step:480/1555 train_time:16481ms step_avg:34.34ms step:481/1555 train_time:16513ms step_avg:34.33ms step:482/1555 train_time:16550ms step_avg:34.34ms step:483/1555 train_time:16581ms step_avg:34.33ms step:484/1555 train_time:16619ms step_avg:34.34ms step:485/1555 train_time:16650ms step_avg:34.33ms step:486/1555 train_time:16687ms step_avg:34.34ms step:487/1555 train_time:16719ms step_avg:34.33ms step:488/1555 train_time:16756ms step_avg:34.34ms step:489/1555 train_time:16787ms step_avg:34.33ms step:490/1555 train_time:16824ms step_avg:34.34ms step:491/1555 train_time:16856ms step_avg:34.33ms step:492/1555 train_time:16893ms step_avg:34.34ms step:493/1555 train_time:16925ms step_avg:34.33ms step:494/1555 train_time:16962ms step_avg:34.34ms step:495/1555 train_time:16993ms step_avg:34.33ms step:496/1555 train_time:17030ms step_avg:34.33ms step:497/1555 train_time:17061ms step_avg:34.33ms step:498/1555 train_time:17098ms step_avg:34.33ms step:499/1555 train_time:17130ms step_avg:34.33ms step:500/1555 train_time:17167ms step_avg:34.33ms step:500/1555 val_loss:4.2700 train_time:17216ms step_avg:34.43ms step:501/1555 train_time:17234ms step_avg:34.40ms step:502/1555 train_time:17252ms step_avg:34.37ms step:503/1555 train_time:17269ms step_avg:34.33ms step:504/1555 train_time:17306ms step_avg:34.34ms step:505/1555 train_time:17337ms step_avg:34.33ms step:506/1555 train_time:17380ms step_avg:34.35ms step:507/1555 train_time:17435ms step_avg:34.39ms step:508/1555 train_time:17500ms step_avg:34.45ms step:509/1555 train_time:17557ms step_avg:34.49ms step:510/1555 train_time:17621ms step_avg:34.55ms step:511/1555 train_time:17679ms step_avg:34.60ms step:512/1555 train_time:17743ms step_avg:34.65ms step:513/1555 train_time:17800ms step_avg:34.70ms step:514/1555 train_time:17864ms step_avg:34.76ms step:515/1555 train_time:17921ms step_avg:34.80ms step:516/1555 train_time:17985ms step_avg:34.85ms step:517/1555 train_time:18042ms step_avg:34.90ms step:518/1555 train_time:18106ms step_avg:34.95ms step:519/1555 train_time:18164ms step_avg:35.00ms step:520/1555 train_time:18230ms step_avg:35.06ms step:521/1555 train_time:18290ms step_avg:35.10ms step:522/1555 train_time:18355ms step_avg:35.16ms step:523/1555 train_time:18413ms step_avg:35.21ms step:524/1555 train_time:18477ms step_avg:35.26ms step:525/1555 train_time:18534ms step_avg:35.30ms step:526/1555 train_time:18599ms step_avg:35.36ms step:527/1555 train_time:18656ms step_avg:35.40ms step:528/1555 train_time:18720ms step_avg:35.46ms step:529/1555 train_time:18777ms step_avg:35.50ms step:530/1555 train_time:18842ms step_avg:35.55ms step:531/1555 train_time:18898ms step_avg:35.59ms step:532/1555 train_time:18962ms step_avg:35.64ms step:533/1555 train_time:19019ms step_avg:35.68ms step:534/1555 train_time:19083ms step_avg:35.74ms step:535/1555 train_time:19140ms step_avg:35.78ms step:536/1555 train_time:19205ms step_avg:35.83ms step:537/1555 train_time:19264ms step_avg:35.87ms step:538/1555 train_time:19330ms step_avg:35.93ms step:539/1555 train_time:19389ms step_avg:35.97ms step:540/1555 train_time:19453ms step_avg:36.02ms step:541/1555 train_time:19511ms step_avg:36.06ms step:542/1555 train_time:19575ms step_avg:36.12ms step:543/1555 train_time:19634ms step_avg:36.16ms step:544/1555 train_time:19697ms step_avg:36.21ms step:545/1555 train_time:19755ms step_avg:36.25ms step:546/1555 train_time:19818ms step_avg:36.30ms step:547/1555 train_time:19875ms step_avg:36.33ms step:548/1555 train_time:19939ms step_avg:36.39ms step:549/1555 train_time:19996ms step_avg:36.42ms step:550/1555 train_time:20059ms step_avg:36.47ms step:551/1555 train_time:20117ms step_avg:36.51ms step:552/1555 train_time:20181ms step_avg:36.56ms step:553/1555 train_time:20240ms step_avg:36.60ms step:554/1555 train_time:20305ms step_avg:36.65ms step:555/1555 train_time:20363ms step_avg:36.69ms step:556/1555 train_time:20429ms step_avg:36.74ms step:557/1555 train_time:20487ms step_avg:36.78ms step:558/1555 train_time:20552ms step_avg:36.83ms step:559/1555 train_time:20609ms step_avg:36.87ms step:560/1555 train_time:20673ms step_avg:36.92ms step:561/1555 train_time:20731ms step_avg:36.95ms step:562/1555 train_time:20795ms step_avg:37.00ms step:563/1555 train_time:20852ms step_avg:37.04ms step:564/1555 train_time:20917ms step_avg:37.09ms step:565/1555 train_time:20974ms step_avg:37.12ms step:566/1555 train_time:21038ms step_avg:37.17ms step:567/1555 train_time:21096ms step_avg:37.21ms step:568/1555 train_time:21160ms step_avg:37.25ms step:569/1555 train_time:21218ms step_avg:37.29ms step:570/1555 train_time:21281ms step_avg:37.34ms step:571/1555 train_time:21339ms step_avg:37.37ms step:572/1555 train_time:21405ms step_avg:37.42ms step:573/1555 train_time:21463ms step_avg:37.46ms step:574/1555 train_time:21528ms step_avg:37.50ms step:575/1555 train_time:21585ms step_avg:37.54ms step:576/1555 train_time:21650ms step_avg:37.59ms step:577/1555 train_time:21708ms step_avg:37.62ms step:578/1555 train_time:21772ms step_avg:37.67ms step:579/1555 train_time:21830ms step_avg:37.70ms step:580/1555 train_time:21894ms step_avg:37.75ms step:581/1555 train_time:21952ms step_avg:37.78ms step:582/1555 train_time:22016ms step_avg:37.83ms step:583/1555 train_time:22074ms step_avg:37.86ms step:584/1555 train_time:22139ms step_avg:37.91ms step:585/1555 train_time:22196ms step_avg:37.94ms step:586/1555 train_time:22260ms step_avg:37.99ms step:587/1555 train_time:22317ms step_avg:38.02ms step:588/1555 train_time:22382ms step_avg:38.06ms step:589/1555 train_time:22439ms step_avg:38.10ms step:590/1555 train_time:22503ms step_avg:38.14ms step:591/1555 train_time:22561ms step_avg:38.17ms step:592/1555 train_time:22627ms step_avg:38.22ms step:593/1555 train_time:22685ms step_avg:38.25ms step:594/1555 train_time:22749ms step_avg:38.30ms step:595/1555 train_time:22806ms step_avg:38.33ms step:596/1555 train_time:22871ms step_avg:38.37ms step:597/1555 train_time:22929ms step_avg:38.41ms step:598/1555 train_time:22994ms step_avg:38.45ms step:599/1555 train_time:23052ms step_avg:38.48ms step:600/1555 train_time:23117ms step_avg:38.53ms step:601/1555 train_time:23175ms step_avg:38.56ms step:602/1555 train_time:23238ms step_avg:38.60ms step:603/1555 train_time:23295ms step_avg:38.63ms step:604/1555 train_time:23359ms step_avg:38.67ms step:605/1555 train_time:23417ms step_avg:38.71ms step:606/1555 train_time:23481ms step_avg:38.75ms step:607/1555 train_time:23540ms step_avg:38.78ms step:608/1555 train_time:23605ms step_avg:38.82ms step:609/1555 train_time:23662ms step_avg:38.85ms step:610/1555 train_time:23727ms step_avg:38.90ms step:611/1555 train_time:23785ms step_avg:38.93ms step:612/1555 train_time:23849ms step_avg:38.97ms step:613/1555 train_time:23907ms step_avg:39.00ms step:614/1555 train_time:23972ms step_avg:39.04ms step:615/1555 train_time:24030ms step_avg:39.07ms step:616/1555 train_time:24094ms step_avg:39.11ms step:617/1555 train_time:24152ms step_avg:39.14ms step:618/1555 train_time:24217ms step_avg:39.19ms step:619/1555 train_time:24274ms step_avg:39.22ms step:620/1555 train_time:24339ms step_avg:39.26ms step:621/1555 train_time:24396ms step_avg:39.29ms step:622/1555 train_time:24460ms step_avg:39.32ms step:623/1555 train_time:24517ms step_avg:39.35ms step:624/1555 train_time:24582ms step_avg:39.39ms step:625/1555 train_time:24640ms step_avg:39.42ms step:626/1555 train_time:24704ms step_avg:39.46ms step:627/1555 train_time:24762ms step_avg:39.49ms step:628/1555 train_time:24827ms step_avg:39.53ms step:629/1555 train_time:24885ms step_avg:39.56ms step:630/1555 train_time:24950ms step_avg:39.60ms step:631/1555 train_time:25008ms step_avg:39.63ms step:632/1555 train_time:25072ms step_avg:39.67ms step:633/1555 train_time:25132ms step_avg:39.70ms step:634/1555 train_time:25195ms step_avg:39.74ms step:635/1555 train_time:25253ms step_avg:39.77ms step:636/1555 train_time:25317ms step_avg:39.81ms step:637/1555 train_time:25375ms step_avg:39.84ms step:638/1555 train_time:25439ms step_avg:39.87ms step:639/1555 train_time:25496ms step_avg:39.90ms step:640/1555 train_time:25560ms step_avg:39.94ms step:641/1555 train_time:25618ms step_avg:39.97ms step:642/1555 train_time:25682ms step_avg:40.00ms step:643/1555 train_time:25740ms step_avg:40.03ms step:644/1555 train_time:25805ms step_avg:40.07ms step:645/1555 train_time:25862ms step_avg:40.10ms step:646/1555 train_time:25927ms step_avg:40.14ms step:647/1555 train_time:25986ms step_avg:40.16ms step:648/1555 train_time:26050ms step_avg:40.20ms step:649/1555 train_time:26108ms step_avg:40.23ms step:650/1555 train_time:26173ms step_avg:40.27ms step:651/1555 train_time:26231ms step_avg:40.29ms step:652/1555 train_time:26295ms step_avg:40.33ms step:653/1555 train_time:26352ms step_avg:40.36ms step:654/1555 train_time:26416ms step_avg:40.39ms step:655/1555 train_time:26475ms step_avg:40.42ms step:656/1555 train_time:26538ms step_avg:40.45ms step:657/1555 train_time:26595ms step_avg:40.48ms step:658/1555 train_time:26660ms step_avg:40.52ms step:659/1555 train_time:26718ms step_avg:40.54ms step:660/1555 train_time:26782ms step_avg:40.58ms step:661/1555 train_time:26840ms step_avg:40.61ms step:662/1555 train_time:26905ms step_avg:40.64ms step:663/1555 train_time:26963ms step_avg:40.67ms step:664/1555 train_time:27029ms step_avg:40.71ms step:665/1555 train_time:27086ms step_avg:40.73ms step:666/1555 train_time:27151ms step_avg:40.77ms step:667/1555 train_time:27209ms step_avg:40.79ms step:668/1555 train_time:27273ms step_avg:40.83ms step:669/1555 train_time:27331ms step_avg:40.85ms step:670/1555 train_time:27395ms step_avg:40.89ms step:671/1555 train_time:27453ms step_avg:40.91ms step:672/1555 train_time:27517ms step_avg:40.95ms step:673/1555 train_time:27574ms step_avg:40.97ms step:674/1555 train_time:27638ms step_avg:41.01ms step:675/1555 train_time:27695ms step_avg:41.03ms step:676/1555 train_time:27758ms step_avg:41.06ms step:677/1555 train_time:27817ms step_avg:41.09ms step:678/1555 train_time:27882ms step_avg:41.12ms step:679/1555 train_time:27940ms step_avg:41.15ms step:680/1555 train_time:28005ms step_avg:41.18ms step:681/1555 train_time:28064ms step_avg:41.21ms step:682/1555 train_time:28129ms step_avg:41.24ms step:683/1555 train_time:28187ms step_avg:41.27ms step:684/1555 train_time:28251ms step_avg:41.30ms step:685/1555 train_time:28309ms step_avg:41.33ms step:686/1555 train_time:28373ms step_avg:41.36ms step:687/1555 train_time:28431ms step_avg:41.38ms step:688/1555 train_time:28495ms step_avg:41.42ms step:689/1555 train_time:28553ms step_avg:41.44ms step:690/1555 train_time:28616ms step_avg:41.47ms step:691/1555 train_time:28674ms step_avg:41.50ms step:692/1555 train_time:28739ms step_avg:41.53ms step:693/1555 train_time:28797ms step_avg:41.55ms step:694/1555 train_time:28862ms step_avg:41.59ms step:695/1555 train_time:28919ms step_avg:41.61ms step:696/1555 train_time:28984ms step_avg:41.64ms step:697/1555 train_time:29041ms step_avg:41.67ms step:698/1555 train_time:29106ms step_avg:41.70ms step:699/1555 train_time:29163ms step_avg:41.72ms step:700/1555 train_time:29228ms step_avg:41.75ms step:701/1555 train_time:29286ms step_avg:41.78ms step:702/1555 train_time:29350ms step_avg:41.81ms step:703/1555 train_time:29408ms step_avg:41.83ms step:704/1555 train_time:29473ms step_avg:41.87ms step:705/1555 train_time:29532ms step_avg:41.89ms step:706/1555 train_time:29595ms step_avg:41.92ms step:707/1555 train_time:29653ms step_avg:41.94ms step:708/1555 train_time:29717ms step_avg:41.97ms step:709/1555 train_time:29775ms step_avg:42.00ms step:710/1555 train_time:29839ms step_avg:42.03ms step:711/1555 train_time:29897ms step_avg:42.05ms step:712/1555 train_time:29960ms step_avg:42.08ms step:713/1555 train_time:30017ms step_avg:42.10ms step:714/1555 train_time:30083ms step_avg:42.13ms step:715/1555 train_time:30140ms step_avg:42.15ms step:716/1555 train_time:30204ms step_avg:42.18ms step:717/1555 train_time:30263ms step_avg:42.21ms step:718/1555 train_time:30327ms step_avg:42.24ms step:719/1555 train_time:30385ms step_avg:42.26ms step:720/1555 train_time:30449ms step_avg:42.29ms step:721/1555 train_time:30507ms step_avg:42.31ms step:722/1555 train_time:30572ms step_avg:42.34ms step:723/1555 train_time:30630ms step_avg:42.37ms step:724/1555 train_time:30694ms step_avg:42.40ms step:725/1555 train_time:30752ms step_avg:42.42ms step:726/1555 train_time:30817ms step_avg:42.45ms step:727/1555 train_time:30874ms step_avg:42.47ms step:728/1555 train_time:30938ms step_avg:42.50ms step:729/1555 train_time:30995ms step_avg:42.52ms step:730/1555 train_time:31060ms step_avg:42.55ms step:731/1555 train_time:31117ms step_avg:42.57ms step:732/1555 train_time:31182ms step_avg:42.60ms step:733/1555 train_time:31239ms step_avg:42.62ms step:734/1555 train_time:31304ms step_avg:42.65ms step:735/1555 train_time:31363ms step_avg:42.67ms step:736/1555 train_time:31429ms step_avg:42.70ms step:737/1555 train_time:31485ms step_avg:42.72ms step:738/1555 train_time:31550ms step_avg:42.75ms step:739/1555 train_time:31608ms step_avg:42.77ms step:740/1555 train_time:31672ms step_avg:42.80ms step:741/1555 train_time:31731ms step_avg:42.82ms step:742/1555 train_time:31795ms step_avg:42.85ms step:743/1555 train_time:31854ms step_avg:42.87ms step:744/1555 train_time:31917ms step_avg:42.90ms step:745/1555 train_time:31975ms step_avg:42.92ms step:746/1555 train_time:32040ms step_avg:42.95ms step:747/1555 train_time:32096ms step_avg:42.97ms step:748/1555 train_time:32161ms step_avg:43.00ms step:749/1555 train_time:32218ms step_avg:43.01ms step:750/1555 train_time:32282ms step_avg:43.04ms step:750/1555 val_loss:3.8685 train_time:32365ms step_avg:43.15ms step:751/1555 train_time:32388ms step_avg:43.13ms step:752/1555 train_time:32410ms step_avg:43.10ms step:753/1555 train_time:32466ms step_avg:43.12ms step:754/1555 train_time:32532ms step_avg:43.15ms step:755/1555 train_time:32592ms step_avg:43.17ms step:756/1555 train_time:32657ms step_avg:43.20ms step:757/1555 train_time:32713ms step_avg:43.21ms step:758/1555 train_time:32777ms step_avg:43.24ms step:759/1555 train_time:32834ms step_avg:43.26ms step:760/1555 train_time:32897ms step_avg:43.29ms step:761/1555 train_time:32954ms step_avg:43.30ms step:762/1555 train_time:33018ms step_avg:43.33ms step:763/1555 train_time:33074ms step_avg:43.35ms step:764/1555 train_time:33138ms step_avg:43.37ms step:765/1555 train_time:33195ms step_avg:43.39ms step:766/1555 train_time:33259ms step_avg:43.42ms step:767/1555 train_time:33317ms step_avg:43.44ms step:768/1555 train_time:33383ms step_avg:43.47ms step:769/1555 train_time:33441ms step_avg:43.49ms step:770/1555 train_time:33508ms step_avg:43.52ms step:771/1555 train_time:33567ms step_avg:43.54ms step:772/1555 train_time:33632ms step_avg:43.56ms step:773/1555 train_time:33690ms step_avg:43.58ms step:774/1555 train_time:33754ms step_avg:43.61ms step:775/1555 train_time:33811ms step_avg:43.63ms step:776/1555 train_time:33875ms step_avg:43.65ms step:777/1555 train_time:33932ms step_avg:43.67ms step:778/1555 train_time:33995ms step_avg:43.70ms step:779/1555 train_time:34053ms step_avg:43.71ms step:780/1555 train_time:34116ms step_avg:43.74ms step:781/1555 train_time:34173ms step_avg:43.76ms step:782/1555 train_time:34238ms step_avg:43.78ms step:783/1555 train_time:34296ms step_avg:43.80ms step:784/1555 train_time:34360ms step_avg:43.83ms step:785/1555 train_time:34418ms step_avg:43.84ms step:786/1555 train_time:34482ms step_avg:43.87ms step:787/1555 train_time:34540ms step_avg:43.89ms step:788/1555 train_time:34606ms step_avg:43.92ms step:789/1555 train_time:34664ms step_avg:43.93ms step:790/1555 train_time:34729ms step_avg:43.96ms step:791/1555 train_time:34787ms step_avg:43.98ms step:792/1555 train_time:34851ms step_avg:44.00ms step:793/1555 train_time:34908ms step_avg:44.02ms step:794/1555 train_time:34972ms step_avg:44.05ms step:795/1555 train_time:35031ms step_avg:44.06ms step:796/1555 train_time:35094ms step_avg:44.09ms step:797/1555 train_time:35152ms step_avg:44.11ms step:798/1555 train_time:35217ms step_avg:44.13ms step:799/1555 train_time:35274ms step_avg:44.15ms step:800/1555 train_time:35339ms step_avg:44.17ms step:801/1555 train_time:35398ms step_avg:44.19ms step:802/1555 train_time:35462ms step_avg:44.22ms step:803/1555 train_time:35520ms step_avg:44.23ms step:804/1555 train_time:35584ms step_avg:44.26ms step:805/1555 train_time:35641ms step_avg:44.27ms step:806/1555 train_time:35705ms step_avg:44.30ms step:807/1555 train_time:35763ms step_avg:44.32ms step:808/1555 train_time:35827ms step_avg:44.34ms step:809/1555 train_time:35885ms step_avg:44.36ms step:810/1555 train_time:35950ms step_avg:44.38ms step:811/1555 train_time:36007ms step_avg:44.40ms step:812/1555 train_time:36072ms step_avg:44.42ms step:813/1555 train_time:36129ms step_avg:44.44ms step:814/1555 train_time:36193ms step_avg:44.46ms step:815/1555 train_time:36252ms step_avg:44.48ms step:816/1555 train_time:36317ms step_avg:44.51ms step:817/1555 train_time:36375ms step_avg:44.52ms step:818/1555 train_time:36439ms step_avg:44.55ms step:819/1555 train_time:36498ms step_avg:44.56ms step:820/1555 train_time:36562ms step_avg:44.59ms step:821/1555 train_time:36619ms step_avg:44.60ms step:822/1555 train_time:36683ms step_avg:44.63ms step:823/1555 train_time:36740ms step_avg:44.64ms step:824/1555 train_time:36805ms step_avg:44.67ms step:825/1555 train_time:36862ms step_avg:44.68ms step:826/1555 train_time:36927ms step_avg:44.71ms step:827/1555 train_time:36984ms step_avg:44.72ms step:828/1555 train_time:37050ms step_avg:44.75ms step:829/1555 train_time:37107ms step_avg:44.76ms step:830/1555 train_time:37171ms step_avg:44.78ms step:831/1555 train_time:37229ms step_avg:44.80ms step:832/1555 train_time:37295ms step_avg:44.83ms step:833/1555 train_time:37353ms step_avg:44.84ms step:834/1555 train_time:37417ms step_avg:44.86ms step:835/1555 train_time:37475ms step_avg:44.88ms step:836/1555 train_time:37539ms step_avg:44.90ms step:837/1555 train_time:37597ms step_avg:44.92ms step:838/1555 train_time:37661ms step_avg:44.94ms step:839/1555 train_time:37718ms step_avg:44.96ms step:840/1555 train_time:37782ms step_avg:44.98ms step:841/1555 train_time:37840ms step_avg:44.99ms step:842/1555 train_time:37903ms step_avg:45.02ms step:843/1555 train_time:37961ms step_avg:45.03ms step:844/1555 train_time:38027ms step_avg:45.06ms step:845/1555 train_time:38085ms step_avg:45.07ms step:846/1555 train_time:38150ms step_avg:45.10ms step:847/1555 train_time:38209ms step_avg:45.11ms step:848/1555 train_time:38273ms step_avg:45.13ms step:849/1555 train_time:38331ms step_avg:45.15ms step:850/1555 train_time:38395ms step_avg:45.17ms step:851/1555 train_time:38453ms step_avg:45.19ms step:852/1555 train_time:38518ms step_avg:45.21ms step:853/1555 train_time:38575ms step_avg:45.22ms step:854/1555 train_time:38640ms step_avg:45.25ms step:855/1555 train_time:38698ms step_avg:45.26ms step:856/1555 train_time:38761ms step_avg:45.28ms step:857/1555 train_time:38819ms step_avg:45.30ms step:858/1555 train_time:38882ms step_avg:45.32ms step:859/1555 train_time:38940ms step_avg:45.33ms step:860/1555 train_time:39005ms step_avg:45.35ms step:861/1555 train_time:39062ms step_avg:45.37ms step:862/1555 train_time:39128ms step_avg:45.39ms step:863/1555 train_time:39185ms step_avg:45.41ms step:864/1555 train_time:39250ms step_avg:45.43ms step:865/1555 train_time:39307ms step_avg:45.44ms step:866/1555 train_time:39372ms step_avg:45.46ms step:867/1555 train_time:39430ms step_avg:45.48ms step:868/1555 train_time:39494ms step_avg:45.50ms step:869/1555 train_time:39552ms step_avg:45.51ms step:870/1555 train_time:39617ms step_avg:45.54ms step:871/1555 train_time:39676ms step_avg:45.55ms step:872/1555 train_time:39740ms step_avg:45.57ms step:873/1555 train_time:39797ms step_avg:45.59ms step:874/1555 train_time:39861ms step_avg:45.61ms step:875/1555 train_time:39919ms step_avg:45.62ms step:876/1555 train_time:39983ms step_avg:45.64ms step:877/1555 train_time:40040ms step_avg:45.66ms step:878/1555 train_time:40104ms step_avg:45.68ms step:879/1555 train_time:40162ms step_avg:45.69ms step:880/1555 train_time:40226ms step_avg:45.71ms step:881/1555 train_time:40284ms step_avg:45.73ms step:882/1555 train_time:40349ms step_avg:45.75ms step:883/1555 train_time:40406ms step_avg:45.76ms step:884/1555 train_time:40471ms step_avg:45.78ms step:885/1555 train_time:40529ms step_avg:45.80ms step:886/1555 train_time:40593ms step_avg:45.82ms step:887/1555 train_time:40652ms step_avg:45.83ms step:888/1555 train_time:40716ms step_avg:45.85ms step:889/1555 train_time:40774ms step_avg:45.86ms step:890/1555 train_time:40838ms step_avg:45.89ms step:891/1555 train_time:40896ms step_avg:45.90ms step:892/1555 train_time:40960ms step_avg:45.92ms step:893/1555 train_time:41017ms step_avg:45.93ms step:894/1555 train_time:41081ms step_avg:45.95ms step:895/1555 train_time:41139ms step_avg:45.97ms step:896/1555 train_time:41203ms step_avg:45.99ms step:897/1555 train_time:41261ms step_avg:46.00ms step:898/1555 train_time:41325ms step_avg:46.02ms step:899/1555 train_time:41383ms step_avg:46.03ms step:900/1555 train_time:41448ms step_avg:46.05ms step:901/1555 train_time:41505ms step_avg:46.07ms step:902/1555 train_time:41570ms step_avg:46.09ms step:903/1555 train_time:41629ms step_avg:46.10ms step:904/1555 train_time:41693ms step_avg:46.12ms step:905/1555 train_time:41752ms step_avg:46.13ms step:906/1555 train_time:41817ms step_avg:46.16ms step:907/1555 train_time:41874ms step_avg:46.17ms step:908/1555 train_time:41939ms step_avg:46.19ms step:909/1555 train_time:41996ms step_avg:46.20ms step:910/1555 train_time:42060ms step_avg:46.22ms step:911/1555 train_time:42118ms step_avg:46.23ms step:912/1555 train_time:42181ms step_avg:46.25ms step:913/1555 train_time:42239ms step_avg:46.26ms step:914/1555 train_time:42304ms step_avg:46.28ms step:915/1555 train_time:42360ms step_avg:46.30ms step:916/1555 train_time:42425ms step_avg:46.32ms step:917/1555 train_time:42482ms step_avg:46.33ms step:918/1555 train_time:42548ms step_avg:46.35ms step:919/1555 train_time:42607ms step_avg:46.36ms step:920/1555 train_time:42672ms step_avg:46.38ms step:921/1555 train_time:42730ms step_avg:46.39ms step:922/1555 train_time:42794ms step_avg:46.41ms step:923/1555 train_time:42852ms step_avg:46.43ms step:924/1555 train_time:42916ms step_avg:46.45ms step:925/1555 train_time:42974ms step_avg:46.46ms step:926/1555 train_time:43038ms step_avg:46.48ms step:927/1555 train_time:43096ms step_avg:46.49ms step:928/1555 train_time:43159ms step_avg:46.51ms step:929/1555 train_time:43217ms step_avg:46.52ms step:930/1555 train_time:43281ms step_avg:46.54ms step:931/1555 train_time:43339ms step_avg:46.55ms step:932/1555 train_time:43403ms step_avg:46.57ms step:933/1555 train_time:43461ms step_avg:46.58ms step:934/1555 train_time:43525ms step_avg:46.60ms step:935/1555 train_time:43584ms step_avg:46.61ms step:936/1555 train_time:43648ms step_avg:46.63ms step:937/1555 train_time:43705ms step_avg:46.64ms step:938/1555 train_time:43770ms step_avg:46.66ms step:939/1555 train_time:43828ms step_avg:46.68ms step:940/1555 train_time:43893ms step_avg:46.69ms step:941/1555 train_time:43952ms step_avg:46.71ms step:942/1555 train_time:44015ms step_avg:46.73ms step:943/1555 train_time:44075ms step_avg:46.74ms step:944/1555 train_time:44137ms step_avg:46.76ms step:945/1555 train_time:44194ms step_avg:46.77ms step:946/1555 train_time:44259ms step_avg:46.79ms step:947/1555 train_time:44317ms step_avg:46.80ms step:948/1555 train_time:44380ms step_avg:46.81ms step:949/1555 train_time:44438ms step_avg:46.83ms step:950/1555 train_time:44502ms step_avg:46.84ms step:951/1555 train_time:44559ms step_avg:46.85ms step:952/1555 train_time:44624ms step_avg:46.87ms step:953/1555 train_time:44680ms step_avg:46.88ms step:954/1555 train_time:44746ms step_avg:46.90ms step:955/1555 train_time:44803ms step_avg:46.91ms step:956/1555 train_time:44869ms step_avg:46.93ms step:957/1555 train_time:44927ms step_avg:46.95ms step:958/1555 train_time:44991ms step_avg:46.96ms step:959/1555 train_time:45050ms step_avg:46.98ms step:960/1555 train_time:45114ms step_avg:46.99ms step:961/1555 train_time:45172ms step_avg:47.01ms step:962/1555 train_time:45238ms step_avg:47.02ms step:963/1555 train_time:45295ms step_avg:47.04ms step:964/1555 train_time:45359ms step_avg:47.05ms step:965/1555 train_time:45417ms step_avg:47.06ms step:966/1555 train_time:45480ms step_avg:47.08ms step:967/1555 train_time:45538ms step_avg:47.09ms step:968/1555 train_time:45602ms step_avg:47.11ms step:969/1555 train_time:45659ms step_avg:47.12ms step:970/1555 train_time:45723ms step_avg:47.14ms step:971/1555 train_time:45781ms step_avg:47.15ms step:972/1555 train_time:45846ms step_avg:47.17ms step:973/1555 train_time:45904ms step_avg:47.18ms step:974/1555 train_time:45970ms step_avg:47.20ms step:975/1555 train_time:46028ms step_avg:47.21ms step:976/1555 train_time:46092ms step_avg:47.23ms step:977/1555 train_time:46150ms step_avg:47.24ms step:978/1555 train_time:46215ms step_avg:47.25ms step:979/1555 train_time:46273ms step_avg:47.27ms step:980/1555 train_time:46337ms step_avg:47.28ms step:981/1555 train_time:46395ms step_avg:47.29ms step:982/1555 train_time:46460ms step_avg:47.31ms step:983/1555 train_time:46517ms step_avg:47.32ms step:984/1555 train_time:46580ms step_avg:47.34ms step:985/1555 train_time:46638ms step_avg:47.35ms step:986/1555 train_time:46702ms step_avg:47.37ms step:987/1555 train_time:46760ms step_avg:47.38ms step:988/1555 train_time:46823ms step_avg:47.39ms step:989/1555 train_time:46881ms step_avg:47.40ms step:990/1555 train_time:46945ms step_avg:47.42ms step:991/1555 train_time:47003ms step_avg:47.43ms step:992/1555 train_time:47068ms step_avg:47.45ms step:993/1555 train_time:47126ms step_avg:47.46ms step:994/1555 train_time:47191ms step_avg:47.48ms step:995/1555 train_time:47250ms step_avg:47.49ms step:996/1555 train_time:47314ms step_avg:47.50ms step:997/1555 train_time:47371ms step_avg:47.51ms step:998/1555 train_time:47436ms step_avg:47.53ms step:999/1555 train_time:47494ms step_avg:47.54ms step:1000/1555 train_time:47558ms step_avg:47.56ms step:1000/1555 val_loss:3.5683 train_time:47640ms step_avg:47.64ms step:1001/1555 train_time:47658ms step_avg:47.61ms step:1002/1555 train_time:47681ms step_avg:47.59ms step:1003/1555 train_time:47738ms step_avg:47.60ms step:1004/1555 train_time:47807ms step_avg:47.62ms step:1005/1555 train_time:47866ms step_avg:47.63ms step:1006/1555 train_time:47931ms step_avg:47.64ms step:1007/1555 train_time:47989ms step_avg:47.66ms step:1008/1555 train_time:48052ms step_avg:47.67ms step:1009/1555 train_time:48111ms step_avg:47.68ms step:1010/1555 train_time:48174ms step_avg:47.70ms step:1011/1555 train_time:48235ms step_avg:47.71ms step:1012/1555 train_time:48320ms step_avg:47.75ms step:1013/1555 train_time:48404ms step_avg:47.78ms step:1014/1555 train_time:48493ms step_avg:47.82ms step:1015/1555 train_time:48577ms step_avg:47.86ms step:1016/1555 train_time:48668ms step_avg:47.90ms step:1017/1555 train_time:48754ms step_avg:47.94ms step:1018/1555 train_time:48847ms step_avg:47.98ms step:1019/1555 train_time:48931ms step_avg:48.02ms step:1020/1555 train_time:49021ms step_avg:48.06ms step:1021/1555 train_time:49105ms step_avg:48.09ms step:1022/1555 train_time:49194ms step_avg:48.13ms step:1023/1555 train_time:49278ms step_avg:48.17ms step:1024/1555 train_time:49367ms step_avg:48.21ms step:1025/1555 train_time:49449ms step_avg:48.24ms step:1026/1555 train_time:49540ms step_avg:48.28ms step:1027/1555 train_time:49625ms step_avg:48.32ms step:1028/1555 train_time:49716ms step_avg:48.36ms step:1029/1555 train_time:49802ms step_avg:48.40ms step:1030/1555 train_time:49891ms step_avg:48.44ms step:1031/1555 train_time:49977ms step_avg:48.47ms step:1032/1555 train_time:50067ms step_avg:48.51ms step:1033/1555 train_time:50151ms step_avg:48.55ms step:1034/1555 train_time:50242ms step_avg:48.59ms step:1035/1555 train_time:50325ms step_avg:48.62ms step:1036/1555 train_time:50414ms step_avg:48.66ms step:1037/1555 train_time:50498ms step_avg:48.70ms step:1038/1555 train_time:50587ms step_avg:48.74ms step:1039/1555 train_time:50672ms step_avg:48.77ms step:1040/1555 train_time:50763ms step_avg:48.81ms step:1041/1555 train_time:50848ms step_avg:48.85ms step:1042/1555 train_time:50939ms step_avg:48.89ms step:1043/1555 train_time:51023ms step_avg:48.92ms step:1044/1555 train_time:51112ms step_avg:48.96ms step:1045/1555 train_time:51196ms step_avg:48.99ms step:1046/1555 train_time:51287ms step_avg:49.03ms step:1047/1555 train_time:51370ms step_avg:49.06ms step:1048/1555 train_time:51460ms step_avg:49.10ms step:1049/1555 train_time:51543ms step_avg:49.14ms step:1050/1555 train_time:51632ms step_avg:49.17ms step:1051/1555 train_time:51718ms step_avg:49.21ms step:1052/1555 train_time:51808ms step_avg:49.25ms step:1053/1555 train_time:51892ms step_avg:49.28ms step:1054/1555 train_time:51983ms step_avg:49.32ms step:1055/1555 train_time:52068ms step_avg:49.35ms step:1056/1555 train_time:52157ms step_avg:49.39ms step:1057/1555 train_time:52242ms step_avg:49.42ms step:1058/1555 train_time:52330ms step_avg:49.46ms step:1059/1555 train_time:52414ms step_avg:49.49ms step:1060/1555 train_time:52505ms step_avg:49.53ms step:1061/1555 train_time:52588ms step_avg:49.56ms step:1062/1555 train_time:52678ms step_avg:49.60ms step:1063/1555 train_time:52762ms step_avg:49.63ms step:1064/1555 train_time:52851ms step_avg:49.67ms step:1065/1555 train_time:52938ms step_avg:49.71ms step:1066/1555 train_time:53029ms step_avg:49.75ms step:1067/1555 train_time:53113ms step_avg:49.78ms step:1068/1555 train_time:53204ms step_avg:49.82ms step:1069/1555 train_time:53287ms step_avg:49.85ms step:1070/1555 train_time:53377ms step_avg:49.89ms step:1071/1555 train_time:53461ms step_avg:49.92ms step:1072/1555 train_time:53551ms step_avg:49.95ms step:1073/1555 train_time:53635ms step_avg:49.99ms step:1074/1555 train_time:53725ms step_avg:50.02ms step:1075/1555 train_time:53809ms step_avg:50.05ms step:1076/1555 train_time:53901ms step_avg:50.09ms step:1077/1555 train_time:53985ms step_avg:50.12ms step:1078/1555 train_time:54076ms step_avg:50.16ms step:1079/1555 train_time:54160ms step_avg:50.19ms step:1080/1555 train_time:54248ms step_avg:50.23ms step:1081/1555 train_time:54333ms step_avg:50.26ms step:1082/1555 train_time:54423ms step_avg:50.30ms step:1083/1555 train_time:54507ms step_avg:50.33ms step:1084/1555 train_time:54597ms step_avg:50.37ms step:1085/1555 train_time:54681ms step_avg:50.40ms step:1086/1555 train_time:54771ms step_avg:50.43ms step:1087/1555 train_time:54857ms step_avg:50.47ms step:1088/1555 train_time:54947ms step_avg:50.50ms step:1089/1555 train_time:55031ms step_avg:50.53ms step:1090/1555 train_time:55123ms step_avg:50.57ms step:1091/1555 train_time:55206ms step_avg:50.60ms step:1092/1555 train_time:55296ms step_avg:50.64ms step:1093/1555 train_time:55380ms step_avg:50.67ms step:1094/1555 train_time:55469ms step_avg:50.70ms step:1095/1555 train_time:55553ms step_avg:50.73ms step:1096/1555 train_time:55644ms step_avg:50.77ms step:1097/1555 train_time:55727ms step_avg:50.80ms step:1098/1555 train_time:55817ms step_avg:50.83ms step:1099/1555 train_time:55901ms step_avg:50.87ms step:1100/1555 train_time:55990ms step_avg:50.90ms step:1101/1555 train_time:56075ms step_avg:50.93ms step:1102/1555 train_time:56165ms step_avg:50.97ms step:1103/1555 train_time:56249ms step_avg:51.00ms step:1104/1555 train_time:56340ms step_avg:51.03ms step:1105/1555 train_time:56424ms step_avg:51.06ms step:1106/1555 train_time:56514ms step_avg:51.10ms step:1107/1555 train_time:56598ms step_avg:51.13ms step:1108/1555 train_time:56688ms step_avg:51.16ms step:1109/1555 train_time:56772ms step_avg:51.19ms step:1110/1555 train_time:56863ms step_avg:51.23ms step:1111/1555 train_time:56946ms step_avg:51.26ms step:1112/1555 train_time:57037ms step_avg:51.29ms step:1113/1555 train_time:57121ms step_avg:51.32ms step:1114/1555 train_time:57211ms step_avg:51.36ms step:1115/1555 train_time:57296ms step_avg:51.39ms step:1116/1555 train_time:57386ms step_avg:51.42ms step:1117/1555 train_time:57470ms step_avg:51.45ms step:1118/1555 train_time:57559ms step_avg:51.48ms step:1119/1555 train_time:57643ms step_avg:51.51ms step:1120/1555 train_time:57733ms step_avg:51.55ms step:1121/1555 train_time:57819ms step_avg:51.58ms step:1122/1555 train_time:57908ms step_avg:51.61ms step:1123/1555 train_time:57992ms step_avg:51.64ms step:1124/1555 train_time:58082ms step_avg:51.67ms step:1125/1555 train_time:58165ms step_avg:51.70ms step:1126/1555 train_time:58256ms step_avg:51.74ms step:1127/1555 train_time:58341ms step_avg:51.77ms step:1128/1555 train_time:58430ms step_avg:51.80ms step:1129/1555 train_time:58514ms step_avg:51.83ms step:1130/1555 train_time:58605ms step_avg:51.86ms step:1131/1555 train_time:58688ms step_avg:51.89ms step:1132/1555 train_time:58777ms step_avg:51.92ms step:1133/1555 train_time:58862ms step_avg:51.95ms step:1134/1555 train_time:58952ms step_avg:51.99ms step:1135/1555 train_time:59036ms step_avg:52.01ms step:1136/1555 train_time:59126ms step_avg:52.05ms step:1137/1555 train_time:59211ms step_avg:52.08ms step:1138/1555 train_time:59302ms step_avg:52.11ms step:1139/1555 train_time:59385ms step_avg:52.14ms step:1140/1555 train_time:59475ms step_avg:52.17ms step:1141/1555 train_time:59559ms step_avg:52.20ms step:1142/1555 train_time:59649ms step_avg:52.23ms step:1143/1555 train_time:59733ms step_avg:52.26ms step:1144/1555 train_time:59825ms step_avg:52.29ms step:1145/1555 train_time:59909ms step_avg:52.32ms step:1146/1555 train_time:60000ms step_avg:52.36ms step:1147/1555 train_time:60084ms step_avg:52.38ms step:1148/1555 train_time:60172ms step_avg:52.41ms step:1149/1555 train_time:60257ms step_avg:52.44ms step:1150/1555 train_time:60348ms step_avg:52.48ms step:1151/1555 train_time:60430ms step_avg:52.50ms step:1152/1555 train_time:60522ms step_avg:52.54ms step:1153/1555 train_time:60605ms step_avg:52.56ms step:1154/1555 train_time:60694ms step_avg:52.59ms step:1155/1555 train_time:60779ms step_avg:52.62ms step:1156/1555 train_time:60868ms step_avg:52.65ms step:1157/1555 train_time:60952ms step_avg:52.68ms step:1158/1555 train_time:61044ms step_avg:52.72ms step:1159/1555 train_time:61127ms step_avg:52.74ms step:1160/1555 train_time:61217ms step_avg:52.77ms step:1161/1555 train_time:61301ms step_avg:52.80ms step:1162/1555 train_time:61391ms step_avg:52.83ms step:1163/1555 train_time:61475ms step_avg:52.86ms step:1164/1555 train_time:61566ms step_avg:52.89ms step:1165/1555 train_time:61649ms step_avg:52.92ms step:1166/1555 train_time:61740ms step_avg:52.95ms step:1167/1555 train_time:61824ms step_avg:52.98ms step:1168/1555 train_time:61915ms step_avg:53.01ms step:1169/1555 train_time:61999ms step_avg:53.04ms step:1170/1555 train_time:62088ms step_avg:53.07ms step:1171/1555 train_time:62173ms step_avg:53.09ms step:1172/1555 train_time:62263ms step_avg:53.13ms step:1173/1555 train_time:62347ms step_avg:53.15ms step:1174/1555 train_time:62437ms step_avg:53.18ms step:1175/1555 train_time:62521ms step_avg:53.21ms step:1176/1555 train_time:62610ms step_avg:53.24ms step:1177/1555 train_time:62694ms step_avg:53.27ms step:1178/1555 train_time:62785ms step_avg:53.30ms step:1179/1555 train_time:62869ms step_avg:53.32ms step:1180/1555 train_time:62958ms step_avg:53.35ms step:1181/1555 train_time:63043ms step_avg:53.38ms step:1182/1555 train_time:63132ms step_avg:53.41ms step:1183/1555 train_time:63216ms step_avg:53.44ms step:1184/1555 train_time:63306ms step_avg:53.47ms step:1185/1555 train_time:63389ms step_avg:53.49ms step:1186/1555 train_time:63480ms step_avg:53.52ms step:1187/1555 train_time:63564ms step_avg:53.55ms step:1188/1555 train_time:63654ms step_avg:53.58ms step:1189/1555 train_time:63738ms step_avg:53.61ms step:1190/1555 train_time:63828ms step_avg:53.64ms step:1191/1555 train_time:63912ms step_avg:53.66ms step:1192/1555 train_time:64004ms step_avg:53.69ms step:1193/1555 train_time:64088ms step_avg:53.72ms step:1194/1555 train_time:64179ms step_avg:53.75ms step:1195/1555 train_time:64263ms step_avg:53.78ms step:1196/1555 train_time:64352ms step_avg:53.81ms step:1197/1555 train_time:64437ms step_avg:53.83ms step:1198/1555 train_time:64527ms step_avg:53.86ms step:1199/1555 train_time:64611ms step_avg:53.89ms step:1200/1555 train_time:64703ms step_avg:53.92ms step:1201/1555 train_time:64786ms step_avg:53.94ms step:1202/1555 train_time:64876ms step_avg:53.97ms step:1203/1555 train_time:64961ms step_avg:54.00ms step:1204/1555 train_time:65049ms step_avg:54.03ms step:1205/1555 train_time:65133ms step_avg:54.05ms step:1206/1555 train_time:65225ms step_avg:54.08ms step:1207/1555 train_time:65308ms step_avg:54.11ms step:1208/1555 train_time:65400ms step_avg:54.14ms step:1209/1555 train_time:65483ms step_avg:54.16ms step:1210/1555 train_time:65572ms step_avg:54.19ms step:1211/1555 train_time:65657ms step_avg:54.22ms step:1212/1555 train_time:65748ms step_avg:54.25ms step:1213/1555 train_time:65832ms step_avg:54.27ms step:1214/1555 train_time:65923ms step_avg:54.30ms step:1215/1555 train_time:66006ms step_avg:54.33ms step:1216/1555 train_time:66095ms step_avg:54.35ms step:1217/1555 train_time:66179ms step_avg:54.38ms step:1218/1555 train_time:66269ms step_avg:54.41ms step:1219/1555 train_time:66353ms step_avg:54.43ms step:1220/1555 train_time:66444ms step_avg:54.46ms step:1221/1555 train_time:66528ms step_avg:54.49ms step:1222/1555 train_time:66618ms step_avg:54.52ms step:1223/1555 train_time:66702ms step_avg:54.54ms step:1224/1555 train_time:66794ms step_avg:54.57ms step:1225/1555 train_time:66876ms step_avg:54.59ms step:1226/1555 train_time:66967ms step_avg:54.62ms step:1227/1555 train_time:67051ms step_avg:54.65ms step:1228/1555 train_time:67142ms step_avg:54.68ms step:1229/1555 train_time:67226ms step_avg:54.70ms step:1230/1555 train_time:67316ms step_avg:54.73ms step:1231/1555 train_time:67401ms step_avg:54.75ms step:1232/1555 train_time:67490ms step_avg:54.78ms step:1233/1555 train_time:67574ms step_avg:54.80ms step:1234/1555 train_time:67665ms step_avg:54.83ms step:1235/1555 train_time:67749ms step_avg:54.86ms step:1236/1555 train_time:67842ms step_avg:54.89ms step:1237/1555 train_time:67925ms step_avg:54.91ms step:1238/1555 train_time:68014ms step_avg:54.94ms step:1239/1555 train_time:68101ms step_avg:54.96ms step:1240/1555 train_time:68190ms step_avg:54.99ms step:1241/1555 train_time:68273ms step_avg:55.01ms step:1242/1555 train_time:68365ms step_avg:55.04ms step:1243/1555 train_time:68448ms step_avg:55.07ms step:1244/1555 train_time:68539ms step_avg:55.10ms step:1245/1555 train_time:68625ms step_avg:55.12ms step:1246/1555 train_time:68714ms step_avg:55.15ms step:1247/1555 train_time:68798ms step_avg:55.17ms step:1248/1555 train_time:68888ms step_avg:55.20ms step:1249/1555 train_time:68972ms step_avg:55.22ms step:1250/1555 train_time:69063ms step_avg:55.25ms step:1250/1555 val_loss:3.3959 train_time:69177ms step_avg:55.34ms step:1251/1555 train_time:69195ms step_avg:55.31ms step:1252/1555 train_time:69237ms step_avg:55.30ms step:1253/1555 train_time:69325ms step_avg:55.33ms step:1254/1555 train_time:69417ms step_avg:55.36ms step:1255/1555 train_time:69502ms step_avg:55.38ms step:1256/1555 train_time:69591ms step_avg:55.41ms step:1257/1555 train_time:69674ms step_avg:55.43ms step:1258/1555 train_time:69762ms step_avg:55.45ms step:1259/1555 train_time:69847ms step_avg:55.48ms step:1260/1555 train_time:69935ms step_avg:55.50ms step:1261/1555 train_time:70018ms step_avg:55.53ms step:1262/1555 train_time:70109ms step_avg:55.55ms step:1263/1555 train_time:70195ms step_avg:55.58ms step:1264/1555 train_time:70288ms step_avg:55.61ms step:1265/1555 train_time:70374ms step_avg:55.63ms step:1266/1555 train_time:70463ms step_avg:55.66ms step:1267/1555 train_time:70548ms step_avg:55.68ms step:1268/1555 train_time:70638ms step_avg:55.71ms step:1269/1555 train_time:70721ms step_avg:55.73ms step:1270/1555 train_time:70812ms step_avg:55.76ms step:1271/1555 train_time:70894ms step_avg:55.78ms step:1272/1555 train_time:70983ms step_avg:55.80ms step:1273/1555 train_time:71067ms step_avg:55.83ms step:1274/1555 train_time:71158ms step_avg:55.85ms step:1275/1555 train_time:71243ms step_avg:55.88ms step:1276/1555 train_time:71335ms step_avg:55.90ms step:1277/1555 train_time:71418ms step_avg:55.93ms step:1278/1555 train_time:71510ms step_avg:55.95ms step:1279/1555 train_time:71593ms step_avg:55.98ms step:1280/1555 train_time:71682ms step_avg:56.00ms step:1281/1555 train_time:71766ms step_avg:56.02ms step:1282/1555 train_time:71856ms step_avg:56.05ms step:1283/1555 train_time:71939ms step_avg:56.07ms step:1284/1555 train_time:72029ms step_avg:56.10ms step:1285/1555 train_time:72114ms step_avg:56.12ms step:1286/1555 train_time:72204ms step_avg:56.15ms step:1287/1555 train_time:72289ms step_avg:56.17ms step:1288/1555 train_time:72380ms step_avg:56.20ms step:1289/1555 train_time:72465ms step_avg:56.22ms step:1290/1555 train_time:72556ms step_avg:56.24ms step:1291/1555 train_time:72640ms step_avg:56.27ms step:1292/1555 train_time:72730ms step_avg:56.29ms step:1293/1555 train_time:72814ms step_avg:56.31ms step:1294/1555 train_time:72903ms step_avg:56.34ms step:1295/1555 train_time:72987ms step_avg:56.36ms step:1296/1555 train_time:73078ms step_avg:56.39ms step:1297/1555 train_time:73161ms step_avg:56.41ms step:1298/1555 train_time:73253ms step_avg:56.44ms step:1299/1555 train_time:73337ms step_avg:56.46ms step:1300/1555 train_time:73427ms step_avg:56.48ms step:1301/1555 train_time:73512ms step_avg:56.50ms step:1302/1555 train_time:73601ms step_avg:56.53ms step:1303/1555 train_time:73686ms step_avg:56.55ms step:1304/1555 train_time:73776ms step_avg:56.58ms step:1305/1555 train_time:73859ms step_avg:56.60ms step:1306/1555 train_time:73949ms step_avg:56.62ms step:1307/1555 train_time:74033ms step_avg:56.64ms step:1308/1555 train_time:74121ms step_avg:56.67ms step:1309/1555 train_time:74207ms step_avg:56.69ms step:1310/1555 train_time:74296ms step_avg:56.71ms step:1311/1555 train_time:74379ms step_avg:56.73ms step:1312/1555 train_time:74471ms step_avg:56.76ms step:1313/1555 train_time:74555ms step_avg:56.78ms step:1314/1555 train_time:74645ms step_avg:56.81ms step:1315/1555 train_time:74729ms step_avg:56.83ms step:1316/1555 train_time:74819ms step_avg:56.85ms step:1317/1555 train_time:74903ms step_avg:56.87ms step:1318/1555 train_time:74993ms step_avg:56.90ms step:1319/1555 train_time:75076ms step_avg:56.92ms step:1320/1555 train_time:75167ms step_avg:56.94ms step:1321/1555 train_time:75252ms step_avg:56.97ms step:1322/1555 train_time:75340ms step_avg:56.99ms step:1323/1555 train_time:75425ms step_avg:57.01ms step:1324/1555 train_time:75516ms step_avg:57.04ms step:1325/1555 train_time:75599ms step_avg:57.06ms step:1326/1555 train_time:75691ms step_avg:57.08ms step:1327/1555 train_time:75775ms step_avg:57.10ms step:1328/1555 train_time:75864ms step_avg:57.13ms step:1329/1555 train_time:75948ms step_avg:57.15ms step:1330/1555 train_time:76038ms step_avg:57.17ms step:1331/1555 train_time:76121ms step_avg:57.19ms step:1332/1555 train_time:76213ms step_avg:57.22ms step:1333/1555 train_time:76296ms step_avg:57.24ms step:1334/1555 train_time:76385ms step_avg:57.26ms step:1335/1555 train_time:76470ms step_avg:57.28ms step:1336/1555 train_time:76559ms step_avg:57.30ms step:1337/1555 train_time:76643ms step_avg:57.32ms step:1338/1555 train_time:76735ms step_avg:57.35ms step:1339/1555 train_time:76819ms step_avg:57.37ms step:1340/1555 train_time:76909ms step_avg:57.39ms step:1341/1555 train_time:76993ms step_avg:57.41ms step:1342/1555 train_time:77082ms step_avg:57.44ms step:1343/1555 train_time:77166ms step_avg:57.46ms step:1344/1555 train_time:77257ms step_avg:57.48ms step:1345/1555 train_time:77341ms step_avg:57.50ms step:1346/1555 train_time:77432ms step_avg:57.53ms step:1347/1555 train_time:77516ms step_avg:57.55ms step:1348/1555 train_time:77605ms step_avg:57.57ms step:1349/1555 train_time:77690ms step_avg:57.59ms step:1350/1555 train_time:77780ms step_avg:57.61ms step:1351/1555 train_time:77863ms step_avg:57.63ms step:1352/1555 train_time:77954ms step_avg:57.66ms step:1353/1555 train_time:78037ms step_avg:57.68ms step:1354/1555 train_time:78127ms step_avg:57.70ms step:1355/1555 train_time:78211ms step_avg:57.72ms step:1356/1555 train_time:78300ms step_avg:57.74ms step:1357/1555 train_time:78384ms step_avg:57.76ms step:1358/1555 train_time:78476ms step_avg:57.79ms step:1359/1555 train_time:78559ms step_avg:57.81ms step:1360/1555 train_time:78650ms step_avg:57.83ms step:1361/1555 train_time:78734ms step_avg:57.85ms step:1362/1555 train_time:78824ms step_avg:57.87ms step:1363/1555 train_time:78909ms step_avg:57.89ms step:1364/1555 train_time:78998ms step_avg:57.92ms step:1365/1555 train_time:79082ms step_avg:57.94ms step:1366/1555 train_time:79173ms step_avg:57.96ms step:1367/1555 train_time:79256ms step_avg:57.98ms step:1368/1555 train_time:79346ms step_avg:58.00ms step:1369/1555 train_time:79431ms step_avg:58.02ms step:1370/1555 train_time:79521ms step_avg:58.04ms step:1371/1555 train_time:79605ms step_avg:58.06ms step:1372/1555 train_time:79696ms step_avg:58.09ms step:1373/1555 train_time:79780ms step_avg:58.11ms step:1374/1555 train_time:79871ms step_avg:58.13ms step:1375/1555 train_time:79954ms step_avg:58.15ms step:1376/1555 train_time:80043ms step_avg:58.17ms step:1377/1555 train_time:80129ms step_avg:58.19ms step:1378/1555 train_time:80219ms step_avg:58.21ms step:1379/1555 train_time:80302ms step_avg:58.23ms step:1380/1555 train_time:80393ms step_avg:58.26ms step:1381/1555 train_time:80476ms step_avg:58.27ms step:1382/1555 train_time:80565ms step_avg:58.30ms step:1383/1555 train_time:80652ms step_avg:58.32ms step:1384/1555 train_time:80741ms step_avg:58.34ms step:1385/1555 train_time:80825ms step_avg:58.36ms step:1386/1555 train_time:80916ms step_avg:58.38ms step:1387/1555 train_time:80999ms step_avg:58.40ms step:1388/1555 train_time:81089ms step_avg:58.42ms step:1389/1555 train_time:81174ms step_avg:58.44ms step:1390/1555 train_time:81263ms step_avg:58.46ms step:1391/1555 train_time:81347ms step_avg:58.48ms step:1392/1555 train_time:81438ms step_avg:58.50ms step:1393/1555 train_time:81522ms step_avg:58.52ms step:1394/1555 train_time:81614ms step_avg:58.55ms step:1395/1555 train_time:81697ms step_avg:58.56ms step:1396/1555 train_time:81788ms step_avg:58.59ms step:1397/1555 train_time:81872ms step_avg:58.61ms step:1398/1555 train_time:81961ms step_avg:58.63ms step:1399/1555 train_time:82045ms step_avg:58.65ms step:1400/1555 train_time:82137ms step_avg:58.67ms step:1401/1555 train_time:82220ms step_avg:58.69ms step:1402/1555 train_time:82311ms step_avg:58.71ms step:1403/1555 train_time:82394ms step_avg:58.73ms step:1404/1555 train_time:82483ms step_avg:58.75ms step:1405/1555 train_time:82569ms step_avg:58.77ms step:1406/1555 train_time:82659ms step_avg:58.79ms step:1407/1555 train_time:82742ms step_avg:58.81ms step:1408/1555 train_time:82833ms step_avg:58.83ms step:1409/1555 train_time:82916ms step_avg:58.85ms step:1410/1555 train_time:83005ms step_avg:58.87ms step:1411/1555 train_time:83089ms step_avg:58.89ms step:1412/1555 train_time:83179ms step_avg:58.91ms step:1413/1555 train_time:83263ms step_avg:58.93ms step:1414/1555 train_time:83353ms step_avg:58.95ms step:1415/1555 train_time:83437ms step_avg:58.97ms step:1416/1555 train_time:83527ms step_avg:58.99ms step:1417/1555 train_time:83611ms step_avg:59.01ms step:1418/1555 train_time:83700ms step_avg:59.03ms step:1419/1555 train_time:83784ms step_avg:59.04ms step:1420/1555 train_time:83875ms step_avg:59.07ms step:1421/1555 train_time:83959ms step_avg:59.08ms step:1422/1555 train_time:84049ms step_avg:59.11ms step:1423/1555 train_time:84134ms step_avg:59.12ms step:1424/1555 train_time:84224ms step_avg:59.15ms step:1425/1555 train_time:84309ms step_avg:59.16ms step:1426/1555 train_time:84399ms step_avg:59.19ms step:1427/1555 train_time:84484ms step_avg:59.20ms step:1428/1555 train_time:84575ms step_avg:59.23ms step:1429/1555 train_time:84659ms step_avg:59.24ms step:1430/1555 train_time:84748ms step_avg:59.26ms step:1431/1555 train_time:84833ms step_avg:59.28ms step:1432/1555 train_time:84922ms step_avg:59.30ms step:1433/1555 train_time:85007ms step_avg:59.32ms step:1434/1555 train_time:85097ms step_avg:59.34ms step:1435/1555 train_time:85182ms step_avg:59.36ms step:1436/1555 train_time:85273ms step_avg:59.38ms step:1437/1555 train_time:85356ms step_avg:59.40ms step:1438/1555 train_time:85446ms step_avg:59.42ms step:1439/1555 train_time:85531ms step_avg:59.44ms step:1440/1555 train_time:85621ms step_avg:59.46ms step:1441/1555 train_time:85706ms step_avg:59.48ms step:1442/1555 train_time:85795ms step_avg:59.50ms step:1443/1555 train_time:85880ms step_avg:59.52ms step:1444/1555 train_time:85970ms step_avg:59.54ms step:1445/1555 train_time:86054ms step_avg:59.55ms step:1446/1555 train_time:86144ms step_avg:59.57ms step:1447/1555 train_time:86230ms step_avg:59.59ms step:1448/1555 train_time:86319ms step_avg:59.61ms step:1449/1555 train_time:86403ms step_avg:59.63ms step:1450/1555 train_time:86493ms step_avg:59.65ms step:1451/1555 train_time:86576ms step_avg:59.67ms step:1452/1555 train_time:86667ms step_avg:59.69ms step:1453/1555 train_time:86751ms step_avg:59.70ms step:1454/1555 train_time:86841ms step_avg:59.73ms step:1455/1555 train_time:86925ms step_avg:59.74ms step:1456/1555 train_time:87017ms step_avg:59.76ms step:1457/1555 train_time:87100ms step_avg:59.78ms step:1458/1555 train_time:87190ms step_avg:59.80ms step:1459/1555 train_time:87274ms step_avg:59.82ms step:1460/1555 train_time:87364ms step_avg:59.84ms step:1461/1555 train_time:87450ms step_avg:59.86ms step:1462/1555 train_time:87539ms step_avg:59.88ms step:1463/1555 train_time:87624ms step_avg:59.89ms step:1464/1555 train_time:87716ms step_avg:59.92ms step:1465/1555 train_time:87798ms step_avg:59.93ms step:1466/1555 train_time:87889ms step_avg:59.95ms step:1467/1555 train_time:87973ms step_avg:59.97ms step:1468/1555 train_time:88062ms step_avg:59.99ms step:1469/1555 train_time:88149ms step_avg:60.01ms step:1470/1555 train_time:88239ms step_avg:60.03ms step:1471/1555 train_time:88322ms step_avg:60.04ms step:1472/1555 train_time:88414ms step_avg:60.06ms step:1473/1555 train_time:88497ms step_avg:60.08ms step:1474/1555 train_time:88587ms step_avg:60.10ms step:1475/1555 train_time:88670ms step_avg:60.12ms step:1476/1555 train_time:88760ms step_avg:60.14ms step:1477/1555 train_time:88846ms step_avg:60.15ms step:1478/1555 train_time:88936ms step_avg:60.17ms step:1479/1555 train_time:89019ms step_avg:60.19ms step:1480/1555 train_time:89111ms step_avg:60.21ms step:1481/1555 train_time:89194ms step_avg:60.23ms step:1482/1555 train_time:89283ms step_avg:60.25ms step:1483/1555 train_time:89369ms step_avg:60.26ms step:1484/1555 train_time:89459ms step_avg:60.28ms step:1485/1555 train_time:89543ms step_avg:60.30ms step:1486/1555 train_time:89633ms step_avg:60.32ms step:1487/1555 train_time:89716ms step_avg:60.33ms step:1488/1555 train_time:89807ms step_avg:60.35ms step:1489/1555 train_time:89890ms step_avg:60.37ms step:1490/1555 train_time:89980ms step_avg:60.39ms step:1491/1555 train_time:90064ms step_avg:60.41ms step:1492/1555 train_time:90155ms step_avg:60.43ms step:1493/1555 train_time:90239ms step_avg:60.44ms step:1494/1555 train_time:90329ms step_avg:60.46ms step:1495/1555 train_time:90413ms step_avg:60.48ms step:1496/1555 train_time:90503ms step_avg:60.50ms step:1497/1555 train_time:90588ms step_avg:60.51ms step:1498/1555 train_time:90677ms step_avg:60.53ms step:1499/1555 train_time:90760ms step_avg:60.55ms step:1500/1555 train_time:90851ms step_avg:60.57ms step:1500/1555 val_loss:3.2927 train_time:90966ms step_avg:60.64ms step:1501/1555 train_time:90986ms step_avg:60.62ms step:1502/1555 train_time:91026ms step_avg:60.60ms step:1503/1555 train_time:91112ms step_avg:60.62ms step:1504/1555 train_time:91206ms step_avg:60.64ms step:1505/1555 train_time:91291ms step_avg:60.66ms step:1506/1555 train_time:91382ms step_avg:60.68ms step:1507/1555 train_time:91465ms step_avg:60.69ms step:1508/1555 train_time:91553ms step_avg:60.71ms step:1509/1555 train_time:91637ms step_avg:60.73ms step:1510/1555 train_time:91726ms step_avg:60.75ms step:1511/1555 train_time:91808ms step_avg:60.76ms step:1512/1555 train_time:91899ms step_avg:60.78ms step:1513/1555 train_time:91985ms step_avg:60.80ms step:1514/1555 train_time:92076ms step_avg:60.82ms step:1515/1555 train_time:92163ms step_avg:60.83ms step:1516/1555 train_time:92258ms step_avg:60.86ms step:1517/1555 train_time:92345ms step_avg:60.87ms step:1518/1555 train_time:92433ms step_avg:60.89ms step:1519/1555 train_time:92518ms step_avg:60.91ms step:1520/1555 train_time:92607ms step_avg:60.93ms step:1521/1555 train_time:92689ms step_avg:60.94ms step:1522/1555 train_time:92779ms step_avg:60.96ms step:1523/1555 train_time:92864ms step_avg:60.97ms step:1524/1555 train_time:92954ms step_avg:60.99ms step:1525/1555 train_time:93041ms step_avg:61.01ms step:1526/1555 train_time:93132ms step_avg:61.03ms step:1527/1555 train_time:93218ms step_avg:61.05ms step:1528/1555 train_time:93308ms step_avg:61.07ms step:1529/1555 train_time:93393ms step_avg:61.08ms step:1530/1555 train_time:93485ms step_avg:61.10ms step:1531/1555 train_time:93568ms step_avg:61.12ms step:1532/1555 train_time:93658ms step_avg:61.13ms step:1533/1555 train_time:93741ms step_avg:61.15ms step:1534/1555 train_time:93831ms step_avg:61.17ms step:1535/1555 train_time:93915ms step_avg:61.18ms step:1536/1555 train_time:94007ms step_avg:61.20ms step:1537/1555 train_time:94092ms step_avg:61.22ms step:1538/1555 train_time:94183ms step_avg:61.24ms step:1539/1555 train_time:94268ms step_avg:61.25ms step:1540/1555 train_time:94358ms step_avg:61.27ms step:1541/1555 train_time:94443ms step_avg:61.29ms step:1542/1555 train_time:94532ms step_avg:61.30ms step:1543/1555 train_time:94617ms step_avg:61.32ms step:1544/1555 train_time:94707ms step_avg:61.34ms step:1545/1555 train_time:94791ms step_avg:61.35ms step:1546/1555 train_time:94881ms step_avg:61.37ms step:1547/1555 train_time:94966ms step_avg:61.39ms step:1548/1555 train_time:95056ms step_avg:61.41ms step:1549/1555 train_time:95141ms step_avg:61.42ms step:1550/1555 train_time:95231ms step_avg:61.44ms step:1551/1555 train_time:95317ms step_avg:61.46ms step:1552/1555 train_time:95408ms step_avg:61.47ms step:1553/1555 train_time:95493ms step_avg:61.49ms step:1554/1555 train_time:95584ms step_avg:61.51ms step:1555/1555 train_time:95667ms step_avg:61.52ms step:1555/1555 val_loss:3.2765 train_time:95782ms step_avg:61.60ms peak memory allocated: 31630 MiB reserved: 46718 MiB