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 10:30:04 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 29C P0 114W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 30C P0 118W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:71:00.0 Off | 0 | | N/A 32C 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 116W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:87:00.0 Off | 0 | | N/A 32C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:8D:00.0 Off | 0 | | N/A 31C P0 119W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:95:00.0 Off | 0 | | N/A 30C P0 116W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ +-----------------------------------------------------------------------------------------+ | Processes: | | GPU GI CI PID Type Process name GPU Memory | | ID ID Usage | |=========================================================================================| | 0 N/A N/A 94 C /usr/local/bin/python 1510MiB | | 1 N/A N/A 95 C /usr/local/bin/python 1510MiB | | 2 N/A N/A 96 C /usr/local/bin/python 1510MiB | | 3 N/A N/A 97 C /usr/local/bin/python 1510MiB | | 4 N/A N/A 98 C /usr/local/bin/python 1510MiB | | 5 N/A N/A 99 C /usr/local/bin/python 1510MiB | | 6 N/A N/A 100 C /usr/local/bin/python 1510MiB | | 7 N/A N/A 101 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:87ms step_avg:87.43ms step:2/1555 train_time:109ms step_avg:54.31ms step:3/1555 train_time:128ms step_avg:42.67ms step:4/1555 train_time:154ms step_avg:38.46ms step:5/1555 train_time:185ms step_avg:36.94ms step:6/1555 train_time:222ms step_avg:37.00ms step:7/1555 train_time:253ms step_avg:36.14ms step:8/1555 train_time:290ms step_avg:36.30ms step:9/1555 train_time:322ms step_avg:35.76ms step:10/1555 train_time:360ms step_avg:35.97ms step:11/1555 train_time:390ms step_avg:35.50ms step:12/1555 train_time:428ms step_avg:35.70ms step:13/1555 train_time:459ms step_avg:35.34ms step:14/1555 train_time:497ms step_avg:35.50ms step:15/1555 train_time:528ms step_avg:35.21ms step:16/1555 train_time:566ms step_avg:35.35ms step:17/1555 train_time:597ms step_avg:35.09ms step:18/1555 train_time:634ms step_avg:35.23ms step:19/1555 train_time:665ms step_avg:35.00ms step:20/1555 train_time:702ms step_avg:35.12ms step:21/1555 train_time:734ms step_avg:34.94ms step:22/1555 train_time:771ms step_avg:35.05ms step:23/1555 train_time:802ms step_avg:34.88ms step:24/1555 train_time:840ms step_avg:34.99ms step:25/1555 train_time:871ms step_avg:34.83ms step:26/1555 train_time:908ms step_avg:34.94ms step:27/1555 train_time:939ms step_avg:34.79ms step:28/1555 train_time:977ms step_avg:34.89ms step:29/1555 train_time:1008ms step_avg:34.75ms step:30/1555 train_time:1046ms step_avg:34.87ms step:31/1555 train_time:1077ms step_avg:34.76ms step:32/1555 train_time:1115ms step_avg:34.85ms step:33/1555 train_time:1146ms step_avg:34.74ms step:34/1555 train_time:1184ms step_avg:34.84ms step:35/1555 train_time:1216ms step_avg:34.74ms step:36/1555 train_time:1254ms step_avg:34.83ms step:37/1555 train_time:1285ms step_avg:34.72ms step:38/1555 train_time:1322ms step_avg:34.79ms step:39/1555 train_time:1353ms step_avg:34.70ms step:40/1555 train_time:1391ms step_avg:34.78ms step:41/1555 train_time:1422ms step_avg:34.68ms step:42/1555 train_time:1460ms step_avg:34.76ms step:43/1555 train_time:1491ms step_avg:34.67ms step:44/1555 train_time:1529ms step_avg:34.75ms step:45/1555 train_time:1559ms step_avg:34.65ms step:46/1555 train_time:1597ms step_avg:34.72ms step:47/1555 train_time:1628ms step_avg:34.64ms step:48/1555 train_time:1666ms step_avg:34.70ms step:49/1555 train_time:1697ms step_avg:34.62ms step:50/1555 train_time:1734ms step_avg:34.69ms step:51/1555 train_time:1765ms step_avg:34.61ms step:52/1555 train_time:1803ms step_avg:34.68ms step:53/1555 train_time:1834ms step_avg:34.61ms step:54/1555 train_time:1871ms step_avg:34.66ms step:55/1555 train_time:1903ms step_avg:34.59ms step:56/1555 train_time:1940ms step_avg:34.65ms step:57/1555 train_time:1971ms step_avg:34.58ms step:58/1555 train_time:2009ms step_avg:34.63ms step:59/1555 train_time:2040ms step_avg:34.57ms step:60/1555 train_time:2077ms step_avg:34.62ms step:61/1555 train_time:2108ms step_avg:34.56ms step:62/1555 train_time:2146ms step_avg:34.61ms step:63/1555 train_time:2177ms step_avg:34.56ms step:64/1555 train_time:2215ms step_avg:34.61ms step:65/1555 train_time:2246ms step_avg:34.55ms step:66/1555 train_time:2284ms step_avg:34.60ms step:67/1555 train_time:2315ms step_avg:34.55ms step:68/1555 train_time:2353ms step_avg:34.60ms step:69/1555 train_time:2383ms step_avg:34.54ms step:70/1555 train_time:2421ms step_avg:34.59ms step:71/1555 train_time:2452ms step_avg:34.54ms step:72/1555 train_time:2490ms step_avg:34.58ms step:73/1555 train_time:2521ms step_avg:34.54ms step:74/1555 train_time:2560ms step_avg:34.59ms step:75/1555 train_time:2591ms step_avg:34.54ms step:76/1555 train_time:2629ms step_avg:34.59ms step:77/1555 train_time:2660ms step_avg:34.54ms step:78/1555 train_time:2698ms step_avg:34.59ms step:79/1555 train_time:2729ms step_avg:34.54ms step:80/1555 train_time:2767ms step_avg:34.58ms step:81/1555 train_time:2798ms step_avg:34.54ms step:82/1555 train_time:2836ms step_avg:34.58ms step:83/1555 train_time:2867ms step_avg:34.54ms step:84/1555 train_time:2905ms step_avg:34.58ms step:85/1555 train_time:2935ms step_avg:34.53ms step:86/1555 train_time:2973ms step_avg:34.57ms step:87/1555 train_time:3004ms step_avg:34.52ms step:88/1555 train_time:3041ms step_avg:34.56ms step:89/1555 train_time:3072ms step_avg:34.51ms step:90/1555 train_time:3109ms step_avg:34.55ms step:91/1555 train_time:3141ms step_avg:34.51ms step:92/1555 train_time:3178ms step_avg:34.54ms step:93/1555 train_time:3209ms step_avg:34.51ms step:94/1555 train_time:3247ms step_avg:34.54ms step:95/1555 train_time:3278ms step_avg:34.51ms step:96/1555 train_time:3316ms step_avg:34.54ms step:97/1555 train_time:3347ms step_avg:34.50ms step:98/1555 train_time:3385ms step_avg:34.54ms step:99/1555 train_time:3416ms step_avg:34.50ms step:100/1555 train_time:3453ms step_avg:34.53ms step:101/1555 train_time:3484ms step_avg:34.50ms step:102/1555 train_time:3522ms step_avg:34.53ms step:103/1555 train_time:3553ms step_avg:34.49ms step:104/1555 train_time:3590ms step_avg:34.52ms step:105/1555 train_time:3621ms step_avg:34.49ms step:106/1555 train_time:3659ms step_avg:34.52ms step:107/1555 train_time:3690ms step_avg:34.49ms step:108/1555 train_time:3727ms step_avg:34.51ms step:109/1555 train_time:3759ms step_avg:34.48ms step:110/1555 train_time:3797ms step_avg:34.51ms step:111/1555 train_time:3827ms step_avg:34.48ms step:112/1555 train_time:3865ms step_avg:34.51ms step:113/1555 train_time:3896ms step_avg:34.48ms step:114/1555 train_time:3935ms step_avg:34.51ms step:115/1555 train_time:3966ms step_avg:34.48ms step:116/1555 train_time:4004ms step_avg:34.51ms step:117/1555 train_time:4035ms step_avg:34.48ms step:118/1555 train_time:4072ms step_avg:34.51ms step:119/1555 train_time:4103ms step_avg:34.48ms step:120/1555 train_time:4141ms step_avg:34.51ms step:121/1555 train_time:4171ms step_avg:34.47ms step:122/1555 train_time:4209ms step_avg:34.50ms step:123/1555 train_time:4240ms step_avg:34.47ms step:124/1555 train_time:4277ms step_avg:34.49ms step:125/1555 train_time:4308ms step_avg:34.46ms step:126/1555 train_time:4346ms step_avg:34.49ms step:127/1555 train_time:4377ms step_avg:34.46ms step:128/1555 train_time:4414ms step_avg:34.49ms step:129/1555 train_time:4445ms step_avg:34.46ms step:130/1555 train_time:4483ms step_avg:34.49ms step:131/1555 train_time:4514ms step_avg:34.45ms step:132/1555 train_time:4551ms step_avg:34.48ms step:133/1555 train_time:4582ms step_avg:34.45ms step:134/1555 train_time:4620ms step_avg:34.48ms step:135/1555 train_time:4651ms step_avg:34.45ms step:136/1555 train_time:4688ms step_avg:34.47ms step:137/1555 train_time:4720ms step_avg:34.45ms step:138/1555 train_time:4758ms step_avg:34.48ms step:139/1555 train_time:4789ms step_avg:34.45ms step:140/1555 train_time:4827ms step_avg:34.48ms step:141/1555 train_time:4858ms step_avg:34.45ms step:142/1555 train_time:4895ms step_avg:34.48ms step:143/1555 train_time:4927ms step_avg:34.45ms step:144/1555 train_time:4965ms step_avg:34.48ms step:145/1555 train_time:4995ms step_avg:34.45ms step:146/1555 train_time:5033ms step_avg:34.47ms step:147/1555 train_time:5064ms step_avg:34.45ms step:148/1555 train_time:5103ms step_avg:34.48ms step:149/1555 train_time:5133ms step_avg:34.45ms step:150/1555 train_time:5171ms step_avg:34.47ms step:151/1555 train_time:5202ms step_avg:34.45ms step:152/1555 train_time:5240ms step_avg:34.47ms step:153/1555 train_time:5271ms step_avg:34.45ms step:154/1555 train_time:5308ms step_avg:34.47ms step:155/1555 train_time:5339ms step_avg:34.45ms step:156/1555 train_time:5377ms step_avg:34.47ms step:157/1555 train_time:5408ms step_avg:34.44ms step:158/1555 train_time:5447ms step_avg:34.47ms step:159/1555 train_time:5476ms step_avg:34.44ms step:160/1555 train_time:5513ms step_avg:34.46ms step:161/1555 train_time:5544ms step_avg:34.43ms step:162/1555 train_time:5582ms step_avg:34.45ms step:163/1555 train_time:5613ms step_avg:34.43ms step:164/1555 train_time:5651ms step_avg:34.46ms step:165/1555 train_time:5682ms step_avg:34.43ms step:166/1555 train_time:5719ms step_avg:34.45ms step:167/1555 train_time:5750ms step_avg:34.43ms step:168/1555 train_time:5788ms step_avg:34.45ms step:169/1555 train_time:5818ms step_avg:34.43ms step:170/1555 train_time:5856ms step_avg:34.44ms step:171/1555 train_time:5886ms step_avg:34.42ms step:172/1555 train_time:5924ms step_avg:34.44ms step:173/1555 train_time:5955ms step_avg:34.42ms step:174/1555 train_time:5992ms step_avg:34.44ms step:175/1555 train_time:6023ms step_avg:34.42ms step:176/1555 train_time:6061ms step_avg:34.44ms step:177/1555 train_time:6092ms step_avg:34.42ms step:178/1555 train_time:6130ms step_avg:34.44ms step:179/1555 train_time:6161ms step_avg:34.42ms step:180/1555 train_time:6198ms step_avg:34.43ms step:181/1555 train_time:6229ms step_avg:34.42ms step:182/1555 train_time:6266ms step_avg:34.43ms step:183/1555 train_time:6297ms step_avg:34.41ms step:184/1555 train_time:6335ms step_avg:34.43ms step:185/1555 train_time:6366ms step_avg:34.41ms step:186/1555 train_time:6404ms step_avg:34.43ms step:187/1555 train_time:6435ms step_avg:34.41ms step:188/1555 train_time:6473ms step_avg:34.43ms step:189/1555 train_time:6503ms step_avg:34.41ms step:190/1555 train_time:6541ms step_avg:34.43ms step:191/1555 train_time:6572ms step_avg:34.41ms step:192/1555 train_time:6610ms step_avg:34.43ms step:193/1555 train_time:6641ms step_avg:34.41ms step:194/1555 train_time:6679ms step_avg:34.43ms step:195/1555 train_time:6709ms step_avg:34.41ms step:196/1555 train_time:6747ms step_avg:34.42ms step:197/1555 train_time:6777ms step_avg:34.40ms step:198/1555 train_time:6815ms step_avg:34.42ms step:199/1555 train_time:6846ms step_avg:34.40ms step:200/1555 train_time:6883ms step_avg:34.42ms step:201/1555 train_time:6914ms step_avg:34.40ms step:202/1555 train_time:6952ms step_avg:34.41ms step:203/1555 train_time:6982ms step_avg:34.40ms step:204/1555 train_time:7020ms step_avg:34.41ms step:205/1555 train_time:7051ms step_avg:34.39ms step:206/1555 train_time:7089ms step_avg:34.41ms step:207/1555 train_time:7120ms step_avg:34.39ms step:208/1555 train_time:7157ms step_avg:34.41ms step:209/1555 train_time:7188ms step_avg:34.39ms step:210/1555 train_time:7226ms step_avg:34.41ms step:211/1555 train_time:7257ms step_avg:34.39ms step:212/1555 train_time:7295ms step_avg:34.41ms step:213/1555 train_time:7326ms step_avg:34.39ms step:214/1555 train_time:7363ms step_avg:34.41ms step:215/1555 train_time:7394ms step_avg:34.39ms step:216/1555 train_time:7432ms step_avg:34.41ms step:217/1555 train_time:7463ms step_avg:34.39ms step:218/1555 train_time:7500ms step_avg:34.40ms step:219/1555 train_time:7531ms step_avg:34.39ms step:220/1555 train_time:7569ms step_avg:34.41ms step:221/1555 train_time:7600ms step_avg:34.39ms step:222/1555 train_time:7637ms step_avg:34.40ms step:223/1555 train_time:7668ms step_avg:34.39ms step:224/1555 train_time:7706ms step_avg:34.40ms step:225/1555 train_time:7737ms step_avg:34.39ms step:226/1555 train_time:7775ms step_avg:34.40ms step:227/1555 train_time:7806ms step_avg:34.39ms step:228/1555 train_time:7843ms step_avg:34.40ms step:229/1555 train_time:7874ms step_avg:34.38ms step:230/1555 train_time:7912ms step_avg:34.40ms step:231/1555 train_time:7943ms step_avg:34.38ms step:232/1555 train_time:7980ms step_avg:34.40ms step:233/1555 train_time:8011ms step_avg:34.38ms step:234/1555 train_time:8049ms step_avg:34.40ms step:235/1555 train_time:8080ms step_avg:34.38ms step:236/1555 train_time:8117ms step_avg:34.40ms step:237/1555 train_time:8148ms step_avg:34.38ms step:238/1555 train_time:8186ms step_avg:34.40ms step:239/1555 train_time:8218ms step_avg:34.38ms step:240/1555 train_time:8255ms step_avg:34.40ms step:241/1555 train_time:8286ms step_avg:34.38ms step:242/1555 train_time:8324ms step_avg:34.40ms step:243/1555 train_time:8355ms step_avg:34.38ms step:244/1555 train_time:8393ms step_avg:34.40ms step:245/1555 train_time:8423ms step_avg:34.38ms step:246/1555 train_time:8461ms step_avg:34.39ms step:247/1555 train_time:8492ms step_avg:34.38ms step:248/1555 train_time:8529ms step_avg:34.39ms step:249/1555 train_time:8560ms step_avg:34.38ms step:250/1555 train_time:8597ms step_avg:34.39ms step:250/1555 val_loss:4.5653 train_time:8647ms step_avg:34.59ms step:251/1555 train_time:8668ms step_avg:34.53ms step:252/1555 train_time:8695ms step_avg:34.50ms step:253/1555 train_time:8714ms step_avg:34.44ms step:254/1555 train_time:8737ms step_avg:34.40ms step:255/1555 train_time:8769ms step_avg:34.39ms step:256/1555 train_time:8808ms step_avg:34.41ms step:257/1555 train_time:8840ms step_avg:34.40ms step:258/1555 train_time:8878ms step_avg:34.41ms step:259/1555 train_time:8909ms step_avg:34.40ms step:260/1555 train_time:8947ms step_avg:34.41ms step:261/1555 train_time:8977ms step_avg:34.40ms step:262/1555 train_time:9015ms step_avg:34.41ms step:263/1555 train_time:9046ms step_avg:34.39ms step:264/1555 train_time:9083ms step_avg:34.41ms step:265/1555 train_time:9114ms step_avg:34.39ms step:266/1555 train_time:9152ms step_avg:34.41ms step:267/1555 train_time:9182ms step_avg:34.39ms step:268/1555 train_time:9220ms step_avg:34.40ms step:269/1555 train_time:9250ms step_avg:34.39ms step:270/1555 train_time:9288ms step_avg:34.40ms step:271/1555 train_time:9319ms step_avg:34.39ms step:272/1555 train_time:9356ms step_avg:34.40ms step:273/1555 train_time:9387ms step_avg:34.38ms step:274/1555 train_time:9425ms step_avg:34.40ms step:275/1555 train_time:9455ms step_avg:34.38ms step:276/1555 train_time:9492ms step_avg:34.39ms step:277/1555 train_time:9523ms step_avg:34.38ms step:278/1555 train_time:9560ms step_avg:34.39ms step:279/1555 train_time:9591ms step_avg:34.38ms step:280/1555 train_time:9630ms step_avg:34.39ms step:281/1555 train_time:9660ms step_avg:34.38ms step:282/1555 train_time:9697ms step_avg:34.39ms step:283/1555 train_time:9729ms step_avg:34.38ms step:284/1555 train_time:9766ms step_avg:34.39ms step:285/1555 train_time:9797ms step_avg:34.38ms step:286/1555 train_time:9835ms step_avg:34.39ms step:287/1555 train_time:9866ms step_avg:34.38ms step:288/1555 train_time:9904ms step_avg:34.39ms step:289/1555 train_time:9935ms step_avg:34.38ms step:290/1555 train_time:9972ms step_avg:34.39ms step:291/1555 train_time:10003ms step_avg:34.37ms step:292/1555 train_time:10041ms step_avg:34.39ms step:293/1555 train_time:10072ms step_avg:34.38ms step:294/1555 train_time:10111ms step_avg:34.39ms step:295/1555 train_time:10140ms step_avg:34.37ms step:296/1555 train_time:10178ms step_avg:34.38ms step:297/1555 train_time:10208ms step_avg:34.37ms step:298/1555 train_time:10246ms step_avg:34.38ms step:299/1555 train_time:10277ms step_avg:34.37ms step:300/1555 train_time:10314ms step_avg:34.38ms step:301/1555 train_time:10345ms step_avg:34.37ms step:302/1555 train_time:10382ms step_avg:34.38ms step:303/1555 train_time:10413ms step_avg:34.37ms step:304/1555 train_time:10451ms step_avg:34.38ms step:305/1555 train_time:10481ms step_avg:34.37ms step:306/1555 train_time:10519ms step_avg:34.38ms step:307/1555 train_time:10550ms step_avg:34.37ms step:308/1555 train_time:10588ms step_avg:34.38ms step:309/1555 train_time:10619ms step_avg:34.37ms step:310/1555 train_time:10657ms step_avg:34.38ms step:311/1555 train_time:10688ms step_avg:34.37ms step:312/1555 train_time:10726ms step_avg:34.38ms step:313/1555 train_time:10757ms step_avg:34.37ms step:314/1555 train_time:10794ms step_avg:34.38ms step:315/1555 train_time:10825ms step_avg:34.37ms step:316/1555 train_time:10863ms step_avg:34.38ms step:317/1555 train_time:10894ms step_avg:34.37ms step:318/1555 train_time:10932ms step_avg:34.38ms step:319/1555 train_time:10963ms step_avg:34.37ms step:320/1555 train_time:11001ms step_avg:34.38ms step:321/1555 train_time:11031ms step_avg:34.37ms step:322/1555 train_time:11070ms step_avg:34.38ms step:323/1555 train_time:11100ms step_avg:34.37ms step:324/1555 train_time:11138ms step_avg:34.38ms step:325/1555 train_time:11169ms step_avg:34.37ms step:326/1555 train_time:11206ms step_avg:34.38ms step:327/1555 train_time:11237ms step_avg:34.36ms step:328/1555 train_time:11275ms step_avg:34.37ms step:329/1555 train_time:11305ms step_avg:34.36ms step:330/1555 train_time:11343ms step_avg:34.37ms step:331/1555 train_time:11374ms step_avg:34.36ms step:332/1555 train_time:11411ms step_avg:34.37ms step:333/1555 train_time:11442ms step_avg:34.36ms step:334/1555 train_time:11480ms step_avg:34.37ms step:335/1555 train_time:11511ms step_avg:34.36ms step:336/1555 train_time:11550ms step_avg:34.37ms step:337/1555 train_time:11580ms step_avg:34.36ms step:338/1555 train_time:11618ms step_avg:34.37ms step:339/1555 train_time:11648ms step_avg:34.36ms step:340/1555 train_time:11686ms step_avg:34.37ms step:341/1555 train_time:11717ms step_avg:34.36ms step:342/1555 train_time:11754ms step_avg:34.37ms step:343/1555 train_time:11785ms step_avg:34.36ms step:344/1555 train_time:11822ms step_avg:34.37ms step:345/1555 train_time:11853ms step_avg:34.36ms step:346/1555 train_time:11891ms step_avg:34.37ms step:347/1555 train_time:11922ms step_avg:34.36ms step:348/1555 train_time:11959ms step_avg:34.37ms step:349/1555 train_time:11991ms step_avg:34.36ms step:350/1555 train_time:12028ms step_avg:34.37ms step:351/1555 train_time:12059ms step_avg:34.36ms step:352/1555 train_time:12096ms step_avg:34.36ms step:353/1555 train_time:12127ms step_avg:34.35ms step:354/1555 train_time:12165ms step_avg:34.36ms step:355/1555 train_time:12196ms step_avg:34.35ms step:356/1555 train_time:12234ms step_avg:34.36ms step:357/1555 train_time:12264ms step_avg:34.35ms step:358/1555 train_time:12302ms step_avg:34.36ms step:359/1555 train_time:12333ms step_avg:34.35ms step:360/1555 train_time:12370ms step_avg:34.36ms step:361/1555 train_time:12401ms step_avg:34.35ms step:362/1555 train_time:12439ms step_avg:34.36ms step:363/1555 train_time:12469ms step_avg:34.35ms step:364/1555 train_time:12507ms step_avg:34.36ms step:365/1555 train_time:12538ms step_avg:34.35ms step:366/1555 train_time:12575ms step_avg:34.36ms step:367/1555 train_time:12606ms step_avg:34.35ms step:368/1555 train_time:12644ms step_avg:34.36ms step:369/1555 train_time:12675ms step_avg:34.35ms step:370/1555 train_time:12712ms step_avg:34.36ms step:371/1555 train_time:12743ms step_avg:34.35ms step:372/1555 train_time:12781ms step_avg:34.36ms step:373/1555 train_time:12812ms step_avg:34.35ms step:374/1555 train_time:12850ms step_avg:34.36ms step:375/1555 train_time:12880ms step_avg:34.35ms step:376/1555 train_time:12918ms step_avg:34.36ms step:377/1555 train_time:12949ms step_avg:34.35ms step:378/1555 train_time:12987ms step_avg:34.36ms step:379/1555 train_time:13018ms step_avg:34.35ms step:380/1555 train_time:13056ms step_avg:34.36ms step:381/1555 train_time:13086ms step_avg:34.35ms step:382/1555 train_time:13123ms step_avg:34.35ms step:383/1555 train_time:13154ms step_avg:34.35ms step:384/1555 train_time:13193ms step_avg:34.36ms step:385/1555 train_time:13223ms step_avg:34.35ms step:386/1555 train_time:13261ms step_avg:34.35ms step:387/1555 train_time:13292ms step_avg:34.35ms step:388/1555 train_time:13329ms step_avg:34.35ms step:389/1555 train_time:13360ms step_avg:34.34ms step:390/1555 train_time:13397ms step_avg:34.35ms step:391/1555 train_time:13428ms step_avg:34.34ms step:392/1555 train_time:13466ms step_avg:34.35ms step:393/1555 train_time:13497ms step_avg:34.34ms step:394/1555 train_time:13534ms step_avg:34.35ms step:395/1555 train_time:13565ms step_avg:34.34ms step:396/1555 train_time:13603ms step_avg:34.35ms step:397/1555 train_time:13634ms step_avg:34.34ms step:398/1555 train_time:13671ms step_avg:34.35ms step:399/1555 train_time:13701ms step_avg:34.34ms step:400/1555 train_time:13739ms step_avg:34.35ms step:401/1555 train_time:13770ms step_avg:34.34ms step:402/1555 train_time:13807ms step_avg:34.35ms step:403/1555 train_time:13838ms step_avg:34.34ms step:404/1555 train_time:13875ms step_avg:34.34ms step:405/1555 train_time:13906ms step_avg:34.34ms step:406/1555 train_time:13944ms step_avg:34.34ms step:407/1555 train_time:13975ms step_avg:34.34ms step:408/1555 train_time:14013ms step_avg:34.34ms step:409/1555 train_time:14044ms step_avg:34.34ms step:410/1555 train_time:14081ms step_avg:34.34ms step:411/1555 train_time:14112ms step_avg:34.34ms step:412/1555 train_time:14150ms step_avg:34.35ms step:413/1555 train_time:14181ms step_avg:34.34ms step:414/1555 train_time:14218ms step_avg:34.34ms step:415/1555 train_time:14250ms step_avg:34.34ms step:416/1555 train_time:14288ms step_avg:34.35ms step:417/1555 train_time:14319ms step_avg:34.34ms step:418/1555 train_time:14357ms step_avg:34.35ms step:419/1555 train_time:14388ms step_avg:34.34ms step:420/1555 train_time:14426ms step_avg:34.35ms step:421/1555 train_time:14456ms step_avg:34.34ms step:422/1555 train_time:14494ms step_avg:34.35ms step:423/1555 train_time:14525ms step_avg:34.34ms step:424/1555 train_time:14563ms step_avg:34.35ms step:425/1555 train_time:14594ms step_avg:34.34ms step:426/1555 train_time:14632ms step_avg:34.35ms step:427/1555 train_time:14663ms step_avg:34.34ms step:428/1555 train_time:14700ms step_avg:34.35ms step:429/1555 train_time:14731ms step_avg:34.34ms step:430/1555 train_time:14768ms step_avg:34.35ms step:431/1555 train_time:14799ms step_avg:34.34ms step:432/1555 train_time:14837ms step_avg:34.34ms step:433/1555 train_time:14867ms step_avg:34.34ms step:434/1555 train_time:14905ms step_avg:34.34ms step:435/1555 train_time:14936ms step_avg:34.34ms step:436/1555 train_time:14973ms step_avg:34.34ms step:437/1555 train_time:15004ms step_avg:34.33ms step:438/1555 train_time:15042ms step_avg:34.34ms step:439/1555 train_time:15073ms step_avg:34.33ms step:440/1555 train_time:15111ms step_avg:34.34ms step:441/1555 train_time:15141ms step_avg:34.33ms step:442/1555 train_time:15179ms step_avg:34.34ms step:443/1555 train_time:15210ms step_avg:34.33ms step:444/1555 train_time:15247ms step_avg:34.34ms step:445/1555 train_time:15278ms step_avg:34.33ms step:446/1555 train_time:15315ms step_avg:34.34ms step:447/1555 train_time:15345ms step_avg:34.33ms step:448/1555 train_time:15382ms step_avg:34.34ms step:449/1555 train_time:15414ms step_avg:34.33ms step:450/1555 train_time:15451ms step_avg:34.34ms step:451/1555 train_time:15482ms step_avg:34.33ms step:452/1555 train_time:15519ms step_avg:34.33ms step:453/1555 train_time:15551ms step_avg:34.33ms step:454/1555 train_time:15589ms step_avg:34.34ms step:455/1555 train_time:15619ms step_avg:34.33ms step:456/1555 train_time:15657ms step_avg:34.33ms step:457/1555 train_time:15687ms step_avg:34.33ms step:458/1555 train_time:15725ms step_avg:34.33ms step:459/1555 train_time:15755ms step_avg:34.32ms step:460/1555 train_time:15793ms step_avg:34.33ms step:461/1555 train_time:15824ms step_avg:34.33ms step:462/1555 train_time:15862ms step_avg:34.33ms step:463/1555 train_time:15892ms step_avg:34.32ms step:464/1555 train_time:15930ms step_avg:34.33ms step:465/1555 train_time:15961ms step_avg:34.33ms step:466/1555 train_time:15999ms step_avg:34.33ms step:467/1555 train_time:16030ms step_avg:34.33ms step:468/1555 train_time:16068ms step_avg:34.33ms step:469/1555 train_time:16099ms step_avg:34.33ms step:470/1555 train_time:16137ms step_avg:34.33ms step:471/1555 train_time:16168ms step_avg:34.33ms step:472/1555 train_time:16206ms step_avg:34.33ms step:473/1555 train_time:16237ms step_avg:34.33ms step:474/1555 train_time:16274ms step_avg:34.33ms step:475/1555 train_time:16305ms step_avg:34.33ms step:476/1555 train_time:16343ms step_avg:34.33ms step:477/1555 train_time:16374ms step_avg:34.33ms step:478/1555 train_time:16411ms step_avg:34.33ms step:479/1555 train_time:16442ms step_avg:34.33ms step:480/1555 train_time:16480ms step_avg:34.33ms step:481/1555 train_time:16511ms step_avg:34.33ms step:482/1555 train_time:16549ms step_avg:34.34ms step:483/1555 train_time:16580ms step_avg:34.33ms step:484/1555 train_time:16618ms step_avg:34.33ms step:485/1555 train_time:16648ms step_avg:34.33ms step:486/1555 train_time:16686ms step_avg:34.33ms step:487/1555 train_time:16716ms step_avg:34.33ms step:488/1555 train_time:16754ms step_avg:34.33ms step:489/1555 train_time:16785ms step_avg:34.32ms step:490/1555 train_time:16822ms step_avg:34.33ms step:491/1555 train_time:16853ms step_avg:34.32ms step:492/1555 train_time:16890ms step_avg:34.33ms step:493/1555 train_time:16922ms step_avg:34.32ms step:494/1555 train_time:16959ms step_avg:34.33ms step:495/1555 train_time:16990ms step_avg:34.32ms step:496/1555 train_time:17028ms step_avg:34.33ms step:497/1555 train_time:17058ms step_avg:34.32ms step:498/1555 train_time:17096ms step_avg:34.33ms step:499/1555 train_time:17127ms step_avg:34.32ms step:500/1555 train_time:17165ms step_avg:34.33ms step:500/1555 val_loss:4.2334 train_time:17214ms step_avg:34.43ms step:501/1555 train_time:17234ms step_avg:34.40ms step:502/1555 train_time:17255ms step_avg:34.37ms step:503/1555 train_time:17274ms step_avg:34.34ms step:504/1555 train_time:17303ms step_avg:34.33ms step:505/1555 train_time:17337ms step_avg:34.33ms step:506/1555 train_time:17379ms step_avg:34.35ms step:507/1555 train_time:17434ms step_avg:34.39ms step:508/1555 train_time:17498ms step_avg:34.45ms step:509/1555 train_time:17556ms step_avg:34.49ms step:510/1555 train_time:17619ms step_avg:34.55ms step:511/1555 train_time:17676ms step_avg:34.59ms step:512/1555 train_time:17739ms step_avg:34.65ms step:513/1555 train_time:17796ms step_avg:34.69ms step:514/1555 train_time:17858ms step_avg:34.74ms step:515/1555 train_time:17915ms step_avg:34.79ms step:516/1555 train_time:17978ms step_avg:34.84ms step:517/1555 train_time:18036ms step_avg:34.89ms step:518/1555 train_time:18100ms step_avg:34.94ms step:519/1555 train_time:18157ms step_avg:34.98ms step:520/1555 train_time:18224ms step_avg:35.05ms step:521/1555 train_time:18284ms step_avg:35.09ms step:522/1555 train_time:18351ms step_avg:35.15ms step:523/1555 train_time:18409ms step_avg:35.20ms step:524/1555 train_time:18475ms step_avg:35.26ms step:525/1555 train_time:18533ms step_avg:35.30ms step:526/1555 train_time:18598ms step_avg:35.36ms step:527/1555 train_time:18655ms step_avg:35.40ms step:528/1555 train_time:18720ms step_avg:35.45ms step:529/1555 train_time:18777ms step_avg:35.49ms step:530/1555 train_time:18840ms step_avg:35.55ms step:531/1555 train_time:18897ms step_avg:35.59ms step:532/1555 train_time:18960ms step_avg:35.64ms step:533/1555 train_time:19017ms step_avg:35.68ms step:534/1555 train_time:19081ms step_avg:35.73ms step:535/1555 train_time:19138ms step_avg:35.77ms step:536/1555 train_time:19202ms step_avg:35.83ms step:537/1555 train_time:19261ms step_avg:35.87ms step:538/1555 train_time:19327ms step_avg:35.92ms step:539/1555 train_time:19386ms step_avg:35.97ms step:540/1555 train_time:19451ms step_avg:36.02ms step:541/1555 train_time:19510ms step_avg:36.06ms step:542/1555 train_time:19575ms step_avg:36.12ms step:543/1555 train_time:19632ms step_avg:36.16ms step:544/1555 train_time:19696ms step_avg:36.21ms step:545/1555 train_time:19754ms step_avg:36.25ms step:546/1555 train_time:19818ms step_avg:36.30ms step:547/1555 train_time:19876ms step_avg:36.34ms step:548/1555 train_time:19939ms step_avg:36.39ms step:549/1555 train_time:19996ms step_avg:36.42ms step:550/1555 train_time:20059ms step_avg:36.47ms step:551/1555 train_time:20117ms step_avg:36.51ms step:552/1555 train_time:20182ms step_avg:36.56ms step:553/1555 train_time:20240ms step_avg:36.60ms step:554/1555 train_time:20304ms step_avg:36.65ms step:555/1555 train_time:20363ms step_avg:36.69ms step:556/1555 train_time:20428ms step_avg:36.74ms step:557/1555 train_time:20486ms step_avg:36.78ms step:558/1555 train_time:20552ms step_avg:36.83ms step:559/1555 train_time:20610ms step_avg:36.87ms step:560/1555 train_time:20675ms step_avg:36.92ms step:561/1555 train_time:20732ms step_avg:36.96ms step:562/1555 train_time:20797ms step_avg:37.01ms step:563/1555 train_time:20855ms step_avg:37.04ms step:564/1555 train_time:20920ms step_avg:37.09ms step:565/1555 train_time:20977ms step_avg:37.13ms step:566/1555 train_time:21040ms step_avg:37.17ms step:567/1555 train_time:21097ms step_avg:37.21ms step:568/1555 train_time:21161ms step_avg:37.26ms step:569/1555 train_time:21218ms step_avg:37.29ms step:570/1555 train_time:21283ms step_avg:37.34ms step:571/1555 train_time:21341ms step_avg:37.37ms step:572/1555 train_time:21407ms step_avg:37.42ms step:573/1555 train_time:21465ms step_avg:37.46ms step:574/1555 train_time:21530ms step_avg:37.51ms step:575/1555 train_time:21588ms step_avg:37.54ms step:576/1555 train_time:21653ms step_avg:37.59ms step:577/1555 train_time:21711ms step_avg:37.63ms step:578/1555 train_time:21775ms step_avg:37.67ms step:579/1555 train_time:21833ms step_avg:37.71ms step:580/1555 train_time:21898ms step_avg:37.76ms step:581/1555 train_time:21956ms step_avg:37.79ms step:582/1555 train_time:22020ms step_avg:37.84ms step:583/1555 train_time:22078ms step_avg:37.87ms step:584/1555 train_time:22141ms step_avg:37.91ms step:585/1555 train_time:22198ms step_avg:37.95ms step:586/1555 train_time:22262ms step_avg:37.99ms step:587/1555 train_time:22320ms step_avg:38.02ms step:588/1555 train_time:22384ms step_avg:38.07ms step:589/1555 train_time:22442ms step_avg:38.10ms step:590/1555 train_time:22507ms step_avg:38.15ms step:591/1555 train_time:22566ms step_avg:38.18ms step:592/1555 train_time:22630ms step_avg:38.23ms step:593/1555 train_time:22689ms step_avg:38.26ms step:594/1555 train_time:22754ms step_avg:38.31ms step:595/1555 train_time:22812ms step_avg:38.34ms step:596/1555 train_time:22877ms step_avg:38.38ms step:597/1555 train_time:22934ms step_avg:38.42ms step:598/1555 train_time:22999ms step_avg:38.46ms step:599/1555 train_time:23056ms step_avg:38.49ms step:600/1555 train_time:23120ms step_avg:38.53ms step:601/1555 train_time:23177ms step_avg:38.56ms step:602/1555 train_time:23241ms step_avg:38.61ms step:603/1555 train_time:23298ms step_avg:38.64ms step:604/1555 train_time:23362ms step_avg:38.68ms step:605/1555 train_time:23420ms step_avg:38.71ms step:606/1555 train_time:23485ms step_avg:38.75ms step:607/1555 train_time:23543ms step_avg:38.79ms step:608/1555 train_time:23608ms step_avg:38.83ms step:609/1555 train_time:23666ms step_avg:38.86ms step:610/1555 train_time:23730ms step_avg:38.90ms step:611/1555 train_time:23789ms step_avg:38.93ms step:612/1555 train_time:23852ms step_avg:38.97ms step:613/1555 train_time:23911ms step_avg:39.01ms step:614/1555 train_time:23976ms step_avg:39.05ms step:615/1555 train_time:24034ms step_avg:39.08ms step:616/1555 train_time:24098ms step_avg:39.12ms step:617/1555 train_time:24156ms step_avg:39.15ms step:618/1555 train_time:24221ms step_avg:39.19ms step:619/1555 train_time:24278ms step_avg:39.22ms step:620/1555 train_time:24342ms step_avg:39.26ms step:621/1555 train_time:24399ms step_avg:39.29ms step:622/1555 train_time:24463ms step_avg:39.33ms step:623/1555 train_time:24520ms step_avg:39.36ms step:624/1555 train_time:24586ms step_avg:39.40ms step:625/1555 train_time:24644ms step_avg:39.43ms step:626/1555 train_time:24709ms step_avg:39.47ms step:627/1555 train_time:24767ms step_avg:39.50ms step:628/1555 train_time:24832ms step_avg:39.54ms step:629/1555 train_time:24889ms step_avg:39.57ms step:630/1555 train_time:24954ms step_avg:39.61ms step:631/1555 train_time:25012ms step_avg:39.64ms step:632/1555 train_time:25077ms step_avg:39.68ms step:633/1555 train_time:25134ms step_avg:39.71ms step:634/1555 train_time:25198ms step_avg:39.75ms step:635/1555 train_time:25256ms step_avg:39.77ms step:636/1555 train_time:25321ms step_avg:39.81ms step:637/1555 train_time:25378ms step_avg:39.84ms step:638/1555 train_time:25441ms step_avg:39.88ms step:639/1555 train_time:25499ms step_avg:39.90ms step:640/1555 train_time:25563ms step_avg:39.94ms step:641/1555 train_time:25620ms step_avg:39.97ms step:642/1555 train_time:25685ms step_avg:40.01ms step:643/1555 train_time:25743ms step_avg:40.04ms step:644/1555 train_time:25808ms step_avg:40.07ms step:645/1555 train_time:25866ms step_avg:40.10ms step:646/1555 train_time:25931ms step_avg:40.14ms step:647/1555 train_time:25989ms step_avg:40.17ms step:648/1555 train_time:26055ms step_avg:40.21ms step:649/1555 train_time:26113ms step_avg:40.24ms step:650/1555 train_time:26178ms step_avg:40.27ms step:651/1555 train_time:26235ms step_avg:40.30ms step:652/1555 train_time:26300ms step_avg:40.34ms step:653/1555 train_time:26357ms step_avg:40.36ms step:654/1555 train_time:26421ms step_avg:40.40ms step:655/1555 train_time:26478ms step_avg:40.42ms step:656/1555 train_time:26542ms step_avg:40.46ms step:657/1555 train_time:26599ms step_avg:40.49ms step:658/1555 train_time:26664ms step_avg:40.52ms step:659/1555 train_time:26721ms step_avg:40.55ms step:660/1555 train_time:26786ms step_avg:40.58ms step:661/1555 train_time:26844ms step_avg:40.61ms step:662/1555 train_time:26908ms step_avg:40.65ms step:663/1555 train_time:26967ms step_avg:40.67ms step:664/1555 train_time:27032ms step_avg:40.71ms step:665/1555 train_time:27090ms step_avg:40.74ms step:666/1555 train_time:27155ms step_avg:40.77ms step:667/1555 train_time:27213ms step_avg:40.80ms step:668/1555 train_time:27278ms step_avg:40.83ms step:669/1555 train_time:27335ms step_avg:40.86ms step:670/1555 train_time:27399ms step_avg:40.89ms step:671/1555 train_time:27456ms step_avg:40.92ms step:672/1555 train_time:27520ms step_avg:40.95ms step:673/1555 train_time:27577ms step_avg:40.98ms step:674/1555 train_time:27641ms step_avg:41.01ms step:675/1555 train_time:27698ms step_avg:41.03ms step:676/1555 train_time:27762ms step_avg:41.07ms step:677/1555 train_time:27819ms step_avg:41.09ms step:678/1555 train_time:27885ms step_avg:41.13ms step:679/1555 train_time:27943ms step_avg:41.15ms step:680/1555 train_time:28008ms step_avg:41.19ms step:681/1555 train_time:28066ms step_avg:41.21ms step:682/1555 train_time:28131ms step_avg:41.25ms step:683/1555 train_time:28190ms step_avg:41.27ms step:684/1555 train_time:28254ms step_avg:41.31ms step:685/1555 train_time:28313ms step_avg:41.33ms step:686/1555 train_time:28376ms step_avg:41.36ms step:687/1555 train_time:28433ms step_avg:41.39ms step:688/1555 train_time:28498ms step_avg:41.42ms step:689/1555 train_time:28556ms step_avg:41.45ms step:690/1555 train_time:28620ms step_avg:41.48ms step:691/1555 train_time:28677ms step_avg:41.50ms step:692/1555 train_time:28740ms step_avg:41.53ms step:693/1555 train_time:28797ms step_avg:41.55ms step:694/1555 train_time:28862ms step_avg:41.59ms step:695/1555 train_time:28920ms step_avg:41.61ms step:696/1555 train_time:28985ms step_avg:41.64ms step:697/1555 train_time:29042ms step_avg:41.67ms step:698/1555 train_time:29108ms step_avg:41.70ms step:699/1555 train_time:29166ms step_avg:41.73ms step:700/1555 train_time:29231ms step_avg:41.76ms step:701/1555 train_time:29289ms step_avg:41.78ms step:702/1555 train_time:29354ms step_avg:41.81ms step:703/1555 train_time:29413ms step_avg:41.84ms step:704/1555 train_time:29476ms step_avg:41.87ms step:705/1555 train_time:29534ms step_avg:41.89ms step:706/1555 train_time:29598ms step_avg:41.92ms step:707/1555 train_time:29656ms step_avg:41.95ms step:708/1555 train_time:29719ms step_avg:41.98ms step:709/1555 train_time:29777ms step_avg:42.00ms step:710/1555 train_time:29841ms step_avg:42.03ms step:711/1555 train_time:29898ms step_avg:42.05ms step:712/1555 train_time:29963ms step_avg:42.08ms step:713/1555 train_time:30020ms step_avg:42.10ms step:714/1555 train_time:30085ms step_avg:42.14ms step:715/1555 train_time:30143ms step_avg:42.16ms step:716/1555 train_time:30208ms step_avg:42.19ms step:717/1555 train_time:30267ms step_avg:42.21ms step:718/1555 train_time:30331ms step_avg:42.24ms step:719/1555 train_time:30389ms step_avg:42.27ms step:720/1555 train_time:30454ms step_avg:42.30ms step:721/1555 train_time:30512ms step_avg:42.32ms step:722/1555 train_time:30576ms step_avg:42.35ms step:723/1555 train_time:30634ms step_avg:42.37ms step:724/1555 train_time:30698ms step_avg:42.40ms step:725/1555 train_time:30757ms step_avg:42.42ms step:726/1555 train_time:30821ms step_avg:42.45ms step:727/1555 train_time:30878ms step_avg:42.47ms step:728/1555 train_time:30941ms step_avg:42.50ms step:729/1555 train_time:30998ms step_avg:42.52ms step:730/1555 train_time:31063ms step_avg:42.55ms step:731/1555 train_time:31120ms step_avg:42.57ms step:732/1555 train_time:31185ms step_avg:42.60ms step:733/1555 train_time:31244ms step_avg:42.62ms step:734/1555 train_time:31308ms step_avg:42.65ms step:735/1555 train_time:31367ms step_avg:42.68ms step:736/1555 train_time:31431ms step_avg:42.70ms step:737/1555 train_time:31489ms step_avg:42.73ms step:738/1555 train_time:31553ms step_avg:42.76ms step:739/1555 train_time:31612ms step_avg:42.78ms step:740/1555 train_time:31676ms step_avg:42.81ms step:741/1555 train_time:31734ms step_avg:42.83ms step:742/1555 train_time:31799ms step_avg:42.86ms step:743/1555 train_time:31856ms step_avg:42.88ms step:744/1555 train_time:31921ms step_avg:42.90ms step:745/1555 train_time:31978ms step_avg:42.92ms step:746/1555 train_time:32041ms step_avg:42.95ms step:747/1555 train_time:32098ms step_avg:42.97ms step:748/1555 train_time:32162ms step_avg:43.00ms step:749/1555 train_time:32220ms step_avg:43.02ms step:750/1555 train_time:32285ms step_avg:43.05ms step:750/1555 val_loss:3.8780 train_time:32368ms step_avg:43.16ms step:751/1555 train_time:32389ms step_avg:43.13ms step:752/1555 train_time:32412ms step_avg:43.10ms step:753/1555 train_time:32466ms step_avg:43.12ms step:754/1555 train_time:32536ms step_avg:43.15ms step:755/1555 train_time:32594ms step_avg:43.17ms step:756/1555 train_time:32658ms step_avg:43.20ms step:757/1555 train_time:32715ms step_avg:43.22ms step:758/1555 train_time:32779ms step_avg:43.24ms step:759/1555 train_time:32835ms step_avg:43.26ms step:760/1555 train_time:32899ms step_avg:43.29ms step:761/1555 train_time:32956ms step_avg:43.31ms step:762/1555 train_time:33020ms step_avg:43.33ms step:763/1555 train_time:33077ms step_avg:43.35ms step:764/1555 train_time:33140ms step_avg:43.38ms step:765/1555 train_time:33198ms step_avg:43.40ms step:766/1555 train_time:33261ms step_avg:43.42ms step:767/1555 train_time:33320ms step_avg:43.44ms step:768/1555 train_time:33385ms step_avg:43.47ms step:769/1555 train_time:33446ms step_avg:43.49ms step:770/1555 train_time:33511ms step_avg:43.52ms step:771/1555 train_time:33570ms step_avg:43.54ms step:772/1555 train_time:33635ms step_avg:43.57ms step:773/1555 train_time:33692ms step_avg:43.59ms step:774/1555 train_time:33756ms step_avg:43.61ms step:775/1555 train_time:33813ms step_avg:43.63ms step:776/1555 train_time:33877ms step_avg:43.66ms step:777/1555 train_time:33934ms step_avg:43.67ms step:778/1555 train_time:33999ms step_avg:43.70ms step:779/1555 train_time:34055ms step_avg:43.72ms step:780/1555 train_time:34119ms step_avg:43.74ms step:781/1555 train_time:34175ms step_avg:43.76ms step:782/1555 train_time:34240ms step_avg:43.79ms step:783/1555 train_time:34297ms step_avg:43.80ms step:784/1555 train_time:34362ms step_avg:43.83ms step:785/1555 train_time:34421ms step_avg:43.85ms step:786/1555 train_time:34486ms step_avg:43.88ms step:787/1555 train_time:34545ms step_avg:43.90ms step:788/1555 train_time:34610ms step_avg:43.92ms step:789/1555 train_time:34668ms step_avg:43.94ms step:790/1555 train_time:34733ms step_avg:43.97ms step:791/1555 train_time:34791ms step_avg:43.98ms step:792/1555 train_time:34854ms step_avg:44.01ms step:793/1555 train_time:34911ms step_avg:44.02ms step:794/1555 train_time:34974ms step_avg:44.05ms step:795/1555 train_time:35031ms step_avg:44.06ms step:796/1555 train_time:35095ms step_avg:44.09ms step:797/1555 train_time:35152ms step_avg:44.11ms step:798/1555 train_time:35216ms step_avg:44.13ms step:799/1555 train_time:35274ms step_avg:44.15ms step:800/1555 train_time:35338ms step_avg:44.17ms step:801/1555 train_time:35395ms step_avg:44.19ms step:802/1555 train_time:35461ms step_avg:44.22ms step:803/1555 train_time:35519ms step_avg:44.23ms step:804/1555 train_time:35584ms step_avg:44.26ms step:805/1555 train_time:35643ms step_avg:44.28ms step:806/1555 train_time:35707ms step_avg:44.30ms step:807/1555 train_time:35766ms step_avg:44.32ms step:808/1555 train_time:35830ms step_avg:44.34ms step:809/1555 train_time:35888ms step_avg:44.36ms step:810/1555 train_time:35951ms step_avg:44.38ms step:811/1555 train_time:36009ms step_avg:44.40ms step:812/1555 train_time:36072ms step_avg:44.42ms step:813/1555 train_time:36130ms step_avg:44.44ms step:814/1555 train_time:36194ms step_avg:44.46ms step:815/1555 train_time:36251ms step_avg:44.48ms step:816/1555 train_time:36315ms step_avg:44.50ms step:817/1555 train_time:36372ms step_avg:44.52ms step:818/1555 train_time:36438ms step_avg:44.54ms step:819/1555 train_time:36494ms step_avg:44.56ms step:820/1555 train_time:36560ms step_avg:44.59ms step:821/1555 train_time:36618ms step_avg:44.60ms step:822/1555 train_time:36682ms step_avg:44.63ms step:823/1555 train_time:36740ms step_avg:44.64ms step:824/1555 train_time:36804ms step_avg:44.67ms step:825/1555 train_time:36863ms step_avg:44.68ms step:826/1555 train_time:36928ms step_avg:44.71ms step:827/1555 train_time:36986ms step_avg:44.72ms step:828/1555 train_time:37050ms step_avg:44.75ms step:829/1555 train_time:37107ms step_avg:44.76ms step:830/1555 train_time:37171ms step_avg:44.78ms step:831/1555 train_time:37230ms step_avg:44.80ms step:832/1555 train_time:37293ms step_avg:44.82ms step:833/1555 train_time:37350ms step_avg:44.84ms step:834/1555 train_time:37415ms step_avg:44.86ms step:835/1555 train_time:37471ms step_avg:44.88ms step:836/1555 train_time:37536ms step_avg:44.90ms step:837/1555 train_time:37594ms step_avg:44.91ms step:838/1555 train_time:37658ms step_avg:44.94ms step:839/1555 train_time:37717ms step_avg:44.95ms step:840/1555 train_time:37782ms step_avg:44.98ms step:841/1555 train_time:37840ms step_avg:44.99ms step:842/1555 train_time:37904ms step_avg:45.02ms step:843/1555 train_time:37962ms step_avg:45.03ms step:844/1555 train_time:38026ms step_avg:45.05ms step:845/1555 train_time:38084ms step_avg:45.07ms step:846/1555 train_time:38148ms step_avg:45.09ms step:847/1555 train_time:38206ms step_avg:45.11ms step:848/1555 train_time:38270ms step_avg:45.13ms step:849/1555 train_time:38329ms step_avg:45.15ms step:850/1555 train_time:38393ms step_avg:45.17ms step:851/1555 train_time:38450ms step_avg:45.18ms step:852/1555 train_time:38514ms step_avg:45.20ms step:853/1555 train_time:38571ms step_avg:45.22ms step:854/1555 train_time:38636ms step_avg:45.24ms step:855/1555 train_time:38693ms step_avg:45.26ms step:856/1555 train_time:38758ms step_avg:45.28ms step:857/1555 train_time:38816ms step_avg:45.29ms step:858/1555 train_time:38880ms step_avg:45.32ms step:859/1555 train_time:38939ms step_avg:45.33ms step:860/1555 train_time:39003ms step_avg:45.35ms step:861/1555 train_time:39060ms step_avg:45.37ms step:862/1555 train_time:39125ms step_avg:45.39ms step:863/1555 train_time:39184ms step_avg:45.40ms step:864/1555 train_time:39249ms step_avg:45.43ms step:865/1555 train_time:39307ms step_avg:45.44ms step:866/1555 train_time:39372ms step_avg:45.46ms step:867/1555 train_time:39430ms step_avg:45.48ms step:868/1555 train_time:39493ms step_avg:45.50ms step:869/1555 train_time:39551ms step_avg:45.51ms step:870/1555 train_time:39615ms step_avg:45.53ms step:871/1555 train_time:39672ms step_avg:45.55ms step:872/1555 train_time:39736ms step_avg:45.57ms step:873/1555 train_time:39794ms step_avg:45.58ms step:874/1555 train_time:39858ms step_avg:45.60ms step:875/1555 train_time:39916ms step_avg:45.62ms step:876/1555 train_time:39980ms step_avg:45.64ms step:877/1555 train_time:40038ms step_avg:45.65ms step:878/1555 train_time:40102ms step_avg:45.67ms step:879/1555 train_time:40161ms step_avg:45.69ms step:880/1555 train_time:40225ms step_avg:45.71ms step:881/1555 train_time:40284ms step_avg:45.73ms step:882/1555 train_time:40349ms step_avg:45.75ms step:883/1555 train_time:40407ms step_avg:45.76ms step:884/1555 train_time:40471ms step_avg:45.78ms step:885/1555 train_time:40530ms step_avg:45.80ms step:886/1555 train_time:40592ms step_avg:45.82ms step:887/1555 train_time:40650ms step_avg:45.83ms step:888/1555 train_time:40714ms step_avg:45.85ms step:889/1555 train_time:40771ms step_avg:45.86ms step:890/1555 train_time:40836ms step_avg:45.88ms step:891/1555 train_time:40893ms step_avg:45.90ms step:892/1555 train_time:40958ms step_avg:45.92ms step:893/1555 train_time:41015ms step_avg:45.93ms step:894/1555 train_time:41080ms step_avg:45.95ms step:895/1555 train_time:41138ms step_avg:45.96ms step:896/1555 train_time:41203ms step_avg:45.99ms step:897/1555 train_time:41261ms step_avg:46.00ms step:898/1555 train_time:41326ms step_avg:46.02ms step:899/1555 train_time:41384ms step_avg:46.03ms step:900/1555 train_time:41448ms step_avg:46.05ms step:901/1555 train_time:41506ms step_avg:46.07ms step:902/1555 train_time:41570ms step_avg:46.09ms step:903/1555 train_time:41628ms step_avg:46.10ms step:904/1555 train_time:41692ms step_avg:46.12ms step:905/1555 train_time:41750ms step_avg:46.13ms step:906/1555 train_time:41814ms step_avg:46.15ms step:907/1555 train_time:41872ms step_avg:46.17ms step:908/1555 train_time:41936ms step_avg:46.18ms step:909/1555 train_time:41992ms step_avg:46.20ms step:910/1555 train_time:42057ms step_avg:46.22ms step:911/1555 train_time:42113ms step_avg:46.23ms step:912/1555 train_time:42178ms step_avg:46.25ms step:913/1555 train_time:42236ms step_avg:46.26ms step:914/1555 train_time:42301ms step_avg:46.28ms step:915/1555 train_time:42359ms step_avg:46.29ms step:916/1555 train_time:42423ms step_avg:46.31ms step:917/1555 train_time:42481ms step_avg:46.33ms step:918/1555 train_time:42545ms step_avg:46.35ms step:919/1555 train_time:42605ms step_avg:46.36ms step:920/1555 train_time:42668ms step_avg:46.38ms step:921/1555 train_time:42727ms step_avg:46.39ms step:922/1555 train_time:42791ms step_avg:46.41ms step:923/1555 train_time:42849ms step_avg:46.42ms step:924/1555 train_time:42913ms step_avg:46.44ms step:925/1555 train_time:42971ms step_avg:46.46ms step:926/1555 train_time:43034ms step_avg:46.47ms step:927/1555 train_time:43091ms step_avg:46.48ms step:928/1555 train_time:43155ms step_avg:46.50ms step:929/1555 train_time:43213ms step_avg:46.52ms step:930/1555 train_time:43278ms step_avg:46.54ms step:931/1555 train_time:43335ms step_avg:46.55ms step:932/1555 train_time:43400ms step_avg:46.57ms step:933/1555 train_time:43458ms step_avg:46.58ms step:934/1555 train_time:43523ms step_avg:46.60ms step:935/1555 train_time:43581ms step_avg:46.61ms step:936/1555 train_time:43645ms step_avg:46.63ms step:937/1555 train_time:43704ms step_avg:46.64ms step:938/1555 train_time:43767ms step_avg:46.66ms step:939/1555 train_time:43825ms step_avg:46.67ms step:940/1555 train_time:43889ms step_avg:46.69ms step:941/1555 train_time:43948ms step_avg:46.70ms step:942/1555 train_time:44011ms step_avg:46.72ms step:943/1555 train_time:44069ms step_avg:46.73ms step:944/1555 train_time:44133ms step_avg:46.75ms step:945/1555 train_time:44190ms step_avg:46.76ms step:946/1555 train_time:44254ms step_avg:46.78ms step:947/1555 train_time:44311ms step_avg:46.79ms step:948/1555 train_time:44376ms step_avg:46.81ms step:949/1555 train_time:44434ms step_avg:46.82ms step:950/1555 train_time:44499ms step_avg:46.84ms step:951/1555 train_time:44556ms step_avg:46.85ms step:952/1555 train_time:44621ms step_avg:46.87ms step:953/1555 train_time:44679ms step_avg:46.88ms step:954/1555 train_time:44743ms step_avg:46.90ms step:955/1555 train_time:44801ms step_avg:46.91ms step:956/1555 train_time:44866ms step_avg:46.93ms step:957/1555 train_time:44924ms step_avg:46.94ms step:958/1555 train_time:44989ms step_avg:46.96ms step:959/1555 train_time:45047ms step_avg:46.97ms step:960/1555 train_time:45112ms step_avg:46.99ms step:961/1555 train_time:45169ms step_avg:47.00ms step:962/1555 train_time:45234ms step_avg:47.02ms step:963/1555 train_time:45291ms step_avg:47.03ms step:964/1555 train_time:45355ms step_avg:47.05ms step:965/1555 train_time:45412ms step_avg:47.06ms step:966/1555 train_time:45475ms step_avg:47.08ms step:967/1555 train_time:45533ms step_avg:47.09ms step:968/1555 train_time:45598ms step_avg:47.11ms step:969/1555 train_time:45656ms step_avg:47.12ms step:970/1555 train_time:45721ms step_avg:47.14ms step:971/1555 train_time:45779ms step_avg:47.15ms step:972/1555 train_time:45844ms step_avg:47.16ms step:973/1555 train_time:45902ms step_avg:47.18ms step:974/1555 train_time:45966ms step_avg:47.19ms step:975/1555 train_time:46024ms step_avg:47.20ms step:976/1555 train_time:46089ms step_avg:47.22ms step:977/1555 train_time:46147ms step_avg:47.23ms step:978/1555 train_time:46211ms step_avg:47.25ms step:979/1555 train_time:46270ms step_avg:47.26ms step:980/1555 train_time:46333ms step_avg:47.28ms step:981/1555 train_time:46391ms step_avg:47.29ms step:982/1555 train_time:46455ms step_avg:47.31ms step:983/1555 train_time:46512ms step_avg:47.32ms step:984/1555 train_time:46576ms step_avg:47.33ms step:985/1555 train_time:46635ms step_avg:47.34ms step:986/1555 train_time:46699ms step_avg:47.36ms step:987/1555 train_time:46757ms step_avg:47.37ms step:988/1555 train_time:46820ms step_avg:47.39ms step:989/1555 train_time:46878ms step_avg:47.40ms step:990/1555 train_time:46944ms step_avg:47.42ms step:991/1555 train_time:47002ms step_avg:47.43ms step:992/1555 train_time:47066ms step_avg:47.45ms step:993/1555 train_time:47125ms step_avg:47.46ms step:994/1555 train_time:47189ms step_avg:47.47ms step:995/1555 train_time:47247ms step_avg:47.48ms step:996/1555 train_time:47311ms step_avg:47.50ms step:997/1555 train_time:47370ms step_avg:47.51ms step:998/1555 train_time:47434ms step_avg:47.53ms step:999/1555 train_time:47492ms step_avg:47.54ms step:1000/1555 train_time:47555ms step_avg:47.56ms step:1000/1555 val_loss:3.5743 train_time:47637ms step_avg:47.64ms step:1001/1555 train_time:47660ms step_avg:47.61ms step:1002/1555 train_time:47684ms step_avg:47.59ms step:1003/1555 train_time:47737ms step_avg:47.59ms step:1004/1555 train_time:47807ms step_avg:47.62ms step:1005/1555 train_time:47866ms step_avg:47.63ms step:1006/1555 train_time:47930ms step_avg:47.64ms step:1007/1555 train_time:47990ms step_avg:47.66ms step:1008/1555 train_time:48053ms step_avg:47.67ms step:1009/1555 train_time:48109ms step_avg:47.68ms step:1010/1555 train_time:48172ms step_avg:47.69ms step:1011/1555 train_time:48233ms step_avg:47.71ms step:1012/1555 train_time:48317ms step_avg:47.74ms step:1013/1555 train_time:48400ms step_avg:47.78ms step:1014/1555 train_time:48491ms step_avg:47.82ms step:1015/1555 train_time:48573ms step_avg:47.85ms step:1016/1555 train_time:48663ms step_avg:47.90ms step:1017/1555 train_time:48750ms step_avg:47.94ms step:1018/1555 train_time:48843ms step_avg:47.98ms step:1019/1555 train_time:48931ms step_avg:48.02ms step:1020/1555 train_time:49021ms step_avg:48.06ms step:1021/1555 train_time:49106ms step_avg:48.10ms step:1022/1555 train_time:49196ms step_avg:48.14ms step:1023/1555 train_time:49280ms step_avg:48.17ms step:1024/1555 train_time:49369ms step_avg:48.21ms step:1025/1555 train_time:49451ms step_avg:48.25ms step:1026/1555 train_time:49540ms step_avg:48.28ms step:1027/1555 train_time:49624ms step_avg:48.32ms step:1028/1555 train_time:49715ms step_avg:48.36ms step:1029/1555 train_time:49800ms step_avg:48.40ms step:1030/1555 train_time:49894ms step_avg:48.44ms step:1031/1555 train_time:49977ms step_avg:48.47ms step:1032/1555 train_time:50067ms step_avg:48.51ms step:1033/1555 train_time:50152ms step_avg:48.55ms step:1034/1555 train_time:50241ms step_avg:48.59ms step:1035/1555 train_time:50325ms step_avg:48.62ms step:1036/1555 train_time:50413ms step_avg:48.66ms step:1037/1555 train_time:50497ms step_avg:48.70ms step:1038/1555 train_time:50586ms step_avg:48.73ms step:1039/1555 train_time:50670ms step_avg:48.77ms step:1040/1555 train_time:50760ms step_avg:48.81ms step:1041/1555 train_time:50846ms step_avg:48.84ms step:1042/1555 train_time:50936ms step_avg:48.88ms step:1043/1555 train_time:51021ms step_avg:48.92ms step:1044/1555 train_time:51111ms step_avg:48.96ms step:1045/1555 train_time:51195ms step_avg:48.99ms step:1046/1555 train_time:51284ms step_avg:49.03ms step:1047/1555 train_time:51368ms step_avg:49.06ms step:1048/1555 train_time:51457ms step_avg:49.10ms step:1049/1555 train_time:51540ms step_avg:49.13ms step:1050/1555 train_time:51631ms step_avg:49.17ms step:1051/1555 train_time:51715ms step_avg:49.21ms step:1052/1555 train_time:51807ms step_avg:49.25ms step:1053/1555 train_time:51891ms step_avg:49.28ms step:1054/1555 train_time:51981ms step_avg:49.32ms step:1055/1555 train_time:52067ms step_avg:49.35ms step:1056/1555 train_time:52156ms step_avg:49.39ms step:1057/1555 train_time:52241ms step_avg:49.42ms step:1058/1555 train_time:52332ms step_avg:49.46ms step:1059/1555 train_time:52414ms step_avg:49.49ms step:1060/1555 train_time:52504ms step_avg:49.53ms step:1061/1555 train_time:52589ms step_avg:49.57ms step:1062/1555 train_time:52678ms step_avg:49.60ms step:1063/1555 train_time:52763ms step_avg:49.64ms step:1064/1555 train_time:52852ms step_avg:49.67ms step:1065/1555 train_time:52936ms step_avg:49.71ms step:1066/1555 train_time:53029ms step_avg:49.75ms step:1067/1555 train_time:53112ms step_avg:49.78ms step:1068/1555 train_time:53202ms step_avg:49.81ms step:1069/1555 train_time:53286ms step_avg:49.85ms step:1070/1555 train_time:53375ms step_avg:49.88ms step:1071/1555 train_time:53458ms step_avg:49.91ms step:1072/1555 train_time:53548ms step_avg:49.95ms step:1073/1555 train_time:53632ms step_avg:49.98ms step:1074/1555 train_time:53723ms step_avg:50.02ms step:1075/1555 train_time:53805ms step_avg:50.05ms step:1076/1555 train_time:53895ms step_avg:50.09ms step:1077/1555 train_time:53979ms step_avg:50.12ms step:1078/1555 train_time:54071ms step_avg:50.16ms step:1079/1555 train_time:54154ms step_avg:50.19ms step:1080/1555 train_time:54245ms step_avg:50.23ms step:1081/1555 train_time:54328ms step_avg:50.26ms step:1082/1555 train_time:54419ms step_avg:50.29ms step:1083/1555 train_time:54503ms step_avg:50.33ms step:1084/1555 train_time:54592ms step_avg:50.36ms step:1085/1555 train_time:54676ms step_avg:50.39ms step:1086/1555 train_time:54766ms step_avg:50.43ms step:1087/1555 train_time:54850ms step_avg:50.46ms step:1088/1555 train_time:54939ms step_avg:50.50ms step:1089/1555 train_time:55024ms step_avg:50.53ms step:1090/1555 train_time:55113ms step_avg:50.56ms step:1091/1555 train_time:55198ms step_avg:50.59ms step:1092/1555 train_time:55288ms step_avg:50.63ms step:1093/1555 train_time:55372ms step_avg:50.66ms step:1094/1555 train_time:55462ms step_avg:50.70ms step:1095/1555 train_time:55547ms step_avg:50.73ms step:1096/1555 train_time:55636ms step_avg:50.76ms step:1097/1555 train_time:55720ms step_avg:50.79ms step:1098/1555 train_time:55811ms step_avg:50.83ms step:1099/1555 train_time:55895ms step_avg:50.86ms step:1100/1555 train_time:55985ms step_avg:50.90ms step:1101/1555 train_time:56070ms step_avg:50.93ms step:1102/1555 train_time:56159ms step_avg:50.96ms step:1103/1555 train_time:56243ms step_avg:50.99ms step:1104/1555 train_time:56333ms step_avg:51.03ms step:1105/1555 train_time:56418ms step_avg:51.06ms step:1106/1555 train_time:56509ms step_avg:51.09ms step:1107/1555 train_time:56592ms step_avg:51.12ms step:1108/1555 train_time:56682ms step_avg:51.16ms step:1109/1555 train_time:56766ms step_avg:51.19ms step:1110/1555 train_time:56855ms step_avg:51.22ms step:1111/1555 train_time:56939ms step_avg:51.25ms step:1112/1555 train_time:57030ms step_avg:51.29ms step:1113/1555 train_time:57114ms step_avg:51.32ms step:1114/1555 train_time:57204ms step_avg:51.35ms step:1115/1555 train_time:57288ms step_avg:51.38ms step:1116/1555 train_time:57378ms step_avg:51.41ms step:1117/1555 train_time:57462ms step_avg:51.44ms step:1118/1555 train_time:57552ms step_avg:51.48ms step:1119/1555 train_time:57636ms step_avg:51.51ms step:1120/1555 train_time:57726ms step_avg:51.54ms step:1121/1555 train_time:57810ms step_avg:51.57ms step:1122/1555 train_time:57899ms step_avg:51.60ms step:1123/1555 train_time:57983ms step_avg:51.63ms step:1124/1555 train_time:58074ms step_avg:51.67ms step:1125/1555 train_time:58157ms step_avg:51.70ms step:1126/1555 train_time:58247ms step_avg:51.73ms step:1127/1555 train_time:58332ms step_avg:51.76ms step:1128/1555 train_time:58421ms step_avg:51.79ms step:1129/1555 train_time:58505ms step_avg:51.82ms step:1130/1555 train_time:58595ms step_avg:51.85ms step:1131/1555 train_time:58679ms step_avg:51.88ms step:1132/1555 train_time:58768ms step_avg:51.92ms step:1133/1555 train_time:58852ms step_avg:51.94ms step:1134/1555 train_time:58942ms step_avg:51.98ms step:1135/1555 train_time:59026ms step_avg:52.01ms step:1136/1555 train_time:59115ms step_avg:52.04ms step:1137/1555 train_time:59198ms step_avg:52.07ms step:1138/1555 train_time:59289ms step_avg:52.10ms step:1139/1555 train_time:59373ms step_avg:52.13ms step:1140/1555 train_time:59464ms step_avg:52.16ms step:1141/1555 train_time:59548ms step_avg:52.19ms step:1142/1555 train_time:59637ms step_avg:52.22ms step:1143/1555 train_time:59721ms step_avg:52.25ms step:1144/1555 train_time:59811ms step_avg:52.28ms step:1145/1555 train_time:59895ms step_avg:52.31ms step:1146/1555 train_time:59986ms step_avg:52.34ms step:1147/1555 train_time:60070ms step_avg:52.37ms step:1148/1555 train_time:60159ms step_avg:52.40ms step:1149/1555 train_time:60244ms step_avg:52.43ms step:1150/1555 train_time:60334ms step_avg:52.46ms step:1151/1555 train_time:60418ms step_avg:52.49ms step:1152/1555 train_time:60509ms step_avg:52.53ms step:1153/1555 train_time:60593ms step_avg:52.55ms step:1154/1555 train_time:60683ms step_avg:52.59ms step:1155/1555 train_time:60767ms step_avg:52.61ms step:1156/1555 train_time:60857ms step_avg:52.64ms step:1157/1555 train_time:60941ms step_avg:52.67ms step:1158/1555 train_time:61031ms step_avg:52.70ms step:1159/1555 train_time:61115ms step_avg:52.73ms step:1160/1555 train_time:61205ms step_avg:52.76ms step:1161/1555 train_time:61289ms step_avg:52.79ms step:1162/1555 train_time:61379ms step_avg:52.82ms step:1163/1555 train_time:61463ms step_avg:52.85ms step:1164/1555 train_time:61553ms step_avg:52.88ms step:1165/1555 train_time:61638ms step_avg:52.91ms step:1166/1555 train_time:61728ms step_avg:52.94ms step:1167/1555 train_time:61811ms step_avg:52.97ms step:1168/1555 train_time:61902ms step_avg:53.00ms step:1169/1555 train_time:61986ms step_avg:53.02ms step:1170/1555 train_time:62075ms step_avg:53.06ms step:1171/1555 train_time:62158ms step_avg:53.08ms step:1172/1555 train_time:62248ms step_avg:53.11ms step:1173/1555 train_time:62332ms step_avg:53.14ms step:1174/1555 train_time:62422ms step_avg:53.17ms step:1175/1555 train_time:62506ms step_avg:53.20ms step:1176/1555 train_time:62596ms step_avg:53.23ms step:1177/1555 train_time:62679ms step_avg:53.25ms step:1178/1555 train_time:62770ms step_avg:53.29ms step:1179/1555 train_time:62854ms step_avg:53.31ms step:1180/1555 train_time:62945ms step_avg:53.34ms step:1181/1555 train_time:63028ms step_avg:53.37ms step:1182/1555 train_time:63118ms step_avg:53.40ms step:1183/1555 train_time:63202ms step_avg:53.43ms step:1184/1555 train_time:63292ms step_avg:53.46ms step:1185/1555 train_time:63376ms step_avg:53.48ms step:1186/1555 train_time:63467ms step_avg:53.51ms step:1187/1555 train_time:63551ms step_avg:53.54ms step:1188/1555 train_time:63640ms step_avg:53.57ms step:1189/1555 train_time:63724ms step_avg:53.59ms step:1190/1555 train_time:63814ms step_avg:53.63ms step:1191/1555 train_time:63898ms step_avg:53.65ms step:1192/1555 train_time:63988ms step_avg:53.68ms step:1193/1555 train_time:64073ms step_avg:53.71ms step:1194/1555 train_time:64162ms step_avg:53.74ms step:1195/1555 train_time:64246ms step_avg:53.76ms step:1196/1555 train_time:64336ms step_avg:53.79ms step:1197/1555 train_time:64420ms step_avg:53.82ms step:1198/1555 train_time:64510ms step_avg:53.85ms step:1199/1555 train_time:64594ms step_avg:53.87ms step:1200/1555 train_time:64684ms step_avg:53.90ms step:1201/1555 train_time:64768ms step_avg:53.93ms step:1202/1555 train_time:64858ms step_avg:53.96ms step:1203/1555 train_time:64943ms step_avg:53.98ms step:1204/1555 train_time:65034ms step_avg:54.01ms step:1205/1555 train_time:65117ms step_avg:54.04ms step:1206/1555 train_time:65209ms step_avg:54.07ms step:1207/1555 train_time:65293ms step_avg:54.10ms step:1208/1555 train_time:65383ms step_avg:54.12ms step:1209/1555 train_time:65467ms step_avg:54.15ms step:1210/1555 train_time:65556ms step_avg:54.18ms step:1211/1555 train_time:65641ms step_avg:54.20ms step:1212/1555 train_time:65732ms step_avg:54.23ms step:1213/1555 train_time:65816ms step_avg:54.26ms step:1214/1555 train_time:65905ms step_avg:54.29ms step:1215/1555 train_time:65990ms step_avg:54.31ms step:1216/1555 train_time:66080ms step_avg:54.34ms step:1217/1555 train_time:66164ms step_avg:54.37ms step:1218/1555 train_time:66254ms step_avg:54.40ms step:1219/1555 train_time:66336ms step_avg:54.42ms step:1220/1555 train_time:66427ms step_avg:54.45ms step:1221/1555 train_time:66511ms step_avg:54.47ms step:1222/1555 train_time:66603ms step_avg:54.50ms step:1223/1555 train_time:66687ms step_avg:54.53ms step:1224/1555 train_time:66776ms step_avg:54.56ms step:1225/1555 train_time:66860ms step_avg:54.58ms step:1226/1555 train_time:66951ms step_avg:54.61ms step:1227/1555 train_time:67035ms step_avg:54.63ms step:1228/1555 train_time:67124ms step_avg:54.66ms step:1229/1555 train_time:67208ms step_avg:54.69ms step:1230/1555 train_time:67299ms step_avg:54.71ms step:1231/1555 train_time:67384ms step_avg:54.74ms step:1232/1555 train_time:67475ms step_avg:54.77ms step:1233/1555 train_time:67558ms step_avg:54.79ms step:1234/1555 train_time:67649ms step_avg:54.82ms step:1235/1555 train_time:67733ms step_avg:54.84ms step:1236/1555 train_time:67822ms step_avg:54.87ms step:1237/1555 train_time:67906ms step_avg:54.90ms step:1238/1555 train_time:67996ms step_avg:54.92ms step:1239/1555 train_time:68081ms step_avg:54.95ms step:1240/1555 train_time:68171ms step_avg:54.98ms step:1241/1555 train_time:68255ms step_avg:55.00ms step:1242/1555 train_time:68344ms step_avg:55.03ms step:1243/1555 train_time:68429ms step_avg:55.05ms step:1244/1555 train_time:68518ms step_avg:55.08ms step:1245/1555 train_time:68603ms step_avg:55.10ms step:1246/1555 train_time:68694ms step_avg:55.13ms step:1247/1555 train_time:68776ms step_avg:55.15ms step:1248/1555 train_time:68867ms step_avg:55.18ms step:1249/1555 train_time:68951ms step_avg:55.20ms step:1250/1555 train_time:69041ms step_avg:55.23ms step:1250/1555 val_loss:3.3996 train_time:69156ms step_avg:55.32ms step:1251/1555 train_time:69176ms step_avg:55.30ms step:1252/1555 train_time:69216ms step_avg:55.28ms step:1253/1555 train_time:69303ms step_avg:55.31ms step:1254/1555 train_time:69397ms step_avg:55.34ms step:1255/1555 train_time:69480ms step_avg:55.36ms step:1256/1555 train_time:69569ms step_avg:55.39ms step:1257/1555 train_time:69652ms step_avg:55.41ms step:1258/1555 train_time:69741ms step_avg:55.44ms step:1259/1555 train_time:69824ms step_avg:55.46ms step:1260/1555 train_time:69914ms step_avg:55.49ms step:1261/1555 train_time:69997ms step_avg:55.51ms step:1262/1555 train_time:70087ms step_avg:55.54ms step:1263/1555 train_time:70172ms step_avg:55.56ms step:1264/1555 train_time:70266ms step_avg:55.59ms step:1265/1555 train_time:70354ms step_avg:55.62ms step:1266/1555 train_time:70444ms step_avg:55.64ms step:1267/1555 train_time:70529ms step_avg:55.67ms step:1268/1555 train_time:70619ms step_avg:55.69ms step:1269/1555 train_time:70701ms step_avg:55.71ms step:1270/1555 train_time:70790ms step_avg:55.74ms step:1271/1555 train_time:70874ms step_avg:55.76ms step:1272/1555 train_time:70963ms step_avg:55.79ms step:1273/1555 train_time:71047ms step_avg:55.81ms step:1274/1555 train_time:71138ms step_avg:55.84ms step:1275/1555 train_time:71223ms step_avg:55.86ms step:1276/1555 train_time:71316ms step_avg:55.89ms step:1277/1555 train_time:71400ms step_avg:55.91ms step:1278/1555 train_time:71491ms step_avg:55.94ms step:1279/1555 train_time:71575ms step_avg:55.96ms step:1280/1555 train_time:71665ms step_avg:55.99ms step:1281/1555 train_time:71749ms step_avg:56.01ms step:1282/1555 train_time:71839ms step_avg:56.04ms step:1283/1555 train_time:71921ms step_avg:56.06ms step:1284/1555 train_time:72011ms step_avg:56.08ms step:1285/1555 train_time:72096ms step_avg:56.11ms step:1286/1555 train_time:72186ms step_avg:56.13ms step:1287/1555 train_time:72272ms step_avg:56.16ms step:1288/1555 train_time:72362ms step_avg:56.18ms step:1289/1555 train_time:72447ms step_avg:56.20ms step:1290/1555 train_time:72538ms step_avg:56.23ms step:1291/1555 train_time:72622ms step_avg:56.25ms step:1292/1555 train_time:72712ms step_avg:56.28ms step:1293/1555 train_time:72796ms step_avg:56.30ms step:1294/1555 train_time:72885ms step_avg:56.32ms step:1295/1555 train_time:72969ms step_avg:56.35ms step:1296/1555 train_time:73058ms step_avg:56.37ms step:1297/1555 train_time:73142ms step_avg:56.39ms step:1298/1555 train_time:73233ms step_avg:56.42ms step:1299/1555 train_time:73317ms step_avg:56.44ms step:1300/1555 train_time:73408ms step_avg:56.47ms step:1301/1555 train_time:73492ms step_avg:56.49ms step:1302/1555 train_time:73582ms step_avg:56.51ms step:1303/1555 train_time:73666ms step_avg:56.54ms step:1304/1555 train_time:73756ms step_avg:56.56ms step:1305/1555 train_time:73839ms step_avg:56.58ms step:1306/1555 train_time:73929ms step_avg:56.61ms step:1307/1555 train_time:74014ms step_avg:56.63ms step:1308/1555 train_time:74103ms step_avg:56.65ms step:1309/1555 train_time:74188ms step_avg:56.68ms step:1310/1555 train_time:74277ms step_avg:56.70ms step:1311/1555 train_time:74361ms step_avg:56.72ms step:1312/1555 train_time:74453ms step_avg:56.75ms step:1313/1555 train_time:74537ms step_avg:56.77ms step:1314/1555 train_time:74627ms step_avg:56.79ms step:1315/1555 train_time:74711ms step_avg:56.81ms step:1316/1555 train_time:74801ms step_avg:56.84ms step:1317/1555 train_time:74884ms step_avg:56.86ms step:1318/1555 train_time:74975ms step_avg:56.89ms step:1319/1555 train_time:75059ms step_avg:56.91ms step:1320/1555 train_time:75149ms step_avg:56.93ms step:1321/1555 train_time:75233ms step_avg:56.95ms step:1322/1555 train_time:75323ms step_avg:56.98ms step:1323/1555 train_time:75407ms step_avg:57.00ms step:1324/1555 train_time:75499ms step_avg:57.02ms step:1325/1555 train_time:75582ms step_avg:57.04ms step:1326/1555 train_time:75672ms step_avg:57.07ms step:1327/1555 train_time:75756ms step_avg:57.09ms step:1328/1555 train_time:75846ms step_avg:57.11ms step:1329/1555 train_time:75931ms step_avg:57.13ms step:1330/1555 train_time:76020ms step_avg:57.16ms step:1331/1555 train_time:76104ms step_avg:57.18ms step:1332/1555 train_time:76195ms step_avg:57.20ms step:1333/1555 train_time:76279ms step_avg:57.22ms step:1334/1555 train_time:76368ms step_avg:57.25ms step:1335/1555 train_time:76454ms step_avg:57.27ms step:1336/1555 train_time:76543ms step_avg:57.29ms step:1337/1555 train_time:76627ms step_avg:57.31ms step:1338/1555 train_time:76717ms step_avg:57.34ms step:1339/1555 train_time:76802ms step_avg:57.36ms step:1340/1555 train_time:76891ms step_avg:57.38ms step:1341/1555 train_time:76975ms step_avg:57.40ms step:1342/1555 train_time:77065ms step_avg:57.43ms step:1343/1555 train_time:77150ms step_avg:57.45ms step:1344/1555 train_time:77240ms step_avg:57.47ms step:1345/1555 train_time:77323ms step_avg:57.49ms step:1346/1555 train_time:77415ms step_avg:57.51ms step:1347/1555 train_time:77499ms step_avg:57.53ms step:1348/1555 train_time:77590ms step_avg:57.56ms step:1349/1555 train_time:77675ms step_avg:57.58ms step:1350/1555 train_time:77764ms step_avg:57.60ms step:1351/1555 train_time:77848ms step_avg:57.62ms step:1352/1555 train_time:77938ms step_avg:57.65ms step:1353/1555 train_time:78021ms step_avg:57.67ms step:1354/1555 train_time:78113ms step_avg:57.69ms step:1355/1555 train_time:78197ms step_avg:57.71ms step:1356/1555 train_time:78286ms step_avg:57.73ms step:1357/1555 train_time:78372ms step_avg:57.75ms step:1358/1555 train_time:78461ms step_avg:57.78ms step:1359/1555 train_time:78546ms step_avg:57.80ms step:1360/1555 train_time:78636ms step_avg:57.82ms step:1361/1555 train_time:78719ms step_avg:57.84ms step:1362/1555 train_time:78810ms step_avg:57.86ms step:1363/1555 train_time:78895ms step_avg:57.88ms step:1364/1555 train_time:78983ms step_avg:57.91ms step:1365/1555 train_time:79068ms step_avg:57.93ms step:1366/1555 train_time:79157ms step_avg:57.95ms step:1367/1555 train_time:79241ms step_avg:57.97ms step:1368/1555 train_time:79331ms step_avg:57.99ms step:1369/1555 train_time:79415ms step_avg:58.01ms step:1370/1555 train_time:79505ms step_avg:58.03ms step:1371/1555 train_time:79590ms step_avg:58.05ms step:1372/1555 train_time:79680ms step_avg:58.08ms step:1373/1555 train_time:79763ms step_avg:58.09ms step:1374/1555 train_time:79854ms step_avg:58.12ms step:1375/1555 train_time:79938ms step_avg:58.14ms step:1376/1555 train_time:80028ms step_avg:58.16ms step:1377/1555 train_time:80111ms step_avg:58.18ms step:1378/1555 train_time:80201ms step_avg:58.20ms step:1379/1555 train_time:80285ms step_avg:58.22ms step:1380/1555 train_time:80375ms step_avg:58.24ms step:1381/1555 train_time:80459ms step_avg:58.26ms step:1382/1555 train_time:80549ms step_avg:58.28ms step:1383/1555 train_time:80633ms step_avg:58.30ms step:1384/1555 train_time:80723ms step_avg:58.33ms step:1385/1555 train_time:80807ms step_avg:58.34ms step:1386/1555 train_time:80897ms step_avg:58.37ms step:1387/1555 train_time:80980ms step_avg:58.39ms step:1388/1555 train_time:81071ms step_avg:58.41ms step:1389/1555 train_time:81155ms step_avg:58.43ms step:1390/1555 train_time:81245ms step_avg:58.45ms step:1391/1555 train_time:81330ms step_avg:58.47ms step:1392/1555 train_time:81420ms step_avg:58.49ms step:1393/1555 train_time:81503ms step_avg:58.51ms step:1394/1555 train_time:81594ms step_avg:58.53ms step:1395/1555 train_time:81678ms step_avg:58.55ms step:1396/1555 train_time:81768ms step_avg:58.57ms step:1397/1555 train_time:81852ms step_avg:58.59ms step:1398/1555 train_time:81942ms step_avg:58.61ms step:1399/1555 train_time:82026ms step_avg:58.63ms step:1400/1555 train_time:82118ms step_avg:58.66ms step:1401/1555 train_time:82202ms step_avg:58.67ms step:1402/1555 train_time:82292ms step_avg:58.70ms step:1403/1555 train_time:82376ms step_avg:58.71ms step:1404/1555 train_time:82465ms step_avg:58.74ms step:1405/1555 train_time:82549ms step_avg:58.75ms step:1406/1555 train_time:82639ms step_avg:58.78ms step:1407/1555 train_time:82723ms step_avg:58.79ms step:1408/1555 train_time:82814ms step_avg:58.82ms step:1409/1555 train_time:82898ms step_avg:58.83ms step:1410/1555 train_time:82988ms step_avg:58.86ms step:1411/1555 train_time:83073ms step_avg:58.88ms step:1412/1555 train_time:83163ms step_avg:58.90ms step:1413/1555 train_time:83248ms step_avg:58.92ms step:1414/1555 train_time:83339ms step_avg:58.94ms step:1415/1555 train_time:83423ms step_avg:58.96ms step:1416/1555 train_time:83513ms step_avg:58.98ms step:1417/1555 train_time:83597ms step_avg:59.00ms step:1418/1555 train_time:83687ms step_avg:59.02ms step:1419/1555 train_time:83773ms step_avg:59.04ms step:1420/1555 train_time:83861ms step_avg:59.06ms step:1421/1555 train_time:83945ms step_avg:59.07ms step:1422/1555 train_time:84036ms step_avg:59.10ms step:1423/1555 train_time:84119ms step_avg:59.11ms step:1424/1555 train_time:84210ms step_avg:59.14ms step:1425/1555 train_time:84295ms step_avg:59.15ms step:1426/1555 train_time:84385ms step_avg:59.18ms step:1427/1555 train_time:84470ms step_avg:59.19ms step:1428/1555 train_time:84559ms step_avg:59.21ms step:1429/1555 train_time:84643ms step_avg:59.23ms step:1430/1555 train_time:84734ms step_avg:59.25ms step:1431/1555 train_time:84818ms step_avg:59.27ms step:1432/1555 train_time:84908ms step_avg:59.29ms step:1433/1555 train_time:84991ms step_avg:59.31ms step:1434/1555 train_time:85081ms step_avg:59.33ms step:1435/1555 train_time:85165ms step_avg:59.35ms step:1436/1555 train_time:85256ms step_avg:59.37ms step:1437/1555 train_time:85339ms step_avg:59.39ms step:1438/1555 train_time:85430ms step_avg:59.41ms step:1439/1555 train_time:85514ms step_avg:59.43ms step:1440/1555 train_time:85604ms step_avg:59.45ms step:1441/1555 train_time:85689ms step_avg:59.47ms step:1442/1555 train_time:85780ms step_avg:59.49ms step:1443/1555 train_time:85863ms step_avg:59.50ms step:1444/1555 train_time:85953ms step_avg:59.52ms step:1445/1555 train_time:86038ms step_avg:59.54ms step:1446/1555 train_time:86128ms step_avg:59.56ms step:1447/1555 train_time:86213ms step_avg:59.58ms step:1448/1555 train_time:86302ms step_avg:59.60ms step:1449/1555 train_time:86387ms step_avg:59.62ms step:1450/1555 train_time:86477ms step_avg:59.64ms step:1451/1555 train_time:86561ms step_avg:59.66ms step:1452/1555 train_time:86651ms step_avg:59.68ms step:1453/1555 train_time:86735ms step_avg:59.69ms step:1454/1555 train_time:86825ms step_avg:59.71ms step:1455/1555 train_time:86909ms step_avg:59.73ms step:1456/1555 train_time:87000ms step_avg:59.75ms step:1457/1555 train_time:87083ms step_avg:59.77ms step:1458/1555 train_time:87174ms step_avg:59.79ms step:1459/1555 train_time:87258ms step_avg:59.81ms step:1460/1555 train_time:87348ms step_avg:59.83ms step:1461/1555 train_time:87432ms step_avg:59.84ms step:1462/1555 train_time:87521ms step_avg:59.86ms step:1463/1555 train_time:87606ms step_avg:59.88ms step:1464/1555 train_time:87698ms step_avg:59.90ms step:1465/1555 train_time:87781ms step_avg:59.92ms step:1466/1555 train_time:87872ms step_avg:59.94ms step:1467/1555 train_time:87955ms step_avg:59.96ms step:1468/1555 train_time:88045ms step_avg:59.98ms step:1469/1555 train_time:88130ms step_avg:59.99ms step:1470/1555 train_time:88219ms step_avg:60.01ms step:1471/1555 train_time:88302ms step_avg:60.03ms step:1472/1555 train_time:88393ms step_avg:60.05ms step:1473/1555 train_time:88478ms step_avg:60.07ms step:1474/1555 train_time:88567ms step_avg:60.09ms step:1475/1555 train_time:88652ms step_avg:60.10ms step:1476/1555 train_time:88741ms step_avg:60.12ms step:1477/1555 train_time:88825ms step_avg:60.14ms step:1478/1555 train_time:88916ms step_avg:60.16ms step:1479/1555 train_time:89000ms step_avg:60.18ms step:1480/1555 train_time:89090ms step_avg:60.20ms step:1481/1555 train_time:89176ms step_avg:60.21ms step:1482/1555 train_time:89264ms step_avg:60.23ms step:1483/1555 train_time:89349ms step_avg:60.25ms step:1484/1555 train_time:89439ms step_avg:60.27ms step:1485/1555 train_time:89523ms step_avg:60.28ms step:1486/1555 train_time:89613ms step_avg:60.31ms step:1487/1555 train_time:89698ms step_avg:60.32ms step:1488/1555 train_time:89786ms step_avg:60.34ms step:1489/1555 train_time:89870ms step_avg:60.36ms step:1490/1555 train_time:89961ms step_avg:60.38ms step:1491/1555 train_time:90044ms step_avg:60.39ms step:1492/1555 train_time:90133ms step_avg:60.41ms step:1493/1555 train_time:90218ms step_avg:60.43ms step:1494/1555 train_time:90308ms step_avg:60.45ms step:1495/1555 train_time:90393ms step_avg:60.46ms step:1496/1555 train_time:90482ms step_avg:60.48ms step:1497/1555 train_time:90568ms step_avg:60.50ms step:1498/1555 train_time:90657ms step_avg:60.52ms step:1499/1555 train_time:90740ms step_avg:60.53ms step:1500/1555 train_time:90831ms step_avg:60.55ms step:1500/1555 val_loss:3.2959 train_time:90947ms step_avg:60.63ms step:1501/1555 train_time:90966ms step_avg:60.60ms step:1502/1555 train_time:91009ms step_avg:60.59ms step:1503/1555 train_time:91100ms step_avg:60.61ms step:1504/1555 train_time:91192ms step_avg:60.63ms step:1505/1555 train_time:91277ms step_avg:60.65ms step:1506/1555 train_time:91368ms step_avg:60.67ms step:1507/1555 train_time:91450ms step_avg:60.68ms step:1508/1555 train_time:91539ms step_avg:60.70ms step:1509/1555 train_time:91621ms step_avg:60.72ms step:1510/1555 train_time:91710ms step_avg:60.74ms step:1511/1555 train_time:91793ms step_avg:60.75ms step:1512/1555 train_time:91884ms step_avg:60.77ms step:1513/1555 train_time:91969ms step_avg:60.79ms step:1514/1555 train_time:92064ms step_avg:60.81ms step:1515/1555 train_time:92149ms step_avg:60.82ms step:1516/1555 train_time:92244ms step_avg:60.85ms step:1517/1555 train_time:92330ms step_avg:60.86ms step:1518/1555 train_time:92418ms step_avg:60.88ms step:1519/1555 train_time:92502ms step_avg:60.90ms step:1520/1555 train_time:92591ms step_avg:60.92ms step:1521/1555 train_time:92675ms step_avg:60.93ms step:1522/1555 train_time:92764ms step_avg:60.95ms step:1523/1555 train_time:92847ms step_avg:60.96ms step:1524/1555 train_time:92938ms step_avg:60.98ms step:1525/1555 train_time:93026ms step_avg:61.00ms step:1526/1555 train_time:93119ms step_avg:61.02ms step:1527/1555 train_time:93204ms step_avg:61.04ms step:1528/1555 train_time:93295ms step_avg:61.06ms step:1529/1555 train_time:93379ms step_avg:61.07ms step:1530/1555 train_time:93469ms step_avg:61.09ms step:1531/1555 train_time:93552ms step_avg:61.11ms step:1532/1555 train_time:93642ms step_avg:61.12ms step:1533/1555 train_time:93726ms step_avg:61.14ms step:1534/1555 train_time:93815ms step_avg:61.16ms step:1535/1555 train_time:93899ms step_avg:61.17ms step:1536/1555 train_time:93991ms step_avg:61.19ms step:1537/1555 train_time:94077ms step_avg:61.21ms step:1538/1555 train_time:94168ms step_avg:61.23ms step:1539/1555 train_time:94253ms step_avg:61.24ms step:1540/1555 train_time:94344ms step_avg:61.26ms step:1541/1555 train_time:94428ms step_avg:61.28ms step:1542/1555 train_time:94519ms step_avg:61.30ms step:1543/1555 train_time:94603ms step_avg:61.31ms step:1544/1555 train_time:94692ms step_avg:61.33ms step:1545/1555 train_time:94777ms step_avg:61.34ms step:1546/1555 train_time:94867ms step_avg:61.36ms step:1547/1555 train_time:94951ms step_avg:61.38ms step:1548/1555 train_time:95042ms step_avg:61.40ms step:1549/1555 train_time:95128ms step_avg:61.41ms step:1550/1555 train_time:95220ms step_avg:61.43ms step:1551/1555 train_time:95304ms step_avg:61.45ms step:1552/1555 train_time:95396ms step_avg:61.47ms step:1553/1555 train_time:95480ms step_avg:61.48ms step:1554/1555 train_time:95570ms step_avg:61.50ms step:1555/1555 train_time:95654ms step_avg:61.51ms step:1555/1555 val_loss:3.2796 train_time:95768ms step_avg:61.59ms peak memory allocated: 31630 MiB reserved: 46498 MiB