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 05:57:40 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 28C P0 113W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 30C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:71:00.0 Off | 0 | | N/A 31C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:79:00.0 Off | 0 | | N/A 30C P0 121W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:7F:00.0 Off | 0 | | N/A 28C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:87:00.0 Off | 0 | | N/A 31C P0 114W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:8D:00.0 Off | 0 | | N/A 30C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:95:00.0 Off | 0 | | N/A 30C P0 115W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 95 C /usr/local/bin/python 1510MiB | | 1 N/A N/A 96 C /usr/local/bin/python 1510MiB | | 2 N/A N/A 97 C /usr/local/bin/python 1510MiB | | 3 N/A N/A 98 C /usr/local/bin/python 1510MiB | | 4 N/A N/A 99 C /usr/local/bin/python 1510MiB | | 5 N/A N/A 100 C /usr/local/bin/python 1510MiB | | 6 N/A N/A 101 C /usr/local/bin/python 1510MiB | | 7 N/A N/A 102 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.8316 train_time:0ms step_avg:0.03ms step:1/1555 train_time:92ms step_avg:92.38ms step:2/1555 train_time:117ms step_avg:58.35ms step:3/1555 train_time:138ms step_avg:46.13ms step:4/1555 train_time:160ms step_avg:40.03ms step:5/1555 train_time:186ms step_avg:37.15ms step:6/1555 train_time:224ms step_avg:37.31ms step:7/1555 train_time:254ms step_avg:36.28ms step:8/1555 train_time:292ms step_avg:36.47ms step:9/1555 train_time:322ms step_avg:35.78ms step:10/1555 train_time:360ms step_avg:35.96ms step:11/1555 train_time:390ms step_avg:35.46ms step:12/1555 train_time:428ms step_avg:35.66ms step:13/1555 train_time:459ms step_avg:35.27ms step:14/1555 train_time:496ms step_avg:35.45ms step:15/1555 train_time:527ms step_avg:35.12ms step:16/1555 train_time:565ms step_avg:35.30ms step:17/1555 train_time:596ms step_avg:35.04ms step:18/1555 train_time:633ms step_avg:35.18ms step:19/1555 train_time:665ms step_avg:35.02ms step:20/1555 train_time:703ms step_avg:35.14ms step:21/1555 train_time:733ms step_avg:34.89ms step:22/1555 train_time:770ms step_avg:35.01ms step:23/1555 train_time:801ms step_avg:34.82ms step:24/1555 train_time:838ms step_avg:34.94ms step:25/1555 train_time:869ms step_avg:34.78ms step:26/1555 train_time:907ms step_avg:34.89ms step:27/1555 train_time:938ms step_avg:34.73ms step:28/1555 train_time:975ms step_avg:34.83ms step:29/1555 train_time:1006ms step_avg:34.69ms step:30/1555 train_time:1044ms step_avg:34.80ms step:31/1555 train_time:1075ms step_avg:34.68ms step:32/1555 train_time:1114ms step_avg:34.80ms step:33/1555 train_time:1145ms step_avg:34.69ms step:34/1555 train_time:1183ms step_avg:34.81ms step:35/1555 train_time:1214ms step_avg:34.68ms 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:1322ms step_avg:34.80ms step:39/1555 train_time:1352ms step_avg:34.68ms step:40/1555 train_time:1390ms step_avg:34.75ms step:41/1555 train_time:1421ms step_avg:34.65ms step:42/1555 train_time:1458ms step_avg:34.72ms step:43/1555 train_time:1489ms step_avg:34.63ms step:44/1555 train_time:1527ms step_avg:34.70ms step:45/1555 train_time:1557ms step_avg:34.61ms step:46/1555 train_time:1595ms step_avg:34.67ms step:47/1555 train_time:1625ms step_avg:34.58ms step:48/1555 train_time:1663ms step_avg:34.65ms step:49/1555 train_time:1694ms step_avg:34.57ms step:50/1555 train_time:1732ms step_avg:34.63ms step:51/1555 train_time:1762ms step_avg:34.55ms step:52/1555 train_time:1799ms step_avg:34.60ms step:53/1555 train_time:1830ms step_avg:34.53ms step:54/1555 train_time:1868ms step_avg:34.59ms step:55/1555 train_time:1899ms step_avg:34.52ms step:56/1555 train_time:1937ms step_avg:34.58ms step:57/1555 train_time:1967ms step_avg:34.51ms step:58/1555 train_time:2005ms step_avg:34.57ms step:59/1555 train_time:2035ms step_avg:34.50ms step:60/1555 train_time:2073ms step_avg:34.55ms step:61/1555 train_time:2104ms step_avg:34.49ms step:62/1555 train_time:2142ms step_avg:34.55ms step:63/1555 train_time:2173ms step_avg:34.49ms step:64/1555 train_time:2211ms step_avg:34.54ms step:65/1555 train_time:2242ms step_avg:34.49ms step:66/1555 train_time:2280ms step_avg:34.55ms step:67/1555 train_time:2311ms step_avg:34.49ms step:68/1555 train_time:2349ms step_avg:34.54ms step:69/1555 train_time:2380ms step_avg:34.49ms step:70/1555 train_time:2418ms step_avg:34.54ms step:71/1555 train_time:2449ms step_avg:34.49ms step:72/1555 train_time:2487ms step_avg:34.54ms step:73/1555 train_time:2518ms step_avg:34.49ms step:74/1555 train_time:2555ms step_avg:34.53ms step:75/1555 train_time:2586ms step_avg:34.48ms step:76/1555 train_time:2624ms step_avg:34.53ms step:77/1555 train_time:2654ms step_avg:34.47ms step:78/1555 train_time:2692ms step_avg:34.51ms step:79/1555 train_time:2723ms step_avg:34.47ms step:80/1555 train_time:2760ms step_avg:34.50ms step:81/1555 train_time:2791ms step_avg:34.46ms step:82/1555 train_time:2830ms step_avg:34.51ms step:83/1555 train_time:2860ms step_avg:34.45ms step:84/1555 train_time:2897ms step_avg:34.49ms step:85/1555 train_time:2928ms step_avg:34.45ms step:86/1555 train_time:2966ms step_avg:34.49ms step:87/1555 train_time:2997ms step_avg:34.45ms step:88/1555 train_time:3035ms step_avg:34.48ms step:89/1555 train_time:3065ms step_avg:34.44ms step:90/1555 train_time:3103ms step_avg:34.48ms step:91/1555 train_time:3134ms step_avg:34.44ms step:92/1555 train_time:3171ms step_avg:34.47ms step:93/1555 train_time:3202ms step_avg:34.43ms step:94/1555 train_time:3240ms step_avg:34.47ms step:95/1555 train_time:3271ms step_avg:34.43ms step:96/1555 train_time:3309ms step_avg:34.47ms step:97/1555 train_time:3340ms step_avg:34.43ms step:98/1555 train_time:3377ms step_avg:34.46ms step:99/1555 train_time:3408ms step_avg:34.43ms step:100/1555 train_time:3446ms step_avg:34.46ms step:101/1555 train_time:3476ms step_avg:34.42ms step:102/1555 train_time:3514ms step_avg:34.45ms step:103/1555 train_time:3545ms step_avg:34.42ms step:104/1555 train_time:3583ms step_avg:34.45ms step:105/1555 train_time:3613ms step_avg:34.41ms step:106/1555 train_time:3651ms step_avg:34.44ms step:107/1555 train_time:3682ms step_avg:34.41ms step:108/1555 train_time:3720ms step_avg:34.44ms step:109/1555 train_time:3750ms step_avg:34.40ms step:110/1555 train_time:3788ms step_avg:34.43ms step:111/1555 train_time:3819ms step_avg:34.40ms step:112/1555 train_time:3856ms step_avg:34.43ms step:113/1555 train_time:3887ms step_avg:34.40ms step:114/1555 train_time:3925ms step_avg:34.43ms step:115/1555 train_time:3956ms step_avg:34.40ms step:116/1555 train_time:3993ms step_avg:34.43ms step:117/1555 train_time:4024ms step_avg:34.39ms step:118/1555 train_time:4062ms step_avg:34.42ms step:119/1555 train_time:4092ms step_avg:34.39ms step:120/1555 train_time:4130ms step_avg:34.42ms step:121/1555 train_time:4160ms step_avg:34.38ms step:122/1555 train_time:4197ms step_avg:34.41ms step:123/1555 train_time:4228ms step_avg:34.38ms step:124/1555 train_time:4266ms step_avg:34.40ms step:125/1555 train_time:4297ms step_avg:34.37ms step:126/1555 train_time:4334ms step_avg:34.40ms step:127/1555 train_time:4365ms step_avg:34.37ms step:128/1555 train_time:4402ms step_avg:34.39ms step:129/1555 train_time:4433ms step_avg:34.37ms step:130/1555 train_time:4471ms step_avg:34.39ms step:131/1555 train_time:4502ms step_avg:34.37ms step:132/1555 train_time:4539ms step_avg:34.39ms step:133/1555 train_time:4570ms step_avg:34.36ms step:134/1555 train_time:4608ms step_avg:34.38ms step:135/1555 train_time:4639ms step_avg:34.36ms step:136/1555 train_time:4676ms step_avg:34.38ms step:137/1555 train_time:4707ms step_avg:34.36ms step:138/1555 train_time:4745ms step_avg:34.39ms step:139/1555 train_time:4776ms step_avg:34.36ms step:140/1555 train_time:4813ms step_avg:34.38ms step:141/1555 train_time:4844ms step_avg:34.36ms step:142/1555 train_time:4883ms step_avg:34.38ms step:143/1555 train_time:4914ms step_avg:34.36ms step:144/1555 train_time:4951ms step_avg:34.38ms step:145/1555 train_time:4982ms step_avg:34.36ms step:146/1555 train_time:5019ms step_avg:34.38ms step:147/1555 train_time:5050ms step_avg:34.35ms step:148/1555 train_time:5088ms step_avg:34.38ms step:149/1555 train_time:5118ms step_avg:34.35ms step:150/1555 train_time:5156ms step_avg:34.37ms step:151/1555 train_time:5186ms step_avg:34.35ms step:152/1555 train_time:5224ms step_avg:34.37ms step:153/1555 train_time:5255ms step_avg:34.34ms step:154/1555 train_time:5292ms step_avg:34.37ms step:155/1555 train_time:5323ms step_avg:34.34ms step:156/1555 train_time:5361ms step_avg:34.36ms step:157/1555 train_time:5392ms step_avg:34.34ms step:158/1555 train_time:5430ms step_avg:34.37ms step:159/1555 train_time:5460ms step_avg:34.34ms step:160/1555 train_time:5498ms step_avg:34.36ms step:161/1555 train_time:5529ms step_avg:34.34ms step:162/1555 train_time:5567ms step_avg:34.36ms step:163/1555 train_time:5597ms step_avg:34.34ms step:164/1555 train_time:5635ms step_avg:34.36ms step:165/1555 train_time:5666ms step_avg:34.34ms step:166/1555 train_time:5703ms step_avg:34.36ms step:167/1555 train_time:5734ms step_avg:34.34ms step:168/1555 train_time:5772ms step_avg:34.36ms step:169/1555 train_time:5803ms step_avg:34.33ms step:170/1555 train_time:5840ms step_avg:34.35ms step:171/1555 train_time:5871ms step_avg:34.33ms step:172/1555 train_time:5908ms step_avg:34.35ms step:173/1555 train_time:5939ms step_avg:34.33ms step:174/1555 train_time:5976ms step_avg:34.35ms step:175/1555 train_time:6008ms step_avg:34.33ms step:176/1555 train_time:6046ms step_avg:34.35ms step:177/1555 train_time:6076ms step_avg:34.33ms step:178/1555 train_time:6114ms step_avg:34.35ms step:179/1555 train_time:6145ms step_avg:34.33ms step:180/1555 train_time:6183ms step_avg:34.35ms step:181/1555 train_time:6213ms step_avg:34.33ms step:182/1555 train_time:6251ms step_avg:34.35ms step:183/1555 train_time:6281ms step_avg:34.32ms step:184/1555 train_time:6319ms step_avg:34.34ms step:185/1555 train_time:6349ms step_avg:34.32ms step:186/1555 train_time:6387ms step_avg:34.34ms step:187/1555 train_time:6418ms step_avg:34.32ms step:188/1555 train_time:6456ms step_avg:34.34ms step:189/1555 train_time:6487ms step_avg:34.32ms step:190/1555 train_time:6525ms step_avg:34.34ms step:191/1555 train_time:6555ms step_avg:34.32ms step:192/1555 train_time:6593ms step_avg:34.34ms step:193/1555 train_time:6624ms step_avg:34.32ms step:194/1555 train_time:6661ms step_avg:34.34ms step:195/1555 train_time:6692ms step_avg:34.32ms step:196/1555 train_time:6729ms step_avg:34.33ms step:197/1555 train_time:6760ms step_avg:34.31ms step:198/1555 train_time:6797ms step_avg:34.33ms step:199/1555 train_time:6828ms step_avg:34.31ms step:200/1555 train_time:6865ms step_avg:34.33ms step:201/1555 train_time:6896ms step_avg:34.31ms step:202/1555 train_time:6933ms step_avg:34.32ms step:203/1555 train_time:6964ms step_avg:34.30ms step:204/1555 train_time:7001ms step_avg:34.32ms step:205/1555 train_time:7032ms step_avg:34.30ms step:206/1555 train_time:7070ms step_avg:34.32ms step:207/1555 train_time:7101ms step_avg:34.30ms step:208/1555 train_time:7138ms step_avg:34.32ms step:209/1555 train_time:7169ms step_avg:34.30ms step:210/1555 train_time:7207ms step_avg:34.32ms step:211/1555 train_time:7238ms step_avg:34.30ms step:212/1555 train_time:7275ms step_avg:34.32ms step:213/1555 train_time:7306ms step_avg:34.30ms step:214/1555 train_time:7344ms step_avg:34.32ms step:215/1555 train_time:7375ms step_avg:34.30ms step:216/1555 train_time:7412ms step_avg:34.32ms step:217/1555 train_time:7443ms step_avg:34.30ms step:218/1555 train_time:7480ms step_avg:34.31ms step:219/1555 train_time:7512ms step_avg:34.30ms step:220/1555 train_time:7549ms step_avg:34.32ms step:221/1555 train_time:7580ms step_avg:34.30ms step:222/1555 train_time:7618ms step_avg:34.31ms step:223/1555 train_time:7648ms step_avg:34.30ms step:224/1555 train_time:7686ms step_avg:34.31ms step:225/1555 train_time:7716ms step_avg:34.30ms step:226/1555 train_time:7754ms step_avg:34.31ms step:227/1555 train_time:7785ms step_avg:34.30ms step:228/1555 train_time:7824ms step_avg:34.32ms step:229/1555 train_time:7853ms step_avg:34.29ms step:230/1555 train_time:7891ms step_avg:34.31ms step:231/1555 train_time:7921ms step_avg:34.29ms step:232/1555 train_time:7959ms step_avg:34.31ms step:233/1555 train_time:7990ms step_avg:34.29ms step:234/1555 train_time:8027ms step_avg:34.30ms step:235/1555 train_time:8057ms step_avg:34.29ms step:236/1555 train_time:8095ms step_avg:34.30ms step:237/1555 train_time:8126ms step_avg:34.29ms step:238/1555 train_time:8165ms step_avg:34.31ms step:239/1555 train_time:8195ms step_avg:34.29ms step:240/1555 train_time:8233ms step_avg:34.30ms step:241/1555 train_time:8263ms step_avg:34.29ms step:242/1555 train_time:8300ms step_avg:34.30ms step:243/1555 train_time:8331ms step_avg:34.28ms step:244/1555 train_time:8368ms step_avg:34.30ms step:245/1555 train_time:8399ms step_avg:34.28ms step:246/1555 train_time:8436ms step_avg:34.29ms step:247/1555 train_time:8467ms step_avg:34.28ms step:248/1555 train_time:8505ms step_avg:34.29ms step:249/1555 train_time:8535ms step_avg:34.28ms step:250/1555 train_time:8573ms step_avg:34.29ms step:250/1555 val_loss:4.5356 train_time:8622ms step_avg:34.49ms step:251/1555 train_time:8642ms step_avg:34.43ms step:252/1555 train_time:8664ms step_avg:34.38ms step:253/1555 train_time:8683ms step_avg:34.32ms step:254/1555 train_time:8711ms step_avg:34.30ms step:255/1555 train_time:8745ms step_avg:34.29ms step:256/1555 train_time:8785ms step_avg:34.32ms step:257/1555 train_time:8816ms step_avg:34.30ms step:258/1555 train_time:8854ms step_avg:34.32ms step:259/1555 train_time:8884ms step_avg:34.30ms step:260/1555 train_time:8922ms step_avg:34.32ms step:261/1555 train_time:8953ms step_avg:34.30ms step:262/1555 train_time:8991ms step_avg:34.32ms step:263/1555 train_time:9022ms step_avg:34.30ms step:264/1555 train_time:9059ms step_avg:34.32ms step:265/1555 train_time:9090ms step_avg:34.30ms step:266/1555 train_time:9127ms step_avg:34.31ms step:267/1555 train_time:9158ms step_avg:34.30ms step:268/1555 train_time:9195ms step_avg:34.31ms step:269/1555 train_time:9226ms step_avg:34.30ms step:270/1555 train_time:9263ms step_avg:34.31ms step:271/1555 train_time:9294ms step_avg:34.29ms step:272/1555 train_time:9331ms step_avg:34.30ms step:273/1555 train_time:9361ms step_avg:34.29ms step:274/1555 train_time:9399ms step_avg:34.30ms step:275/1555 train_time:9430ms step_avg:34.29ms step:276/1555 train_time:9467ms step_avg:34.30ms step:277/1555 train_time:9497ms step_avg:34.29ms step:278/1555 train_time:9535ms step_avg:34.30ms step:279/1555 train_time:9566ms step_avg:34.29ms step:280/1555 train_time:9603ms step_avg:34.30ms step:281/1555 train_time:9634ms step_avg:34.28ms step:282/1555 train_time:9672ms step_avg:34.30ms step:283/1555 train_time:9703ms step_avg:34.29ms step:284/1555 train_time:9740ms step_avg:34.30ms step:285/1555 train_time:9771ms step_avg:34.28ms step:286/1555 train_time:9809ms step_avg:34.30ms step:287/1555 train_time:9840ms step_avg:34.28ms step:288/1555 train_time:9878ms step_avg:34.30ms step:289/1555 train_time:9909ms step_avg:34.29ms step:290/1555 train_time:9946ms step_avg:34.30ms step:291/1555 train_time:9977ms step_avg:34.28ms step:292/1555 train_time:10015ms step_avg:34.30ms step:293/1555 train_time:10045ms step_avg:34.28ms step:294/1555 train_time:10083ms step_avg:34.30ms step:295/1555 train_time:10113ms step_avg:34.28ms step:296/1555 train_time:10151ms step_avg:34.29ms step:297/1555 train_time:10182ms step_avg:34.28ms step:298/1555 train_time:10219ms step_avg:34.29ms step:299/1555 train_time:10250ms step_avg:34.28ms step:300/1555 train_time:10287ms step_avg:34.29ms step:301/1555 train_time:10318ms step_avg:34.28ms step:302/1555 train_time:10356ms step_avg:34.29ms step:303/1555 train_time:10387ms step_avg:34.28ms step:304/1555 train_time:10424ms step_avg:34.29ms step:305/1555 train_time:10455ms step_avg:34.28ms step:306/1555 train_time:10493ms step_avg:34.29ms step:307/1555 train_time:10524ms step_avg:34.28ms step:308/1555 train_time:10561ms step_avg:34.29ms step:309/1555 train_time:10593ms step_avg:34.28ms step:310/1555 train_time:10630ms step_avg:34.29ms step:311/1555 train_time:10661ms step_avg:34.28ms step:312/1555 train_time:10699ms step_avg:34.29ms step:313/1555 train_time:10729ms step_avg:34.28ms step:314/1555 train_time:10767ms step_avg:34.29ms step:315/1555 train_time:10798ms step_avg:34.28ms step:316/1555 train_time:10836ms step_avg:34.29ms step:317/1555 train_time:10867ms step_avg:34.28ms step:318/1555 train_time:10905ms step_avg:34.29ms step:319/1555 train_time:10936ms step_avg:34.28ms step:320/1555 train_time:10974ms step_avg:34.29ms step:321/1555 train_time:11005ms step_avg:34.28ms step:322/1555 train_time:11042ms step_avg:34.29ms step:323/1555 train_time:11073ms step_avg:34.28ms step:324/1555 train_time:11111ms step_avg:34.29ms step:325/1555 train_time:11142ms step_avg:34.28ms step:326/1555 train_time:11181ms step_avg:34.30ms step:327/1555 train_time:11211ms step_avg:34.29ms step:328/1555 train_time:11249ms step_avg:34.30ms step:329/1555 train_time:11280ms step_avg:34.29ms step:330/1555 train_time:11318ms step_avg:34.30ms step:331/1555 train_time:11349ms step_avg:34.29ms step:332/1555 train_time:11386ms step_avg:34.30ms step:333/1555 train_time:11417ms step_avg:34.28ms step:334/1555 train_time:11455ms step_avg:34.30ms step:335/1555 train_time:11485ms step_avg:34.28ms step:336/1555 train_time:11522ms step_avg:34.29ms step:337/1555 train_time:11553ms step_avg:34.28ms step:338/1555 train_time:11591ms step_avg:34.29ms step:339/1555 train_time:11621ms step_avg:34.28ms step:340/1555 train_time:11659ms step_avg:34.29ms step:341/1555 train_time:11690ms step_avg:34.28ms step:342/1555 train_time:11727ms step_avg:34.29ms step:343/1555 train_time:11758ms step_avg:34.28ms step:344/1555 train_time:11795ms step_avg:34.29ms step:345/1555 train_time:11826ms step_avg:34.28ms step:346/1555 train_time:11864ms step_avg:34.29ms step:347/1555 train_time:11895ms step_avg:34.28ms step:348/1555 train_time:11933ms step_avg:34.29ms step:349/1555 train_time:11964ms step_avg:34.28ms step:350/1555 train_time:12001ms step_avg:34.29ms step:351/1555 train_time:12032ms step_avg:34.28ms step:352/1555 train_time:12070ms step_avg:34.29ms step:353/1555 train_time:12100ms step_avg:34.28ms step:354/1555 train_time:12138ms step_avg:34.29ms step:355/1555 train_time:12169ms step_avg:34.28ms step:356/1555 train_time:12206ms step_avg:34.29ms step:357/1555 train_time:12237ms step_avg:34.28ms step:358/1555 train_time:12274ms step_avg:34.29ms step:359/1555 train_time:12305ms step_avg:34.28ms step:360/1555 train_time:12343ms step_avg:34.29ms step:361/1555 train_time:12373ms step_avg:34.27ms step:362/1555 train_time:12411ms step_avg:34.28ms step:363/1555 train_time:12442ms step_avg:34.28ms step:364/1555 train_time:12479ms step_avg:34.28ms step:365/1555 train_time:12510ms step_avg:34.27ms step:366/1555 train_time:12547ms step_avg:34.28ms step:367/1555 train_time:12578ms step_avg:34.27ms step:368/1555 train_time:12616ms step_avg:34.28ms step:369/1555 train_time:12646ms step_avg:34.27ms step:370/1555 train_time:12684ms step_avg:34.28ms step:371/1555 train_time:12714ms step_avg:34.27ms step:372/1555 train_time:12752ms step_avg:34.28ms step:373/1555 train_time:12783ms step_avg:34.27ms step:374/1555 train_time:12821ms step_avg:34.28ms step:375/1555 train_time:12851ms step_avg:34.27ms step:376/1555 train_time:12889ms step_avg:34.28ms step:377/1555 train_time:12920ms step_avg:34.27ms step:378/1555 train_time:12958ms step_avg:34.28ms step:379/1555 train_time:12989ms step_avg:34.27ms step:380/1555 train_time:13026ms step_avg:34.28ms step:381/1555 train_time:13057ms step_avg:34.27ms step:382/1555 train_time:13096ms step_avg:34.28ms step:383/1555 train_time:13126ms step_avg:34.27ms step:384/1555 train_time:13164ms step_avg:34.28ms step:385/1555 train_time:13194ms step_avg:34.27ms step:386/1555 train_time:13232ms step_avg:34.28ms step:387/1555 train_time:13263ms step_avg:34.27ms step:388/1555 train_time:13300ms step_avg:34.28ms step:389/1555 train_time:13332ms step_avg:34.27ms step:390/1555 train_time:13369ms step_avg:34.28ms step:391/1555 train_time:13400ms step_avg:34.27ms step:392/1555 train_time:13438ms step_avg:34.28ms step:393/1555 train_time:13469ms step_avg:34.27ms step:394/1555 train_time:13506ms step_avg:34.28ms step:395/1555 train_time:13537ms step_avg:34.27ms step:396/1555 train_time:13575ms step_avg:34.28ms step:397/1555 train_time:13605ms step_avg:34.27ms step:398/1555 train_time:13642ms step_avg:34.28ms step:399/1555 train_time:13674ms step_avg:34.27ms step:400/1555 train_time:13711ms step_avg:34.28ms step:401/1555 train_time:13742ms step_avg:34.27ms step:402/1555 train_time:13779ms step_avg:34.28ms step:403/1555 train_time:13810ms step_avg:34.27ms step:404/1555 train_time:13848ms step_avg:34.28ms step:405/1555 train_time:13879ms step_avg:34.27ms step:406/1555 train_time:13917ms step_avg:34.28ms step:407/1555 train_time:13948ms step_avg:34.27ms step:408/1555 train_time:13985ms step_avg:34.28ms step:409/1555 train_time:14016ms step_avg:34.27ms step:410/1555 train_time:14054ms step_avg:34.28ms step:411/1555 train_time:14084ms step_avg:34.27ms step:412/1555 train_time:14122ms step_avg:34.28ms step:413/1555 train_time:14153ms step_avg:34.27ms step:414/1555 train_time:14190ms step_avg:34.28ms step:415/1555 train_time:14220ms step_avg:34.27ms step:416/1555 train_time:14258ms step_avg:34.27ms step:417/1555 train_time:14288ms step_avg:34.26ms step:418/1555 train_time:14326ms step_avg:34.27ms step:419/1555 train_time:14357ms step_avg:34.26ms step:420/1555 train_time:14395ms step_avg:34.27ms step:421/1555 train_time:14425ms step_avg:34.26ms step:422/1555 train_time:14463ms step_avg:34.27ms step:423/1555 train_time:14494ms step_avg:34.26ms step:424/1555 train_time:14531ms step_avg:34.27ms step:425/1555 train_time:14562ms step_avg:34.26ms step:426/1555 train_time:14600ms step_avg:34.27ms step:427/1555 train_time:14631ms step_avg:34.26ms step:428/1555 train_time:14668ms step_avg:34.27ms step:429/1555 train_time:14699ms step_avg:34.26ms step:430/1555 train_time:14736ms step_avg:34.27ms step:431/1555 train_time:14767ms step_avg:34.26ms step:432/1555 train_time:14805ms step_avg:34.27ms step:433/1555 train_time:14835ms step_avg:34.26ms step:434/1555 train_time:14873ms step_avg:34.27ms step:435/1555 train_time:14904ms step_avg:34.26ms step:436/1555 train_time:14942ms step_avg:34.27ms step:437/1555 train_time:14973ms step_avg:34.26ms step:438/1555 train_time:15011ms step_avg:34.27ms step:439/1555 train_time:15042ms step_avg:34.26ms step:440/1555 train_time:15079ms step_avg:34.27ms step:441/1555 train_time:15110ms step_avg:34.26ms step:442/1555 train_time:15147ms step_avg:34.27ms step:443/1555 train_time:15177ms step_avg:34.26ms step:444/1555 train_time:15215ms step_avg:34.27ms step:445/1555 train_time:15246ms step_avg:34.26ms step:446/1555 train_time:15283ms step_avg:34.27ms step:447/1555 train_time:15314ms step_avg:34.26ms step:448/1555 train_time:15353ms step_avg:34.27ms step:449/1555 train_time:15383ms step_avg:34.26ms step:450/1555 train_time:15420ms step_avg:34.27ms step:451/1555 train_time:15451ms step_avg:34.26ms step:452/1555 train_time:15490ms step_avg:34.27ms step:453/1555 train_time:15520ms step_avg:34.26ms step:454/1555 train_time:15558ms step_avg:34.27ms step:455/1555 train_time:15589ms step_avg:34.26ms step:456/1555 train_time:15626ms step_avg:34.27ms step:457/1555 train_time:15657ms step_avg:34.26ms step:458/1555 train_time:15695ms step_avg:34.27ms step:459/1555 train_time:15725ms step_avg:34.26ms step:460/1555 train_time:15763ms step_avg:34.27ms step:461/1555 train_time:15794ms step_avg:34.26ms step:462/1555 train_time:15832ms step_avg:34.27ms step:463/1555 train_time:15862ms step_avg:34.26ms step:464/1555 train_time:15901ms step_avg:34.27ms step:465/1555 train_time:15931ms step_avg:34.26ms step:466/1555 train_time:15969ms step_avg:34.27ms step:467/1555 train_time:16000ms step_avg:34.26ms step:468/1555 train_time:16038ms step_avg:34.27ms step:469/1555 train_time:16068ms step_avg:34.26ms step:470/1555 train_time:16106ms step_avg:34.27ms step:471/1555 train_time:16137ms step_avg:34.26ms step:472/1555 train_time:16175ms step_avg:34.27ms step:473/1555 train_time:16205ms step_avg:34.26ms step:474/1555 train_time:16243ms step_avg:34.27ms step:475/1555 train_time:16274ms step_avg:34.26ms step:476/1555 train_time:16312ms step_avg:34.27ms step:477/1555 train_time:16343ms step_avg:34.26ms step:478/1555 train_time:16381ms step_avg:34.27ms step:479/1555 train_time:16411ms step_avg:34.26ms step:480/1555 train_time:16449ms step_avg:34.27ms step:481/1555 train_time:16480ms step_avg:34.26ms step:482/1555 train_time:16517ms step_avg:34.27ms step:483/1555 train_time:16548ms step_avg:34.26ms step:484/1555 train_time:16585ms step_avg:34.27ms step:485/1555 train_time:16615ms step_avg:34.26ms step:486/1555 train_time:16653ms step_avg:34.27ms step:487/1555 train_time:16684ms step_avg:34.26ms step:488/1555 train_time:16721ms step_avg:34.26ms step:489/1555 train_time:16752ms step_avg:34.26ms step:490/1555 train_time:16790ms step_avg:34.27ms step:491/1555 train_time:16821ms step_avg:34.26ms step:492/1555 train_time:16858ms step_avg:34.27ms step:493/1555 train_time:16889ms step_avg:34.26ms step:494/1555 train_time:16927ms step_avg:34.26ms step:495/1555 train_time:16958ms step_avg:34.26ms step:496/1555 train_time:16995ms step_avg:34.26ms step:497/1555 train_time:17026ms step_avg:34.26ms step:498/1555 train_time:17063ms step_avg:34.26ms step:499/1555 train_time:17094ms step_avg:34.26ms step:500/1555 train_time:17132ms step_avg:34.26ms step:500/1555 val_loss:4.2220 train_time:17181ms step_avg:34.36ms step:501/1555 train_time:17204ms step_avg:34.34ms step:502/1555 train_time:17226ms step_avg:34.31ms step:503/1555 train_time:17245ms step_avg:34.28ms step:504/1555 train_time:17272ms step_avg:34.27ms step:505/1555 train_time:17305ms step_avg:34.27ms step:506/1555 train_time:17351ms step_avg:34.29ms step:507/1555 train_time:17401ms step_avg:34.32ms step:508/1555 train_time:17466ms step_avg:34.38ms step:509/1555 train_time:17523ms step_avg:34.43ms step:510/1555 train_time:17587ms step_avg:34.48ms step:511/1555 train_time:17644ms step_avg:34.53ms step:512/1555 train_time:17707ms step_avg:34.58ms step:513/1555 train_time:17763ms step_avg:34.63ms step:514/1555 train_time:17827ms step_avg:34.68ms step:515/1555 train_time:17883ms step_avg:34.72ms step:516/1555 train_time:17946ms step_avg:34.78ms step:517/1555 train_time:18003ms step_avg:34.82ms step:518/1555 train_time:18068ms step_avg:34.88ms step:519/1555 train_time:18127ms step_avg:34.93ms step:520/1555 train_time:18192ms step_avg:34.98ms step:521/1555 train_time:18252ms step_avg:35.03ms step:522/1555 train_time:18317ms step_avg:35.09ms step:523/1555 train_time:18375ms step_avg:35.13ms step:524/1555 train_time:18440ms step_avg:35.19ms step:525/1555 train_time:18497ms step_avg:35.23ms step:526/1555 train_time:18562ms step_avg:35.29ms step:527/1555 train_time:18620ms step_avg:35.33ms step:528/1555 train_time:18684ms step_avg:35.39ms step:529/1555 train_time:18742ms step_avg:35.43ms step:530/1555 train_time:18805ms step_avg:35.48ms step:531/1555 train_time:18863ms step_avg:35.52ms step:532/1555 train_time:18927ms step_avg:35.58ms step:533/1555 train_time:18984ms step_avg:35.62ms step:534/1555 train_time:19048ms step_avg:35.67ms step:535/1555 train_time:19106ms step_avg:35.71ms step:536/1555 train_time:19172ms step_avg:35.77ms step:537/1555 train_time:19230ms step_avg:35.81ms step:538/1555 train_time:19294ms step_avg:35.86ms step:539/1555 train_time:19351ms step_avg:35.90ms step:540/1555 train_time:19415ms step_avg:35.95ms step:541/1555 train_time:19474ms step_avg:36.00ms step:542/1555 train_time:19539ms step_avg:36.05ms step:543/1555 train_time:19596ms step_avg:36.09ms step:544/1555 train_time:19661ms step_avg:36.14ms step:545/1555 train_time:19720ms step_avg:36.18ms step:546/1555 train_time:19783ms step_avg:36.23ms step:547/1555 train_time:19840ms step_avg:36.27ms step:548/1555 train_time:19904ms step_avg:36.32ms step:549/1555 train_time:19961ms step_avg:36.36ms step:550/1555 train_time:20026ms step_avg:36.41ms step:551/1555 train_time:20084ms step_avg:36.45ms step:552/1555 train_time:20148ms step_avg:36.50ms step:553/1555 train_time:20207ms step_avg:36.54ms step:554/1555 train_time:20271ms step_avg:36.59ms step:555/1555 train_time:20329ms step_avg:36.63ms step:556/1555 train_time:20393ms step_avg:36.68ms step:557/1555 train_time:20450ms step_avg:36.71ms step:558/1555 train_time:20514ms step_avg:36.76ms step:559/1555 train_time:20571ms step_avg:36.80ms step:560/1555 train_time:20636ms step_avg:36.85ms step:561/1555 train_time:20694ms step_avg:36.89ms step:562/1555 train_time:20759ms step_avg:36.94ms step:563/1555 train_time:20817ms step_avg:36.98ms step:564/1555 train_time:20881ms step_avg:37.02ms step:565/1555 train_time:20939ms step_avg:37.06ms step:566/1555 train_time:21004ms step_avg:37.11ms step:567/1555 train_time:21061ms step_avg:37.15ms step:568/1555 train_time:21126ms step_avg:37.19ms step:569/1555 train_time:21185ms step_avg:37.23ms step:570/1555 train_time:21248ms step_avg:37.28ms step:571/1555 train_time:21306ms step_avg:37.31ms step:572/1555 train_time:21371ms step_avg:37.36ms step:573/1555 train_time:21429ms step_avg:37.40ms step:574/1555 train_time:21493ms step_avg:37.44ms step:575/1555 train_time:21550ms step_avg:37.48ms step:576/1555 train_time:21614ms step_avg:37.52ms step:577/1555 train_time:21672ms step_avg:37.56ms step:578/1555 train_time:21737ms step_avg:37.61ms step:579/1555 train_time:21794ms step_avg:37.64ms step:580/1555 train_time:21859ms step_avg:37.69ms step:581/1555 train_time:21918ms step_avg:37.72ms step:582/1555 train_time:21982ms step_avg:37.77ms step:583/1555 train_time:22039ms step_avg:37.80ms step:584/1555 train_time:22104ms step_avg:37.85ms step:585/1555 train_time:22162ms step_avg:37.88ms step:586/1555 train_time:22228ms step_avg:37.93ms step:587/1555 train_time:22285ms step_avg:37.96ms step:588/1555 train_time:22349ms step_avg:38.01ms step:589/1555 train_time:22408ms step_avg:38.04ms step:590/1555 train_time:22472ms step_avg:38.09ms step:591/1555 train_time:22530ms step_avg:38.12ms step:592/1555 train_time:22593ms step_avg:38.16ms step:593/1555 train_time:22650ms step_avg:38.20ms step:594/1555 train_time:22714ms step_avg:38.24ms step:595/1555 train_time:22772ms step_avg:38.27ms step:596/1555 train_time:22837ms step_avg:38.32ms step:597/1555 train_time:22894ms step_avg:38.35ms step:598/1555 train_time:22960ms step_avg:38.39ms step:599/1555 train_time:23018ms step_avg:38.43ms step:600/1555 train_time:23083ms step_avg:38.47ms step:601/1555 train_time:23141ms step_avg:38.50ms step:602/1555 train_time:23205ms step_avg:38.55ms step:603/1555 train_time:23263ms step_avg:38.58ms step:604/1555 train_time:23327ms step_avg:38.62ms step:605/1555 train_time:23385ms step_avg:38.65ms step:606/1555 train_time:23449ms step_avg:38.69ms step:607/1555 train_time:23506ms step_avg:38.73ms step:608/1555 train_time:23570ms step_avg:38.77ms step:609/1555 train_time:23629ms step_avg:38.80ms step:610/1555 train_time:23692ms step_avg:38.84ms step:611/1555 train_time:23750ms step_avg:38.87ms step:612/1555 train_time:23815ms step_avg:38.91ms step:613/1555 train_time:23873ms step_avg:38.94ms step:614/1555 train_time:23938ms step_avg:38.99ms step:615/1555 train_time:23996ms step_avg:39.02ms step:616/1555 train_time:24061ms step_avg:39.06ms step:617/1555 train_time:24119ms step_avg:39.09ms step:618/1555 train_time:24184ms step_avg:39.13ms step:619/1555 train_time:24241ms step_avg:39.16ms step:620/1555 train_time:24306ms step_avg:39.20ms step:621/1555 train_time:24364ms step_avg:39.23ms step:622/1555 train_time:24428ms step_avg:39.27ms step:623/1555 train_time:24486ms step_avg:39.30ms step:624/1555 train_time:24550ms step_avg:39.34ms step:625/1555 train_time:24607ms step_avg:39.37ms step:626/1555 train_time:24670ms step_avg:39.41ms step:627/1555 train_time:24728ms step_avg:39.44ms step:628/1555 train_time:24792ms step_avg:39.48ms step:629/1555 train_time:24850ms step_avg:39.51ms step:630/1555 train_time:24915ms step_avg:39.55ms step:631/1555 train_time:24973ms step_avg:39.58ms step:632/1555 train_time:25038ms step_avg:39.62ms step:633/1555 train_time:25096ms step_avg:39.65ms step:634/1555 train_time:25160ms step_avg:39.69ms step:635/1555 train_time:25220ms step_avg:39.72ms step:636/1555 train_time:25284ms step_avg:39.76ms step:637/1555 train_time:25342ms step_avg:39.78ms step:638/1555 train_time:25406ms step_avg:39.82ms step:639/1555 train_time:25464ms step_avg:39.85ms step:640/1555 train_time:25528ms step_avg:39.89ms step:641/1555 train_time:25585ms step_avg:39.91ms step:642/1555 train_time:25649ms step_avg:39.95ms step:643/1555 train_time:25707ms step_avg:39.98ms step:644/1555 train_time:25771ms step_avg:40.02ms step:645/1555 train_time:25830ms step_avg:40.05ms step:646/1555 train_time:25893ms step_avg:40.08ms step:647/1555 train_time:25950ms step_avg:40.11ms step:648/1555 train_time:26014ms step_avg:40.15ms step:649/1555 train_time:26071ms step_avg:40.17ms step:650/1555 train_time:26136ms step_avg:40.21ms step:651/1555 train_time:26194ms step_avg:40.24ms step:652/1555 train_time:26260ms step_avg:40.28ms step:653/1555 train_time:26318ms step_avg:40.30ms step:654/1555 train_time:26383ms step_avg:40.34ms step:655/1555 train_time:26441ms step_avg:40.37ms step:656/1555 train_time:26505ms step_avg:40.40ms step:657/1555 train_time:26563ms step_avg:40.43ms step:658/1555 train_time:26627ms step_avg:40.47ms step:659/1555 train_time:26684ms step_avg:40.49ms step:660/1555 train_time:26748ms step_avg:40.53ms step:661/1555 train_time:26806ms step_avg:40.55ms step:662/1555 train_time:26870ms step_avg:40.59ms step:663/1555 train_time:26928ms step_avg:40.62ms step:664/1555 train_time:26992ms step_avg:40.65ms step:665/1555 train_time:27049ms step_avg:40.68ms step:666/1555 train_time:27114ms step_avg:40.71ms step:667/1555 train_time:27172ms step_avg:40.74ms step:668/1555 train_time:27237ms step_avg:40.77ms step:669/1555 train_time:27294ms step_avg:40.80ms step:670/1555 train_time:27360ms step_avg:40.84ms step:671/1555 train_time:27418ms step_avg:40.86ms step:672/1555 train_time:27483ms step_avg:40.90ms step:673/1555 train_time:27540ms step_avg:40.92ms step:674/1555 train_time:27605ms step_avg:40.96ms step:675/1555 train_time:27663ms step_avg:40.98ms step:676/1555 train_time:27728ms step_avg:41.02ms step:677/1555 train_time:27785ms step_avg:41.04ms step:678/1555 train_time:27849ms step_avg:41.08ms step:679/1555 train_time:27906ms step_avg:41.10ms step:680/1555 train_time:27970ms step_avg:41.13ms step:681/1555 train_time:28029ms step_avg:41.16ms step:682/1555 train_time:28092ms step_avg:41.19ms step:683/1555 train_time:28150ms step_avg:41.21ms step:684/1555 train_time:28214ms step_avg:41.25ms step:685/1555 train_time:28272ms step_avg:41.27ms step:686/1555 train_time:28338ms step_avg:41.31ms step:687/1555 train_time:28395ms step_avg:41.33ms step:688/1555 train_time:28460ms step_avg:41.37ms step:689/1555 train_time:28518ms step_avg:41.39ms step:690/1555 train_time:28582ms step_avg:41.42ms step:691/1555 train_time:28640ms step_avg:41.45ms step:692/1555 train_time:28705ms step_avg:41.48ms step:693/1555 train_time:28763ms step_avg:41.51ms step:694/1555 train_time:28827ms step_avg:41.54ms step:695/1555 train_time:28885ms step_avg:41.56ms step:696/1555 train_time:28949ms step_avg:41.59ms step:697/1555 train_time:29007ms step_avg:41.62ms step:698/1555 train_time:29071ms step_avg:41.65ms step:699/1555 train_time:29129ms step_avg:41.67ms step:700/1555 train_time:29193ms step_avg:41.70ms step:701/1555 train_time:29250ms step_avg:41.73ms step:702/1555 train_time:29315ms step_avg:41.76ms step:703/1555 train_time:29373ms step_avg:41.78ms step:704/1555 train_time:29438ms step_avg:41.81ms step:705/1555 train_time:29496ms step_avg:41.84ms step:706/1555 train_time:29559ms step_avg:41.87ms step:707/1555 train_time:29618ms step_avg:41.89ms step:708/1555 train_time:29683ms step_avg:41.93ms step:709/1555 train_time:29741ms step_avg:41.95ms step:710/1555 train_time:29805ms step_avg:41.98ms step:711/1555 train_time:29863ms step_avg:42.00ms step:712/1555 train_time:29927ms step_avg:42.03ms step:713/1555 train_time:29985ms step_avg:42.05ms step:714/1555 train_time:30049ms step_avg:42.09ms step:715/1555 train_time:30107ms step_avg:42.11ms step:716/1555 train_time:30170ms step_avg:42.14ms step:717/1555 train_time:30229ms step_avg:42.16ms step:718/1555 train_time:30292ms step_avg:42.19ms step:719/1555 train_time:30350ms step_avg:42.21ms step:720/1555 train_time:30414ms step_avg:42.24ms step:721/1555 train_time:30472ms step_avg:42.26ms step:722/1555 train_time:30537ms step_avg:42.30ms step:723/1555 train_time:30595ms step_avg:42.32ms step:724/1555 train_time:30661ms step_avg:42.35ms step:725/1555 train_time:30720ms step_avg:42.37ms step:726/1555 train_time:30784ms step_avg:42.40ms step:727/1555 train_time:30841ms step_avg:42.42ms step:728/1555 train_time:30905ms step_avg:42.45ms step:729/1555 train_time:30962ms step_avg:42.47ms step:730/1555 train_time:31028ms step_avg:42.50ms step:731/1555 train_time:31085ms step_avg:42.52ms step:732/1555 train_time:31149ms step_avg:42.55ms step:733/1555 train_time:31207ms step_avg:42.57ms step:734/1555 train_time:31271ms step_avg:42.60ms step:735/1555 train_time:31329ms step_avg:42.62ms step:736/1555 train_time:31393ms step_avg:42.65ms step:737/1555 train_time:31451ms step_avg:42.67ms step:738/1555 train_time:31516ms step_avg:42.70ms step:739/1555 train_time:31573ms step_avg:42.72ms step:740/1555 train_time:31639ms step_avg:42.76ms step:741/1555 train_time:31697ms step_avg:42.78ms step:742/1555 train_time:31761ms step_avg:42.80ms step:743/1555 train_time:31819ms step_avg:42.83ms step:744/1555 train_time:31884ms step_avg:42.85ms step:745/1555 train_time:31942ms step_avg:42.87ms step:746/1555 train_time:32005ms step_avg:42.90ms step:747/1555 train_time:32063ms step_avg:42.92ms step:748/1555 train_time:32128ms step_avg:42.95ms step:749/1555 train_time:32186ms step_avg:42.97ms step:750/1555 train_time:32249ms step_avg:43.00ms step:750/1555 val_loss:3.8682 train_time:32331ms step_avg:43.11ms step:751/1555 train_time:32354ms step_avg:43.08ms step:752/1555 train_time:32376ms step_avg:43.05ms step:753/1555 train_time:32431ms step_avg:43.07ms step:754/1555 train_time:32502ms step_avg:43.11ms step:755/1555 train_time:32562ms step_avg:43.13ms step:756/1555 train_time:32626ms step_avg:43.16ms step:757/1555 train_time:32683ms step_avg:43.17ms step:758/1555 train_time:32746ms step_avg:43.20ms step:759/1555 train_time:32804ms step_avg:43.22ms step:760/1555 train_time:32867ms step_avg:43.25ms step:761/1555 train_time:32924ms step_avg:43.26ms step:762/1555 train_time:32988ms step_avg:43.29ms step:763/1555 train_time:33045ms step_avg:43.31ms step:764/1555 train_time:33109ms step_avg:43.34ms step:765/1555 train_time:33166ms step_avg:43.35ms step:766/1555 train_time:33230ms step_avg:43.38ms step:767/1555 train_time:33287ms step_avg:43.40ms step:768/1555 train_time:33352ms step_avg:43.43ms step:769/1555 train_time:33412ms step_avg:43.45ms step:770/1555 train_time:33480ms step_avg:43.48ms step:771/1555 train_time:33539ms step_avg:43.50ms step:772/1555 train_time:33603ms step_avg:43.53ms step:773/1555 train_time:33660ms step_avg:43.55ms step:774/1555 train_time:33724ms step_avg:43.57ms step:775/1555 train_time:33782ms step_avg:43.59ms step:776/1555 train_time:33845ms step_avg:43.61ms step:777/1555 train_time:33902ms step_avg:43.63ms step:778/1555 train_time:33965ms step_avg:43.66ms step:779/1555 train_time:34022ms step_avg:43.67ms step:780/1555 train_time:34086ms step_avg:43.70ms step:781/1555 train_time:34143ms step_avg:43.72ms step:782/1555 train_time:34207ms step_avg:43.74ms step:783/1555 train_time:34264ms step_avg:43.76ms step:784/1555 train_time:34330ms step_avg:43.79ms step:785/1555 train_time:34388ms step_avg:43.81ms step:786/1555 train_time:34454ms step_avg:43.84ms step:787/1555 train_time:34513ms step_avg:43.85ms step:788/1555 train_time:34578ms step_avg:43.88ms step:789/1555 train_time:34636ms step_avg:43.90ms step:790/1555 train_time:34700ms step_avg:43.92ms step:791/1555 train_time:34757ms step_avg:43.94ms step:792/1555 train_time:34821ms step_avg:43.97ms step:793/1555 train_time:34879ms step_avg:43.98ms step:794/1555 train_time:34943ms step_avg:44.01ms step:795/1555 train_time:35001ms step_avg:44.03ms step:796/1555 train_time:35064ms step_avg:44.05ms step:797/1555 train_time:35122ms step_avg:44.07ms step:798/1555 train_time:35186ms step_avg:44.09ms step:799/1555 train_time:35244ms step_avg:44.11ms step:800/1555 train_time:35308ms step_avg:44.13ms step:801/1555 train_time:35365ms step_avg:44.15ms step:802/1555 train_time:35431ms step_avg:44.18ms step:803/1555 train_time:35489ms step_avg:44.20ms step:804/1555 train_time:35555ms step_avg:44.22ms step:805/1555 train_time:35613ms step_avg:44.24ms step:806/1555 train_time:35677ms step_avg:44.26ms step:807/1555 train_time:35735ms step_avg:44.28ms step:808/1555 train_time:35799ms step_avg:44.31ms step:809/1555 train_time:35858ms step_avg:44.32ms step:810/1555 train_time:35921ms step_avg:44.35ms step:811/1555 train_time:35979ms step_avg:44.36ms step:812/1555 train_time:36043ms step_avg:44.39ms step:813/1555 train_time:36100ms step_avg:44.40ms step:814/1555 train_time:36164ms step_avg:44.43ms step:815/1555 train_time:36223ms step_avg:44.45ms step:816/1555 train_time:36286ms step_avg:44.47ms step:817/1555 train_time:36343ms step_avg:44.48ms step:818/1555 train_time:36407ms step_avg:44.51ms step:819/1555 train_time:36466ms step_avg:44.53ms step:820/1555 train_time:36531ms step_avg:44.55ms step:821/1555 train_time:36589ms step_avg:44.57ms step:822/1555 train_time:36655ms step_avg:44.59ms step:823/1555 train_time:36713ms step_avg:44.61ms step:824/1555 train_time:36778ms step_avg:44.63ms step:825/1555 train_time:36836ms step_avg:44.65ms step:826/1555 train_time:36900ms step_avg:44.67ms step:827/1555 train_time:36957ms step_avg:44.69ms step:828/1555 train_time:37021ms step_avg:44.71ms step:829/1555 train_time:37077ms step_avg:44.73ms step:830/1555 train_time:37142ms step_avg:44.75ms step:831/1555 train_time:37200ms step_avg:44.77ms step:832/1555 train_time:37263ms step_avg:44.79ms step:833/1555 train_time:37321ms step_avg:44.80ms step:834/1555 train_time:37386ms step_avg:44.83ms step:835/1555 train_time:37443ms step_avg:44.84ms step:836/1555 train_time:37508ms step_avg:44.87ms step:837/1555 train_time:37566ms step_avg:44.88ms step:838/1555 train_time:37632ms step_avg:44.91ms step:839/1555 train_time:37689ms step_avg:44.92ms step:840/1555 train_time:37755ms step_avg:44.95ms step:841/1555 train_time:37813ms step_avg:44.96ms step:842/1555 train_time:37878ms step_avg:44.99ms step:843/1555 train_time:37937ms step_avg:45.00ms step:844/1555 train_time:38000ms step_avg:45.02ms step:845/1555 train_time:38057ms step_avg:45.04ms step:846/1555 train_time:38122ms step_avg:45.06ms step:847/1555 train_time:38179ms step_avg:45.08ms step:848/1555 train_time:38242ms step_avg:45.10ms step:849/1555 train_time:38300ms step_avg:45.11ms step:850/1555 train_time:38364ms step_avg:45.13ms step:851/1555 train_time:38422ms step_avg:45.15ms step:852/1555 train_time:38487ms step_avg:45.17ms step:853/1555 train_time:38545ms step_avg:45.19ms step:854/1555 train_time:38609ms step_avg:45.21ms step:855/1555 train_time:38667ms step_avg:45.22ms step:856/1555 train_time:38731ms step_avg:45.25ms step:857/1555 train_time:38789ms step_avg:45.26ms step:858/1555 train_time:38855ms step_avg:45.29ms step:859/1555 train_time:38913ms step_avg:45.30ms step:860/1555 train_time:38977ms step_avg:45.32ms step:861/1555 train_time:39034ms step_avg:45.34ms step:862/1555 train_time:39099ms step_avg:45.36ms step:863/1555 train_time:39156ms step_avg:45.37ms step:864/1555 train_time:39220ms step_avg:45.39ms step:865/1555 train_time:39277ms step_avg:45.41ms step:866/1555 train_time:39343ms step_avg:45.43ms step:867/1555 train_time:39401ms step_avg:45.44ms step:868/1555 train_time:39464ms step_avg:45.47ms step:869/1555 train_time:39522ms step_avg:45.48ms step:870/1555 train_time:39588ms step_avg:45.50ms step:871/1555 train_time:39645ms step_avg:45.52ms step:872/1555 train_time:39709ms step_avg:45.54ms step:873/1555 train_time:39766ms step_avg:45.55ms step:874/1555 train_time:39831ms step_avg:45.57ms step:875/1555 train_time:39889ms step_avg:45.59ms step:876/1555 train_time:39954ms step_avg:45.61ms step:877/1555 train_time:40012ms step_avg:45.62ms step:878/1555 train_time:40077ms step_avg:45.65ms step:879/1555 train_time:40135ms step_avg:45.66ms step:880/1555 train_time:40199ms step_avg:45.68ms step:881/1555 train_time:40256ms step_avg:45.69ms step:882/1555 train_time:40320ms step_avg:45.71ms step:883/1555 train_time:40378ms step_avg:45.73ms step:884/1555 train_time:40444ms step_avg:45.75ms step:885/1555 train_time:40502ms step_avg:45.77ms step:886/1555 train_time:40565ms step_avg:45.78ms step:887/1555 train_time:40623ms step_avg:45.80ms step:888/1555 train_time:40688ms step_avg:45.82ms step:889/1555 train_time:40745ms step_avg:45.83ms step:890/1555 train_time:40809ms step_avg:45.85ms step:891/1555 train_time:40866ms step_avg:45.87ms step:892/1555 train_time:40931ms step_avg:45.89ms step:893/1555 train_time:40988ms step_avg:45.90ms step:894/1555 train_time:41054ms step_avg:45.92ms step:895/1555 train_time:41113ms step_avg:45.94ms step:896/1555 train_time:41177ms step_avg:45.96ms step:897/1555 train_time:41236ms step_avg:45.97ms step:898/1555 train_time:41300ms step_avg:45.99ms step:899/1555 train_time:41357ms step_avg:46.00ms step:900/1555 train_time:41422ms step_avg:46.02ms step:901/1555 train_time:41480ms step_avg:46.04ms step:902/1555 train_time:41544ms step_avg:46.06ms step:903/1555 train_time:41602ms step_avg:46.07ms step:904/1555 train_time:41665ms step_avg:46.09ms step:905/1555 train_time:41723ms step_avg:46.10ms step:906/1555 train_time:41787ms step_avg:46.12ms step:907/1555 train_time:41845ms step_avg:46.14ms step:908/1555 train_time:41909ms step_avg:46.16ms step:909/1555 train_time:41967ms step_avg:46.17ms step:910/1555 train_time:42032ms step_avg:46.19ms step:911/1555 train_time:42089ms step_avg:46.20ms step:912/1555 train_time:42155ms step_avg:46.22ms step:913/1555 train_time:42213ms step_avg:46.24ms step:914/1555 train_time:42277ms step_avg:46.26ms step:915/1555 train_time:42336ms step_avg:46.27ms step:916/1555 train_time:42399ms step_avg:46.29ms step:917/1555 train_time:42457ms step_avg:46.30ms step:918/1555 train_time:42521ms step_avg:46.32ms step:919/1555 train_time:42579ms step_avg:46.33ms step:920/1555 train_time:42643ms step_avg:46.35ms step:921/1555 train_time:42701ms step_avg:46.36ms step:922/1555 train_time:42764ms step_avg:46.38ms step:923/1555 train_time:42822ms step_avg:46.39ms step:924/1555 train_time:42886ms step_avg:46.41ms step:925/1555 train_time:42944ms step_avg:46.43ms step:926/1555 train_time:43009ms step_avg:46.45ms step:927/1555 train_time:43067ms step_avg:46.46ms step:928/1555 train_time:43132ms step_avg:46.48ms step:929/1555 train_time:43190ms step_avg:46.49ms step:930/1555 train_time:43255ms step_avg:46.51ms step:931/1555 train_time:43313ms step_avg:46.52ms step:932/1555 train_time:43377ms step_avg:46.54ms step:933/1555 train_time:43435ms step_avg:46.55ms step:934/1555 train_time:43500ms step_avg:46.57ms step:935/1555 train_time:43557ms step_avg:46.59ms step:936/1555 train_time:43622ms step_avg:46.60ms step:937/1555 train_time:43679ms step_avg:46.62ms step:938/1555 train_time:43743ms step_avg:46.63ms step:939/1555 train_time:43801ms step_avg:46.65ms step:940/1555 train_time:43864ms step_avg:46.66ms step:941/1555 train_time:43922ms step_avg:46.68ms step:942/1555 train_time:43986ms step_avg:46.69ms step:943/1555 train_time:44044ms step_avg:46.71ms step:944/1555 train_time:44108ms step_avg:46.72ms step:945/1555 train_time:44166ms step_avg:46.74ms step:946/1555 train_time:44231ms step_avg:46.76ms step:947/1555 train_time:44288ms step_avg:46.77ms step:948/1555 train_time:44354ms step_avg:46.79ms step:949/1555 train_time:44412ms step_avg:46.80ms step:950/1555 train_time:44476ms step_avg:46.82ms step:951/1555 train_time:44535ms step_avg:46.83ms step:952/1555 train_time:44599ms step_avg:46.85ms step:953/1555 train_time:44657ms step_avg:46.86ms step:954/1555 train_time:44721ms step_avg:46.88ms step:955/1555 train_time:44778ms step_avg:46.89ms step:956/1555 train_time:44843ms step_avg:46.91ms step:957/1555 train_time:44900ms step_avg:46.92ms step:958/1555 train_time:44964ms step_avg:46.94ms step:959/1555 train_time:45022ms step_avg:46.95ms step:960/1555 train_time:45086ms step_avg:46.96ms step:961/1555 train_time:45144ms step_avg:46.98ms step:962/1555 train_time:45208ms step_avg:46.99ms step:963/1555 train_time:45266ms step_avg:47.01ms step:964/1555 train_time:45330ms step_avg:47.02ms step:965/1555 train_time:45388ms step_avg:47.03ms step:966/1555 train_time:45453ms step_avg:47.05ms step:967/1555 train_time:45512ms step_avg:47.06ms step:968/1555 train_time:45576ms step_avg:47.08ms step:969/1555 train_time:45635ms step_avg:47.09ms step:970/1555 train_time:45699ms step_avg:47.11ms step:971/1555 train_time:45756ms step_avg:47.12ms step:972/1555 train_time:45820ms step_avg:47.14ms step:973/1555 train_time:45878ms step_avg:47.15ms step:974/1555 train_time:45942ms step_avg:47.17ms step:975/1555 train_time:46000ms step_avg:47.18ms step:976/1555 train_time:46064ms step_avg:47.20ms step:977/1555 train_time:46122ms step_avg:47.21ms step:978/1555 train_time:46186ms step_avg:47.22ms step:979/1555 train_time:46244ms step_avg:47.24ms step:980/1555 train_time:46307ms step_avg:47.25ms step:981/1555 train_time:46366ms step_avg:47.26ms step:982/1555 train_time:46430ms step_avg:47.28ms step:983/1555 train_time:46488ms step_avg:47.29ms step:984/1555 train_time:46553ms step_avg:47.31ms step:985/1555 train_time:46611ms step_avg:47.32ms step:986/1555 train_time:46676ms step_avg:47.34ms step:987/1555 train_time:46734ms step_avg:47.35ms step:988/1555 train_time:46797ms step_avg:47.37ms step:989/1555 train_time:46855ms step_avg:47.38ms step:990/1555 train_time:46919ms step_avg:47.39ms step:991/1555 train_time:46977ms step_avg:47.40ms step:992/1555 train_time:47042ms step_avg:47.42ms step:993/1555 train_time:47101ms step_avg:47.43ms step:994/1555 train_time:47164ms step_avg:47.45ms step:995/1555 train_time:47222ms step_avg:47.46ms step:996/1555 train_time:47285ms step_avg:47.47ms step:997/1555 train_time:47343ms step_avg:47.49ms step:998/1555 train_time:47408ms step_avg:47.50ms step:999/1555 train_time:47466ms step_avg:47.51ms step:1000/1555 train_time:47531ms step_avg:47.53ms step:1000/1555 val_loss:3.5680 train_time:47614ms step_avg:47.61ms step:1001/1555 train_time:47636ms step_avg:47.59ms step:1002/1555 train_time:47657ms step_avg:47.56ms step:1003/1555 train_time:47717ms step_avg:47.57ms step:1004/1555 train_time:47782ms step_avg:47.59ms step:1005/1555 train_time:47842ms step_avg:47.60ms step:1006/1555 train_time:47906ms step_avg:47.62ms step:1007/1555 train_time:47964ms step_avg:47.63ms step:1008/1555 train_time:48028ms step_avg:47.65ms step:1009/1555 train_time:48085ms step_avg:47.66ms step:1010/1555 train_time:48149ms step_avg:47.67ms step:1011/1555 train_time:48208ms step_avg:47.68ms step:1012/1555 train_time:48294ms step_avg:47.72ms step:1013/1555 train_time:48378ms step_avg:47.76ms step:1014/1555 train_time:48468ms step_avg:47.80ms step:1015/1555 train_time:48552ms step_avg:47.83ms step:1016/1555 train_time:48644ms step_avg:47.88ms step:1017/1555 train_time:48731ms step_avg:47.92ms step:1018/1555 train_time:48823ms step_avg:47.96ms step:1019/1555 train_time:48907ms step_avg:48.00ms step:1020/1555 train_time:48999ms step_avg:48.04ms step:1021/1555 train_time:49082ms step_avg:48.07ms step:1022/1555 train_time:49170ms step_avg:48.11ms step:1023/1555 train_time:49254ms step_avg:48.15ms step:1024/1555 train_time:49343ms step_avg:48.19ms step:1025/1555 train_time:49425ms step_avg:48.22ms step:1026/1555 train_time:49514ms step_avg:48.26ms step:1027/1555 train_time:49600ms step_avg:48.30ms step:1028/1555 train_time:49690ms step_avg:48.34ms step:1029/1555 train_time:49777ms step_avg:48.37ms step:1030/1555 train_time:49867ms step_avg:48.41ms step:1031/1555 train_time:49952ms step_avg:48.45ms step:1032/1555 train_time:50043ms step_avg:48.49ms step:1033/1555 train_time:50126ms step_avg:48.52ms step:1034/1555 train_time:50215ms step_avg:48.56ms step:1035/1555 train_time:50299ms step_avg:48.60ms step:1036/1555 train_time:50388ms step_avg:48.64ms step:1037/1555 train_time:50471ms step_avg:48.67ms step:1038/1555 train_time:50563ms step_avg:48.71ms step:1039/1555 train_time:50646ms step_avg:48.75ms step:1040/1555 train_time:50737ms step_avg:48.79ms step:1041/1555 train_time:50822ms step_avg:48.82ms step:1042/1555 train_time:50913ms step_avg:48.86ms step:1043/1555 train_time:50998ms step_avg:48.90ms step:1044/1555 train_time:51088ms step_avg:48.93ms step:1045/1555 train_time:51172ms step_avg:48.97ms step:1046/1555 train_time:51263ms step_avg:49.01ms step:1047/1555 train_time:51346ms step_avg:49.04ms step:1048/1555 train_time:51436ms step_avg:49.08ms step:1049/1555 train_time:51519ms step_avg:49.11ms step:1050/1555 train_time:51608ms step_avg:49.15ms step:1051/1555 train_time:51695ms step_avg:49.19ms step:1052/1555 train_time:51785ms step_avg:49.23ms step:1053/1555 train_time:51870ms step_avg:49.26ms step:1054/1555 train_time:51961ms step_avg:49.30ms step:1055/1555 train_time:52044ms step_avg:49.33ms step:1056/1555 train_time:52134ms step_avg:49.37ms step:1057/1555 train_time:52220ms step_avg:49.40ms step:1058/1555 train_time:52310ms step_avg:49.44ms step:1059/1555 train_time:52392ms step_avg:49.47ms step:1060/1555 train_time:52483ms step_avg:49.51ms step:1061/1555 train_time:52564ms step_avg:49.54ms step:1062/1555 train_time:52656ms step_avg:49.58ms step:1063/1555 train_time:52741ms step_avg:49.62ms step:1064/1555 train_time:52830ms step_avg:49.65ms step:1065/1555 train_time:52915ms step_avg:49.69ms step:1066/1555 train_time:53007ms step_avg:49.73ms step:1067/1555 train_time:53089ms step_avg:49.76ms step:1068/1555 train_time:53181ms step_avg:49.79ms step:1069/1555 train_time:53264ms step_avg:49.83ms step:1070/1555 train_time:53355ms step_avg:49.86ms step:1071/1555 train_time:53438ms step_avg:49.90ms step:1072/1555 train_time:53527ms step_avg:49.93ms step:1073/1555 train_time:53611ms step_avg:49.96ms step:1074/1555 train_time:53702ms step_avg:50.00ms step:1075/1555 train_time:53786ms step_avg:50.03ms step:1076/1555 train_time:53876ms step_avg:50.07ms step:1077/1555 train_time:53960ms step_avg:50.10ms step:1078/1555 train_time:54049ms step_avg:50.14ms step:1079/1555 train_time:54134ms step_avg:50.17ms step:1080/1555 train_time:54224ms step_avg:50.21ms step:1081/1555 train_time:54308ms step_avg:50.24ms step:1082/1555 train_time:54399ms step_avg:50.28ms step:1083/1555 train_time:54482ms step_avg:50.31ms step:1084/1555 train_time:54572ms step_avg:50.34ms step:1085/1555 train_time:54656ms step_avg:50.37ms step:1086/1555 train_time:54746ms step_avg:50.41ms step:1087/1555 train_time:54831ms step_avg:50.44ms step:1088/1555 train_time:54922ms step_avg:50.48ms step:1089/1555 train_time:55005ms step_avg:50.51ms step:1090/1555 train_time:55095ms step_avg:50.55ms step:1091/1555 train_time:55180ms step_avg:50.58ms step:1092/1555 train_time:55269ms step_avg:50.61ms step:1093/1555 train_time:55353ms step_avg:50.64ms step:1094/1555 train_time:55443ms step_avg:50.68ms step:1095/1555 train_time:55526ms step_avg:50.71ms step:1096/1555 train_time:55617ms step_avg:50.75ms step:1097/1555 train_time:55701ms step_avg:50.78ms step:1098/1555 train_time:55790ms step_avg:50.81ms step:1099/1555 train_time:55875ms step_avg:50.84ms step:1100/1555 train_time:55964ms step_avg:50.88ms step:1101/1555 train_time:56048ms step_avg:50.91ms step:1102/1555 train_time:56140ms step_avg:50.94ms step:1103/1555 train_time:56224ms step_avg:50.97ms step:1104/1555 train_time:56313ms step_avg:51.01ms step:1105/1555 train_time:56397ms step_avg:51.04ms step:1106/1555 train_time:56486ms step_avg:51.07ms step:1107/1555 train_time:56570ms step_avg:51.10ms step:1108/1555 train_time:56660ms step_avg:51.14ms step:1109/1555 train_time:56743ms step_avg:51.17ms step:1110/1555 train_time:56834ms step_avg:51.20ms step:1111/1555 train_time:56917ms step_avg:51.23ms step:1112/1555 train_time:57007ms step_avg:51.26ms step:1113/1555 train_time:57091ms step_avg:51.29ms step:1114/1555 train_time:57183ms step_avg:51.33ms step:1115/1555 train_time:57265ms step_avg:51.36ms step:1116/1555 train_time:57355ms step_avg:51.39ms step:1117/1555 train_time:57439ms step_avg:51.42ms step:1118/1555 train_time:57528ms step_avg:51.46ms step:1119/1555 train_time:57613ms step_avg:51.49ms step:1120/1555 train_time:57704ms step_avg:51.52ms step:1121/1555 train_time:57787ms step_avg:51.55ms step:1122/1555 train_time:57878ms step_avg:51.58ms step:1123/1555 train_time:57962ms step_avg:51.61ms step:1124/1555 train_time:58051ms step_avg:51.65ms step:1125/1555 train_time:58136ms step_avg:51.68ms step:1126/1555 train_time:58225ms step_avg:51.71ms step:1127/1555 train_time:58309ms step_avg:51.74ms step:1128/1555 train_time:58399ms step_avg:51.77ms step:1129/1555 train_time:58483ms step_avg:51.80ms step:1130/1555 train_time:58573ms step_avg:51.83ms step:1131/1555 train_time:58655ms step_avg:51.86ms step:1132/1555 train_time:58745ms step_avg:51.90ms step:1133/1555 train_time:58828ms step_avg:51.92ms step:1134/1555 train_time:58920ms step_avg:51.96ms step:1135/1555 train_time:59003ms step_avg:51.98ms step:1136/1555 train_time:59093ms step_avg:52.02ms step:1137/1555 train_time:59178ms step_avg:52.05ms step:1138/1555 train_time:59268ms step_avg:52.08ms step:1139/1555 train_time:59353ms step_avg:52.11ms step:1140/1555 train_time:59443ms step_avg:52.14ms step:1141/1555 train_time:59526ms step_avg:52.17ms step:1142/1555 train_time:59618ms step_avg:52.20ms step:1143/1555 train_time:59702ms step_avg:52.23ms step:1144/1555 train_time:59791ms step_avg:52.27ms step:1145/1555 train_time:59877ms step_avg:52.29ms step:1146/1555 train_time:59965ms step_avg:52.33ms step:1147/1555 train_time:60049ms step_avg:52.35ms step:1148/1555 train_time:60140ms step_avg:52.39ms step:1149/1555 train_time:60224ms step_avg:52.41ms step:1150/1555 train_time:60314ms step_avg:52.45ms step:1151/1555 train_time:60398ms step_avg:52.47ms step:1152/1555 train_time:60487ms step_avg:52.51ms step:1153/1555 train_time:60571ms step_avg:52.53ms step:1154/1555 train_time:60663ms step_avg:52.57ms step:1155/1555 train_time:60745ms step_avg:52.59ms step:1156/1555 train_time:60836ms step_avg:52.63ms step:1157/1555 train_time:60920ms step_avg:52.65ms step:1158/1555 train_time:61009ms step_avg:52.68ms step:1159/1555 train_time:61094ms step_avg:52.71ms step:1160/1555 train_time:61185ms step_avg:52.75ms step:1161/1555 train_time:61270ms step_avg:52.77ms step:1162/1555 train_time:61359ms step_avg:52.80ms step:1163/1555 train_time:61442ms step_avg:52.83ms step:1164/1555 train_time:61533ms step_avg:52.86ms step:1165/1555 train_time:61617ms step_avg:52.89ms step:1166/1555 train_time:61707ms step_avg:52.92ms step:1167/1555 train_time:61792ms step_avg:52.95ms step:1168/1555 train_time:61883ms step_avg:52.98ms step:1169/1555 train_time:61966ms step_avg:53.01ms step:1170/1555 train_time:62056ms step_avg:53.04ms step:1171/1555 train_time:62140ms step_avg:53.07ms step:1172/1555 train_time:62230ms step_avg:53.10ms step:1173/1555 train_time:62315ms step_avg:53.12ms step:1174/1555 train_time:62404ms step_avg:53.16ms step:1175/1555 train_time:62488ms step_avg:53.18ms step:1176/1555 train_time:62581ms step_avg:53.22ms step:1177/1555 train_time:62664ms step_avg:53.24ms step:1178/1555 train_time:62754ms step_avg:53.27ms step:1179/1555 train_time:62839ms step_avg:53.30ms step:1180/1555 train_time:62929ms step_avg:53.33ms step:1181/1555 train_time:63014ms step_avg:53.36ms step:1182/1555 train_time:63104ms step_avg:53.39ms step:1183/1555 train_time:63187ms step_avg:53.41ms step:1184/1555 train_time:63278ms step_avg:53.44ms step:1185/1555 train_time:63361ms step_avg:53.47ms step:1186/1555 train_time:63450ms step_avg:53.50ms step:1187/1555 train_time:63535ms step_avg:53.53ms step:1188/1555 train_time:63625ms step_avg:53.56ms step:1189/1555 train_time:63708ms step_avg:53.58ms step:1190/1555 train_time:63800ms step_avg:53.61ms step:1191/1555 train_time:63883ms step_avg:53.64ms step:1192/1555 train_time:63973ms step_avg:53.67ms step:1193/1555 train_time:64058ms step_avg:53.69ms step:1194/1555 train_time:64148ms step_avg:53.72ms step:1195/1555 train_time:64232ms step_avg:53.75ms step:1196/1555 train_time:64324ms step_avg:53.78ms step:1197/1555 train_time:64408ms step_avg:53.81ms step:1198/1555 train_time:64498ms step_avg:53.84ms step:1199/1555 train_time:64582ms step_avg:53.86ms step:1200/1555 train_time:64671ms step_avg:53.89ms step:1201/1555 train_time:64756ms step_avg:53.92ms step:1202/1555 train_time:64845ms step_avg:53.95ms step:1203/1555 train_time:64929ms step_avg:53.97ms step:1204/1555 train_time:65021ms step_avg:54.00ms step:1205/1555 train_time:65104ms step_avg:54.03ms step:1206/1555 train_time:65193ms step_avg:54.06ms step:1207/1555 train_time:65278ms step_avg:54.08ms step:1208/1555 train_time:65367ms step_avg:54.11ms step:1209/1555 train_time:65451ms step_avg:54.14ms step:1210/1555 train_time:65541ms step_avg:54.17ms step:1211/1555 train_time:65624ms step_avg:54.19ms step:1212/1555 train_time:65715ms step_avg:54.22ms step:1213/1555 train_time:65799ms step_avg:54.25ms step:1214/1555 train_time:65888ms step_avg:54.27ms step:1215/1555 train_time:65973ms step_avg:54.30ms step:1216/1555 train_time:66063ms step_avg:54.33ms step:1217/1555 train_time:66147ms step_avg:54.35ms step:1218/1555 train_time:66238ms step_avg:54.38ms step:1219/1555 train_time:66321ms step_avg:54.41ms step:1220/1555 train_time:66410ms step_avg:54.43ms step:1221/1555 train_time:66497ms step_avg:54.46ms step:1222/1555 train_time:66586ms step_avg:54.49ms step:1223/1555 train_time:66669ms step_avg:54.51ms step:1224/1555 train_time:66760ms step_avg:54.54ms step:1225/1555 train_time:66844ms step_avg:54.57ms step:1226/1555 train_time:66934ms step_avg:54.60ms step:1227/1555 train_time:67018ms step_avg:54.62ms step:1228/1555 train_time:67108ms step_avg:54.65ms step:1229/1555 train_time:67193ms step_avg:54.67ms step:1230/1555 train_time:67284ms step_avg:54.70ms step:1231/1555 train_time:67367ms step_avg:54.73ms step:1232/1555 train_time:67458ms step_avg:54.75ms step:1233/1555 train_time:67542ms step_avg:54.78ms step:1234/1555 train_time:67631ms step_avg:54.81ms step:1235/1555 train_time:67715ms step_avg:54.83ms step:1236/1555 train_time:67805ms step_avg:54.86ms step:1237/1555 train_time:67889ms step_avg:54.88ms step:1238/1555 train_time:67981ms step_avg:54.91ms step:1239/1555 train_time:68064ms step_avg:54.93ms step:1240/1555 train_time:68153ms step_avg:54.96ms step:1241/1555 train_time:68237ms step_avg:54.99ms step:1242/1555 train_time:68326ms step_avg:55.01ms step:1243/1555 train_time:68410ms step_avg:55.04ms step:1244/1555 train_time:68501ms step_avg:55.06ms step:1245/1555 train_time:68585ms step_avg:55.09ms step:1246/1555 train_time:68675ms step_avg:55.12ms step:1247/1555 train_time:68759ms step_avg:55.14ms step:1248/1555 train_time:68848ms step_avg:55.17ms step:1249/1555 train_time:68933ms step_avg:55.19ms step:1250/1555 train_time:69024ms step_avg:55.22ms step:1250/1555 val_loss:3.3954 train_time:69139ms step_avg:55.31ms step:1251/1555 train_time:69161ms step_avg:55.28ms step:1252/1555 train_time:69201ms step_avg:55.27ms step:1253/1555 train_time:69291ms step_avg:55.30ms step:1254/1555 train_time:69382ms step_avg:55.33ms step:1255/1555 train_time:69466ms step_avg:55.35ms step:1256/1555 train_time:69555ms step_avg:55.38ms step:1257/1555 train_time:69639ms step_avg:55.40ms step:1258/1555 train_time:69728ms step_avg:55.43ms step:1259/1555 train_time:69811ms step_avg:55.45ms step:1260/1555 train_time:69899ms step_avg:55.48ms step:1261/1555 train_time:69983ms step_avg:55.50ms step:1262/1555 train_time:70071ms step_avg:55.52ms step:1263/1555 train_time:70156ms step_avg:55.55ms step:1264/1555 train_time:70250ms step_avg:55.58ms step:1265/1555 train_time:70336ms step_avg:55.60ms step:1266/1555 train_time:70428ms step_avg:55.63ms step:1267/1555 train_time:70512ms step_avg:55.65ms step:1268/1555 train_time:70601ms step_avg:55.68ms step:1269/1555 train_time:70685ms step_avg:55.70ms step:1270/1555 train_time:70773ms step_avg:55.73ms step:1271/1555 train_time:70856ms step_avg:55.75ms step:1272/1555 train_time:70947ms step_avg:55.78ms step:1273/1555 train_time:71030ms step_avg:55.80ms step:1274/1555 train_time:71120ms step_avg:55.82ms step:1275/1555 train_time:71208ms step_avg:55.85ms step:1276/1555 train_time:71297ms step_avg:55.88ms step:1277/1555 train_time:71383ms step_avg:55.90ms step:1278/1555 train_time:71473ms step_avg:55.93ms step:1279/1555 train_time:71558ms step_avg:55.95ms step:1280/1555 train_time:71648ms step_avg:55.97ms step:1281/1555 train_time:71731ms step_avg:56.00ms step:1282/1555 train_time:71820ms step_avg:56.02ms step:1283/1555 train_time:71904ms step_avg:56.04ms step:1284/1555 train_time:71993ms step_avg:56.07ms step:1285/1555 train_time:72077ms step_avg:56.09ms step:1286/1555 train_time:72168ms step_avg:56.12ms step:1287/1555 train_time:72252ms step_avg:56.14ms step:1288/1555 train_time:72344ms step_avg:56.17ms step:1289/1555 train_time:72429ms step_avg:56.19ms step:1290/1555 train_time:72519ms step_avg:56.22ms step:1291/1555 train_time:72603ms step_avg:56.24ms step:1292/1555 train_time:72693ms step_avg:56.26ms step:1293/1555 train_time:72776ms step_avg:56.28ms step:1294/1555 train_time:72866ms step_avg:56.31ms step:1295/1555 train_time:72949ms step_avg:56.33ms step:1296/1555 train_time:73039ms step_avg:56.36ms step:1297/1555 train_time:73124ms step_avg:56.38ms step:1298/1555 train_time:73215ms step_avg:56.41ms step:1299/1555 train_time:73301ms step_avg:56.43ms step:1300/1555 train_time:73394ms step_avg:56.46ms step:1301/1555 train_time:73476ms step_avg:56.48ms step:1302/1555 train_time:73568ms step_avg:56.50ms step:1303/1555 train_time:73651ms step_avg:56.52ms step:1304/1555 train_time:73740ms step_avg:56.55ms step:1305/1555 train_time:73824ms step_avg:56.57ms step:1306/1555 train_time:73913ms step_avg:56.59ms step:1307/1555 train_time:73996ms step_avg:56.62ms step:1308/1555 train_time:74089ms step_avg:56.64ms step:1309/1555 train_time:74170ms step_avg:56.66ms step:1310/1555 train_time:74261ms step_avg:56.69ms step:1311/1555 train_time:74346ms step_avg:56.71ms step:1312/1555 train_time:74436ms step_avg:56.73ms step:1313/1555 train_time:74522ms step_avg:56.76ms step:1314/1555 train_time:74612ms step_avg:56.78ms step:1315/1555 train_time:74696ms step_avg:56.80ms step:1316/1555 train_time:74787ms step_avg:56.83ms step:1317/1555 train_time:74870ms step_avg:56.85ms step:1318/1555 train_time:74959ms step_avg:56.87ms step:1319/1555 train_time:75044ms step_avg:56.89ms step:1320/1555 train_time:75134ms step_avg:56.92ms step:1321/1555 train_time:75218ms step_avg:56.94ms step:1322/1555 train_time:75310ms step_avg:56.97ms step:1323/1555 train_time:75393ms step_avg:56.99ms step:1324/1555 train_time:75483ms step_avg:57.01ms step:1325/1555 train_time:75567ms step_avg:57.03ms step:1326/1555 train_time:75657ms step_avg:57.06ms step:1327/1555 train_time:75742ms step_avg:57.08ms step:1328/1555 train_time:75832ms step_avg:57.10ms step:1329/1555 train_time:75916ms step_avg:57.12ms step:1330/1555 train_time:76007ms step_avg:57.15ms step:1331/1555 train_time:76090ms step_avg:57.17ms step:1332/1555 train_time:76180ms step_avg:57.19ms step:1333/1555 train_time:76265ms step_avg:57.21ms step:1334/1555 train_time:76354ms step_avg:57.24ms step:1335/1555 train_time:76438ms step_avg:57.26ms step:1336/1555 train_time:76529ms step_avg:57.28ms step:1337/1555 train_time:76612ms step_avg:57.30ms step:1338/1555 train_time:76702ms step_avg:57.33ms step:1339/1555 train_time:76786ms step_avg:57.35ms step:1340/1555 train_time:76874ms step_avg:57.37ms step:1341/1555 train_time:76958ms step_avg:57.39ms step:1342/1555 train_time:77049ms step_avg:57.41ms step:1343/1555 train_time:77132ms step_avg:57.43ms step:1344/1555 train_time:77223ms step_avg:57.46ms step:1345/1555 train_time:77308ms step_avg:57.48ms step:1346/1555 train_time:77398ms step_avg:57.50ms step:1347/1555 train_time:77483ms step_avg:57.52ms step:1348/1555 train_time:77572ms step_avg:57.55ms step:1349/1555 train_time:77656ms step_avg:57.57ms step:1350/1555 train_time:77747ms step_avg:57.59ms step:1351/1555 train_time:77830ms step_avg:57.61ms step:1352/1555 train_time:77920ms step_avg:57.63ms step:1353/1555 train_time:78005ms step_avg:57.65ms step:1354/1555 train_time:78094ms step_avg:57.68ms step:1355/1555 train_time:78178ms step_avg:57.70ms step:1356/1555 train_time:78268ms step_avg:57.72ms step:1357/1555 train_time:78352ms step_avg:57.74ms step:1358/1555 train_time:78442ms step_avg:57.76ms step:1359/1555 train_time:78527ms step_avg:57.78ms step:1360/1555 train_time:78617ms step_avg:57.81ms step:1361/1555 train_time:78702ms step_avg:57.83ms step:1362/1555 train_time:78793ms step_avg:57.85ms step:1363/1555 train_time:78876ms step_avg:57.87ms step:1364/1555 train_time:78967ms step_avg:57.89ms step:1365/1555 train_time:79050ms step_avg:57.91ms step:1366/1555 train_time:79139ms step_avg:57.94ms step:1367/1555 train_time:79223ms step_avg:57.95ms step:1368/1555 train_time:79313ms step_avg:57.98ms step:1369/1555 train_time:79398ms step_avg:58.00ms step:1370/1555 train_time:79490ms step_avg:58.02ms step:1371/1555 train_time:79574ms step_avg:58.04ms step:1372/1555 train_time:79664ms step_avg:58.06ms step:1373/1555 train_time:79748ms step_avg:58.08ms step:1374/1555 train_time:79838ms step_avg:58.11ms step:1375/1555 train_time:79922ms step_avg:58.12ms step:1376/1555 train_time:80012ms step_avg:58.15ms step:1377/1555 train_time:80095ms step_avg:58.17ms step:1378/1555 train_time:80187ms step_avg:58.19ms step:1379/1555 train_time:80270ms step_avg:58.21ms step:1380/1555 train_time:80360ms step_avg:58.23ms step:1381/1555 train_time:80446ms step_avg:58.25ms step:1382/1555 train_time:80535ms step_avg:58.27ms step:1383/1555 train_time:80619ms step_avg:58.29ms step:1384/1555 train_time:80710ms step_avg:58.32ms step:1385/1555 train_time:80794ms step_avg:58.33ms step:1386/1555 train_time:80883ms step_avg:58.36ms step:1387/1555 train_time:80967ms step_avg:58.38ms step:1388/1555 train_time:81057ms step_avg:58.40ms step:1389/1555 train_time:81141ms step_avg:58.42ms step:1390/1555 train_time:81231ms step_avg:58.44ms step:1391/1555 train_time:81315ms step_avg:58.46ms step:1392/1555 train_time:81406ms step_avg:58.48ms step:1393/1555 train_time:81490ms step_avg:58.50ms step:1394/1555 train_time:81580ms step_avg:58.52ms step:1395/1555 train_time:81665ms step_avg:58.54ms step:1396/1555 train_time:81754ms step_avg:58.56ms step:1397/1555 train_time:81839ms step_avg:58.58ms step:1398/1555 train_time:81929ms step_avg:58.60ms step:1399/1555 train_time:82013ms step_avg:58.62ms step:1400/1555 train_time:82102ms step_avg:58.64ms step:1401/1555 train_time:82186ms step_avg:58.66ms step:1402/1555 train_time:82275ms step_avg:58.68ms step:1403/1555 train_time:82359ms step_avg:58.70ms step:1404/1555 train_time:82451ms step_avg:58.73ms step:1405/1555 train_time:82534ms step_avg:58.74ms step:1406/1555 train_time:82624ms step_avg:58.77ms step:1407/1555 train_time:82709ms step_avg:58.78ms step:1408/1555 train_time:82798ms step_avg:58.81ms step:1409/1555 train_time:82883ms step_avg:58.82ms step:1410/1555 train_time:82972ms step_avg:58.85ms step:1411/1555 train_time:83056ms step_avg:58.86ms step:1412/1555 train_time:83147ms step_avg:58.89ms step:1413/1555 train_time:83230ms step_avg:58.90ms step:1414/1555 train_time:83320ms step_avg:58.92ms step:1415/1555 train_time:83405ms step_avg:58.94ms step:1416/1555 train_time:83494ms step_avg:58.96ms step:1417/1555 train_time:83577ms step_avg:58.98ms step:1418/1555 train_time:83669ms step_avg:59.00ms step:1419/1555 train_time:83754ms step_avg:59.02ms step:1420/1555 train_time:83843ms step_avg:59.04ms step:1421/1555 train_time:83927ms step_avg:59.06ms step:1422/1555 train_time:84016ms step_avg:59.08ms step:1423/1555 train_time:84100ms step_avg:59.10ms step:1424/1555 train_time:84191ms step_avg:59.12ms step:1425/1555 train_time:84274ms step_avg:59.14ms step:1426/1555 train_time:84364ms step_avg:59.16ms step:1427/1555 train_time:84449ms step_avg:59.18ms step:1428/1555 train_time:84537ms step_avg:59.20ms step:1429/1555 train_time:84623ms step_avg:59.22ms step:1430/1555 train_time:84713ms step_avg:59.24ms step:1431/1555 train_time:84796ms step_avg:59.26ms step:1432/1555 train_time:84887ms step_avg:59.28ms step:1433/1555 train_time:84970ms step_avg:59.30ms step:1434/1555 train_time:85060ms step_avg:59.32ms step:1435/1555 train_time:85144ms step_avg:59.33ms step:1436/1555 train_time:85233ms step_avg:59.35ms step:1437/1555 train_time:85316ms step_avg:59.37ms step:1438/1555 train_time:85406ms step_avg:59.39ms step:1439/1555 train_time:85491ms step_avg:59.41ms step:1440/1555 train_time:85580ms step_avg:59.43ms step:1441/1555 train_time:85663ms step_avg:59.45ms step:1442/1555 train_time:85753ms step_avg:59.47ms step:1443/1555 train_time:85837ms step_avg:59.49ms step:1444/1555 train_time:85928ms step_avg:59.51ms step:1445/1555 train_time:86012ms step_avg:59.52ms step:1446/1555 train_time:86101ms step_avg:59.54ms step:1447/1555 train_time:86186ms step_avg:59.56ms step:1448/1555 train_time:86275ms step_avg:59.58ms step:1449/1555 train_time:86360ms step_avg:59.60ms step:1450/1555 train_time:86450ms step_avg:59.62ms step:1451/1555 train_time:86532ms step_avg:59.64ms step:1452/1555 train_time:86623ms step_avg:59.66ms step:1453/1555 train_time:86708ms step_avg:59.67ms step:1454/1555 train_time:86797ms step_avg:59.70ms step:1455/1555 train_time:86882ms step_avg:59.71ms step:1456/1555 train_time:86972ms step_avg:59.73ms step:1457/1555 train_time:87056ms step_avg:59.75ms step:1458/1555 train_time:87147ms step_avg:59.77ms step:1459/1555 train_time:87230ms step_avg:59.79ms step:1460/1555 train_time:87319ms step_avg:59.81ms step:1461/1555 train_time:87405ms step_avg:59.83ms step:1462/1555 train_time:87494ms step_avg:59.85ms step:1463/1555 train_time:87577ms step_avg:59.86ms step:1464/1555 train_time:87668ms step_avg:59.88ms step:1465/1555 train_time:87751ms step_avg:59.90ms step:1466/1555 train_time:87841ms step_avg:59.92ms step:1467/1555 train_time:87927ms step_avg:59.94ms step:1468/1555 train_time:88018ms step_avg:59.96ms step:1469/1555 train_time:88102ms step_avg:59.97ms step:1470/1555 train_time:88192ms step_avg:59.99ms step:1471/1555 train_time:88275ms step_avg:60.01ms step:1472/1555 train_time:88365ms step_avg:60.03ms step:1473/1555 train_time:88450ms step_avg:60.05ms step:1474/1555 train_time:88538ms step_avg:60.07ms step:1475/1555 train_time:88623ms step_avg:60.08ms step:1476/1555 train_time:88712ms step_avg:60.10ms step:1477/1555 train_time:88796ms step_avg:60.12ms step:1478/1555 train_time:88888ms step_avg:60.14ms step:1479/1555 train_time:88970ms step_avg:60.16ms step:1480/1555 train_time:89060ms step_avg:60.18ms step:1481/1555 train_time:89146ms step_avg:60.19ms step:1482/1555 train_time:89236ms step_avg:60.21ms step:1483/1555 train_time:89320ms step_avg:60.23ms step:1484/1555 train_time:89411ms step_avg:60.25ms step:1485/1555 train_time:89495ms step_avg:60.27ms step:1486/1555 train_time:89585ms step_avg:60.29ms step:1487/1555 train_time:89668ms step_avg:60.30ms step:1488/1555 train_time:89757ms step_avg:60.32ms step:1489/1555 train_time:89842ms step_avg:60.34ms step:1490/1555 train_time:89933ms step_avg:60.36ms step:1491/1555 train_time:90017ms step_avg:60.37ms step:1492/1555 train_time:90106ms step_avg:60.39ms step:1493/1555 train_time:90191ms step_avg:60.41ms step:1494/1555 train_time:90280ms step_avg:60.43ms step:1495/1555 train_time:90363ms step_avg:60.44ms step:1496/1555 train_time:90453ms step_avg:60.46ms step:1497/1555 train_time:90537ms step_avg:60.48ms step:1498/1555 train_time:90629ms step_avg:60.50ms step:1499/1555 train_time:90712ms step_avg:60.51ms step:1500/1555 train_time:90802ms step_avg:60.53ms step:1500/1555 val_loss:3.2932 train_time:90916ms step_avg:60.61ms step:1501/1555 train_time:90942ms step_avg:60.59ms step:1502/1555 train_time:90978ms step_avg:60.57ms step:1503/1555 train_time:91061ms step_avg:60.59ms step:1504/1555 train_time:91153ms step_avg:60.61ms step:1505/1555 train_time:91238ms step_avg:60.62ms step:1506/1555 train_time:91329ms step_avg:60.64ms step:1507/1555 train_time:91411ms step_avg:60.66ms step:1508/1555 train_time:91500ms step_avg:60.68ms step:1509/1555 train_time:91583ms step_avg:60.69ms step:1510/1555 train_time:91673ms step_avg:60.71ms step:1511/1555 train_time:91757ms step_avg:60.73ms step:1512/1555 train_time:91844ms step_avg:60.74ms step:1513/1555 train_time:91931ms step_avg:60.76ms step:1514/1555 train_time:92022ms step_avg:60.78ms step:1515/1555 train_time:92108ms step_avg:60.80ms step:1516/1555 train_time:92203ms step_avg:60.82ms step:1517/1555 train_time:92288ms step_avg:60.84ms step:1518/1555 train_time:92377ms step_avg:60.85ms step:1519/1555 train_time:92461ms step_avg:60.87ms step:1520/1555 train_time:92550ms step_avg:60.89ms step:1521/1555 train_time:92633ms step_avg:60.90ms step:1522/1555 train_time:92722ms step_avg:60.92ms step:1523/1555 train_time:92806ms step_avg:60.94ms step:1524/1555 train_time:92898ms step_avg:60.96ms step:1525/1555 train_time:92983ms step_avg:60.97ms step:1526/1555 train_time:93077ms step_avg:60.99ms step:1527/1555 train_time:93161ms step_avg:61.01ms step:1528/1555 train_time:93250ms step_avg:61.03ms step:1529/1555 train_time:93335ms step_avg:61.04ms step:1530/1555 train_time:93424ms step_avg:61.06ms step:1531/1555 train_time:93508ms step_avg:61.08ms step:1532/1555 train_time:93600ms step_avg:61.10ms step:1533/1555 train_time:93683ms step_avg:61.11ms step:1534/1555 train_time:93773ms step_avg:61.13ms step:1535/1555 train_time:93857ms step_avg:61.14ms step:1536/1555 train_time:93947ms step_avg:61.16ms step:1537/1555 train_time:94034ms step_avg:61.18ms step:1538/1555 train_time:94123ms step_avg:61.20ms step:1539/1555 train_time:94209ms step_avg:61.21ms step:1540/1555 train_time:94300ms step_avg:61.23ms step:1541/1555 train_time:94383ms step_avg:61.25ms step:1542/1555 train_time:94475ms step_avg:61.27ms step:1543/1555 train_time:94558ms step_avg:61.28ms step:1544/1555 train_time:94648ms step_avg:61.30ms step:1545/1555 train_time:94732ms step_avg:61.32ms step:1546/1555 train_time:94821ms step_avg:61.33ms step:1547/1555 train_time:94907ms step_avg:61.35ms step:1548/1555 train_time:94998ms step_avg:61.37ms step:1549/1555 train_time:95082ms step_avg:61.38ms step:1550/1555 train_time:95174ms step_avg:61.40ms step:1551/1555 train_time:95258ms step_avg:61.42ms step:1552/1555 train_time:95347ms step_avg:61.44ms step:1553/1555 train_time:95432ms step_avg:61.45ms step:1554/1555 train_time:95523ms step_avg:61.47ms step:1555/1555 train_time:95608ms step_avg:61.48ms step:1555/1555 val_loss:3.2771 train_time:95723ms step_avg:61.56ms peak memory allocated: 31630 MiB reserved: 46538 MiB