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:21:51 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 113W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 1 NVIDIA H100 80GB HBM3 On | 00000000:6B:00.0 Off | 0 | | N/A 30C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 2 NVIDIA H100 80GB HBM3 On | 00000000:71:00.0 Off | 0 | | N/A 31C P0 120W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 3 NVIDIA H100 80GB HBM3 On | 00000000:79:00.0 Off | 0 | | N/A 30C P0 122W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 4 NVIDIA H100 80GB HBM3 On | 00000000:7F:00.0 Off | 0 | | N/A 29C P0 117W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 5 NVIDIA H100 80GB HBM3 On | 00000000:87:00.0 Off | 0 | | N/A 32C P0 116W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 6 NVIDIA H100 80GB HBM3 On | 00000000:8D:00.0 Off | 0 | | N/A 31C P0 118W / 700W | 1519MiB / 81559MiB | 0% Default | | | | Disabled | +-----------------------------------------+------------------------+----------------------+ | 7 NVIDIA H100 80GB HBM3 On | 00000000:95:00.0 Off | 0 | | N/A 31C 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 95 C /usr/local/bin/python 1510MiB | | 1 N/A N/A 96 C /usr/local/bin/python 1510MiB | | 2 N/A N/A 97 C /usr/local/bin/python 1510MiB | | 3 N/A N/A 98 C /usr/local/bin/python 1510MiB | | 4 N/A N/A 99 C /usr/local/bin/python 1510MiB | | 5 N/A N/A 100 C /usr/local/bin/python 1510MiB | | 6 N/A N/A 101 C /usr/local/bin/python 1510MiB | | 7 N/A N/A 102 C /usr/local/bin/python 1510MiB | +-----------------------------------------------------------------------------------------+ ==================================================================================================== Compiling model and warming up kernels (~7 minutes on first execution) Sampling steps [0, 1, 2, 504, 505, 506, 1009, 1010, 1011, 1514, 1515, 1516] for warmup Resetting Model step:0/1555 val_loss:10.8314 train_time:0ms step_avg:0.03ms step:1/1555 train_time:93ms step_avg:93.13ms step:2/1555 train_time:127ms step_avg:63.41ms step:3/1555 train_time:155ms step_avg:51.75ms step:4/1555 train_time:180ms step_avg:45.11ms step:5/1555 train_time:203ms step_avg:40.56ms step:6/1555 train_time:234ms step_avg:39.06ms step:7/1555 train_time:265ms step_avg:37.90ms step:8/1555 train_time:303ms step_avg:37.84ms step:9/1555 train_time:334ms step_avg:37.07ms step:10/1555 train_time:371ms step_avg:37.13ms step:11/1555 train_time:402ms step_avg:36.54ms step:12/1555 train_time:440ms step_avg:36.64ms step:13/1555 train_time:471ms step_avg:36.21ms step:14/1555 train_time:509ms step_avg:36.32ms step:15/1555 train_time:539ms step_avg:35.96ms step:16/1555 train_time:577ms step_avg:36.04ms step:17/1555 train_time:607ms step_avg:35.73ms step:18/1555 train_time:645ms step_avg:35.82ms step:19/1555 train_time:676ms step_avg:35.57ms step:20/1555 train_time:714ms step_avg:35.68ms step:21/1555 train_time:745ms step_avg:35.45ms step:22/1555 train_time:782ms step_avg:35.55ms step:23/1555 train_time:813ms step_avg:35.36ms step:24/1555 train_time:851ms step_avg:35.46ms step:25/1555 train_time:882ms step_avg:35.29ms step:26/1555 train_time:920ms step_avg:35.37ms step:27/1555 train_time:950ms step_avg:35.19ms step:28/1555 train_time:987ms step_avg:35.26ms step:29/1555 train_time:1019ms step_avg:35.13ms step:30/1555 train_time:1057ms step_avg:35.23ms step:31/1555 train_time:1088ms step_avg:35.10ms step:32/1555 train_time:1126ms step_avg:35.17ms step:33/1555 train_time:1157ms step_avg:35.05ms step:34/1555 train_time:1194ms step_avg:35.13ms step:35/1555 train_time:1225ms step_avg:35.01ms step:36/1555 train_time:1263ms step_avg:35.08ms step:37/1555 train_time:1294ms step_avg:34.99ms step:38/1555 train_time:1332ms step_avg:35.06ms step:39/1555 train_time:1364ms step_avg:34.97ms step:40/1555 train_time:1401ms step_avg:35.03ms step:41/1555 train_time:1433ms step_avg:34.95ms step:42/1555 train_time:1471ms step_avg:35.01ms step:43/1555 train_time:1501ms step_avg:34.92ms step:44/1555 train_time:1539ms step_avg:34.97ms step:45/1555 train_time:1570ms step_avg:34.89ms step:46/1555 train_time:1607ms step_avg:34.94ms step:47/1555 train_time:1638ms step_avg:34.86ms step:48/1555 train_time:1676ms step_avg:34.91ms step:49/1555 train_time:1707ms step_avg:34.84ms step:50/1555 train_time:1744ms step_avg:34.89ms step:51/1555 train_time:1776ms step_avg:34.81ms step:52/1555 train_time:1814ms step_avg:34.88ms step:53/1555 train_time:1845ms step_avg:34.81ms step:54/1555 train_time:1882ms step_avg:34.85ms step:55/1555 train_time:1913ms step_avg:34.78ms step:56/1555 train_time:1951ms step_avg:34.84ms step:57/1555 train_time:1982ms step_avg:34.77ms step:58/1555 train_time:2020ms step_avg:34.82ms step:59/1555 train_time:2051ms step_avg:34.76ms step:60/1555 train_time:2088ms step_avg:34.80ms step:61/1555 train_time:2119ms step_avg:34.74ms step:62/1555 train_time:2157ms step_avg:34.79ms step:63/1555 train_time:2188ms step_avg:34.73ms step:64/1555 train_time:2226ms step_avg:34.78ms step:65/1555 train_time:2257ms step_avg:34.72ms step:66/1555 train_time:2294ms step_avg:34.76ms step:67/1555 train_time:2325ms step_avg:34.70ms step:68/1555 train_time:2362ms step_avg:34.74ms step:69/1555 train_time:2394ms step_avg:34.70ms step:70/1555 train_time:2432ms step_avg:34.74ms step:71/1555 train_time:2463ms step_avg:34.69ms step:72/1555 train_time:2500ms step_avg:34.73ms step:73/1555 train_time:2532ms step_avg:34.68ms step:74/1555 train_time:2569ms step_avg:34.72ms step:75/1555 train_time:2600ms step_avg:34.67ms step:76/1555 train_time:2638ms step_avg:34.71ms step:77/1555 train_time:2669ms step_avg:34.66ms step:78/1555 train_time:2707ms step_avg:34.70ms step:79/1555 train_time:2738ms step_avg:34.65ms step:80/1555 train_time:2775ms step_avg:34.69ms step:81/1555 train_time:2806ms step_avg:34.64ms step:82/1555 train_time:2843ms step_avg:34.67ms step:83/1555 train_time:2874ms step_avg:34.63ms step:84/1555 train_time:2912ms step_avg:34.67ms step:85/1555 train_time:2943ms step_avg:34.62ms step:86/1555 train_time:2980ms step_avg:34.65ms step:87/1555 train_time:3012ms step_avg:34.62ms step:88/1555 train_time:3049ms step_avg:34.65ms step:89/1555 train_time:3080ms step_avg:34.61ms step:90/1555 train_time:3118ms step_avg:34.64ms step:91/1555 train_time:3148ms step_avg:34.60ms step:92/1555 train_time:3186ms step_avg:34.63ms step:93/1555 train_time:3216ms step_avg:34.59ms step:94/1555 train_time:3254ms step_avg:34.62ms step:95/1555 train_time:3286ms step_avg:34.59ms step:96/1555 train_time:3323ms step_avg:34.62ms step:97/1555 train_time:3355ms step_avg:34.58ms step:98/1555 train_time:3392ms step_avg:34.62ms step:99/1555 train_time:3424ms step_avg:34.58ms step:100/1555 train_time:3461ms step_avg:34.61ms step:101/1555 train_time:3492ms step_avg:34.58ms step:102/1555 train_time:3530ms step_avg:34.61ms step:103/1555 train_time:3561ms step_avg:34.58ms step:104/1555 train_time:3599ms step_avg:34.60ms step:105/1555 train_time:3629ms step_avg:34.57ms step:106/1555 train_time:3667ms step_avg:34.59ms step:107/1555 train_time:3698ms step_avg:34.56ms step:108/1555 train_time:3736ms step_avg:34.59ms step:109/1555 train_time:3766ms step_avg:34.55ms step:110/1555 train_time:3804ms step_avg:34.58ms step:111/1555 train_time:3835ms step_avg:34.55ms step:112/1555 train_time:3873ms step_avg:34.58ms step:113/1555 train_time:3904ms step_avg:34.55ms step:114/1555 train_time:3941ms step_avg:34.57ms step:115/1555 train_time:3973ms step_avg:34.54ms step:116/1555 train_time:4011ms step_avg:34.58ms step:117/1555 train_time:4042ms step_avg:34.55ms step:118/1555 train_time:4080ms step_avg:34.57ms step:119/1555 train_time:4111ms step_avg:34.55ms step:120/1555 train_time:4149ms step_avg:34.57ms step:121/1555 train_time:4179ms step_avg:34.54ms step:122/1555 train_time:4217ms step_avg:34.57ms step:123/1555 train_time:4248ms step_avg:34.54ms step:124/1555 train_time:4286ms step_avg:34.56ms step:125/1555 train_time:4318ms step_avg:34.54ms step:126/1555 train_time:4355ms step_avg:34.57ms step:127/1555 train_time:4386ms step_avg:34.54ms step:128/1555 train_time:4424ms step_avg:34.56ms step:129/1555 train_time:4455ms step_avg:34.54ms step:130/1555 train_time:4493ms step_avg:34.56ms step:131/1555 train_time:4524ms step_avg:34.53ms step:132/1555 train_time:4562ms step_avg:34.56ms step:133/1555 train_time:4593ms step_avg:34.53ms step:134/1555 train_time:4630ms step_avg:34.56ms step:135/1555 train_time:4661ms step_avg:34.53ms step:136/1555 train_time:4699ms step_avg:34.55ms step:137/1555 train_time:4730ms step_avg:34.52ms step:138/1555 train_time:4767ms step_avg:34.55ms step:139/1555 train_time:4798ms step_avg:34.52ms step:140/1555 train_time:4836ms step_avg:34.54ms step:141/1555 train_time:4866ms step_avg:34.51ms step:142/1555 train_time:4904ms step_avg:34.53ms step:143/1555 train_time:4935ms step_avg:34.51ms step:144/1555 train_time:4973ms step_avg:34.53ms step:145/1555 train_time:5004ms step_avg:34.51ms step:146/1555 train_time:5041ms step_avg:34.53ms step:147/1555 train_time:5072ms step_avg:34.50ms step:148/1555 train_time:5110ms step_avg:34.53ms step:149/1555 train_time:5141ms step_avg:34.50ms step:150/1555 train_time:5179ms step_avg:34.52ms step:151/1555 train_time:5209ms step_avg:34.50ms step:152/1555 train_time:5247ms step_avg:34.52ms step:153/1555 train_time:5279ms step_avg:34.50ms step:154/1555 train_time:5316ms step_avg:34.52ms step:155/1555 train_time:5347ms step_avg:34.49ms step:156/1555 train_time:5384ms step_avg:34.51ms step:157/1555 train_time:5415ms step_avg:34.49ms step:158/1555 train_time:5453ms step_avg:34.51ms step:159/1555 train_time:5484ms step_avg:34.49ms step:160/1555 train_time:5522ms step_avg:34.51ms step:161/1555 train_time:5553ms step_avg:34.49ms step:162/1555 train_time:5590ms step_avg:34.51ms step:163/1555 train_time:5622ms step_avg:34.49ms step:164/1555 train_time:5659ms step_avg:34.51ms step:165/1555 train_time:5690ms step_avg:34.48ms step:166/1555 train_time:5727ms step_avg:34.50ms step:167/1555 train_time:5758ms step_avg:34.48ms step:168/1555 train_time:5795ms step_avg:34.50ms step:169/1555 train_time:5826ms step_avg:34.47ms step:170/1555 train_time:5863ms step_avg:34.49ms step:171/1555 train_time:5894ms step_avg:34.47ms step:172/1555 train_time:5932ms step_avg:34.49ms step:173/1555 train_time:5963ms step_avg:34.47ms step:174/1555 train_time:6001ms step_avg:34.49ms step:175/1555 train_time:6032ms step_avg:34.47ms step:176/1555 train_time:6069ms step_avg:34.48ms step:177/1555 train_time:6100ms step_avg:34.46ms step:178/1555 train_time:6138ms step_avg:34.48ms step:179/1555 train_time:6168ms step_avg:34.46ms step:180/1555 train_time:6206ms step_avg:34.48ms step:181/1555 train_time:6237ms step_avg:34.46ms step:182/1555 train_time:6275ms step_avg:34.48ms step:183/1555 train_time:6306ms step_avg:34.46ms step:184/1555 train_time:6343ms step_avg:34.47ms step:185/1555 train_time:6374ms step_avg:34.46ms step:186/1555 train_time:6412ms step_avg:34.47ms step:187/1555 train_time:6443ms step_avg:34.45ms step:188/1555 train_time:6480ms step_avg:34.47ms step:189/1555 train_time:6511ms step_avg:34.45ms step:190/1555 train_time:6549ms step_avg:34.47ms step:191/1555 train_time:6580ms step_avg:34.45ms step:192/1555 train_time:6617ms step_avg:34.46ms step:193/1555 train_time:6648ms step_avg:34.44ms step:194/1555 train_time:6685ms step_avg:34.46ms step:195/1555 train_time:6717ms step_avg:34.44ms step:196/1555 train_time:6754ms step_avg:34.46ms step:197/1555 train_time:6785ms step_avg:34.44ms step:198/1555 train_time:6822ms step_avg:34.46ms step:199/1555 train_time:6854ms step_avg:34.44ms step:200/1555 train_time:6891ms step_avg:34.46ms step:201/1555 train_time:6922ms step_avg:34.44ms step:202/1555 train_time:6959ms step_avg:34.45ms step:203/1555 train_time:6990ms step_avg:34.43ms step:204/1555 train_time:7028ms step_avg:34.45ms step:205/1555 train_time:7058ms step_avg:34.43ms step:206/1555 train_time:7096ms step_avg:34.44ms step:207/1555 train_time:7127ms step_avg:34.43ms step:208/1555 train_time:7164ms step_avg:34.44ms step:209/1555 train_time:7195ms step_avg:34.43ms step:210/1555 train_time:7232ms step_avg:34.44ms step:211/1555 train_time:7264ms step_avg:34.42ms step:212/1555 train_time:7301ms step_avg:34.44ms step:213/1555 train_time:7332ms step_avg:34.42ms step:214/1555 train_time:7369ms step_avg:34.43ms step:215/1555 train_time:7400ms step_avg:34.42ms step:216/1555 train_time:7437ms step_avg:34.43ms step:217/1555 train_time:7468ms step_avg:34.41ms step:218/1555 train_time:7506ms step_avg:34.43ms step:219/1555 train_time:7537ms step_avg:34.42ms step:220/1555 train_time:7575ms step_avg:34.43ms step:221/1555 train_time:7606ms step_avg:34.42ms step:222/1555 train_time:7643ms step_avg:34.43ms step:223/1555 train_time:7674ms step_avg:34.41ms step:224/1555 train_time:7711ms step_avg:34.43ms step:225/1555 train_time:7742ms step_avg:34.41ms step:226/1555 train_time:7780ms step_avg:34.42ms step:227/1555 train_time:7812ms step_avg:34.41ms step:228/1555 train_time:7849ms step_avg:34.43ms step:229/1555 train_time:7880ms step_avg:34.41ms step:230/1555 train_time:7917ms step_avg:34.42ms step:231/1555 train_time:7948ms step_avg:34.41ms step:232/1555 train_time:7985ms step_avg:34.42ms step:233/1555 train_time:8016ms step_avg:34.40ms step:234/1555 train_time:8053ms step_avg:34.42ms step:235/1555 train_time:8084ms step_avg:34.40ms step:236/1555 train_time:8121ms step_avg:34.41ms step:237/1555 train_time:8152ms step_avg:34.40ms step:238/1555 train_time:8190ms step_avg:34.41ms step:239/1555 train_time:8221ms step_avg:34.40ms step:240/1555 train_time:8258ms step_avg:34.41ms step:241/1555 train_time:8289ms step_avg:34.39ms step:242/1555 train_time:8326ms step_avg:34.40ms step:243/1555 train_time:8357ms step_avg:34.39ms step:244/1555 train_time:8394ms step_avg:34.40ms step:245/1555 train_time:8425ms step_avg:34.39ms step:246/1555 train_time:8462ms step_avg:34.40ms step:247/1555 train_time:8494ms step_avg:34.39ms step:248/1555 train_time:8532ms step_avg:34.40ms step:249/1555 train_time:8563ms step_avg:34.39ms step:250/1555 train_time:8600ms step_avg:34.40ms step:250/1555 val_loss:4.5535 train_time:8650ms step_avg:34.60ms step:251/1555 train_time:8669ms step_avg:34.54ms step:252/1555 train_time:8688ms step_avg:34.48ms step:253/1555 train_time:8706ms step_avg:34.41ms step:254/1555 train_time:8738ms step_avg:34.40ms step:255/1555 train_time:8771ms step_avg:34.39ms step:256/1555 train_time:8809ms step_avg:34.41ms step:257/1555 train_time:8841ms step_avg:34.40ms step:258/1555 train_time:8879ms step_avg:34.41ms step:259/1555 train_time:8910ms step_avg:34.40ms step:260/1555 train_time:8947ms step_avg:34.41ms step:261/1555 train_time:8979ms step_avg:34.40ms step:262/1555 train_time:9016ms step_avg:34.41ms step:263/1555 train_time:9047ms step_avg:34.40ms step:264/1555 train_time:9084ms step_avg:34.41ms step:265/1555 train_time:9115ms step_avg:34.40ms step:266/1555 train_time:9152ms step_avg:34.41ms step:267/1555 train_time:9183ms step_avg:34.39ms step:268/1555 train_time:9220ms step_avg:34.40ms step:269/1555 train_time:9251ms 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.39ms 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:9524ms step_avg:34.38ms step:278/1555 train_time:9561ms step_avg:34.39ms step:279/1555 train_time:9592ms step_avg:34.38ms step:280/1555 train_time:9629ms step_avg:34.39ms step:281/1555 train_time:9660ms step_avg:34.38ms step:282/1555 train_time:9698ms step_avg:34.39ms step:283/1555 train_time:9729ms step_avg:34.38ms step:284/1555 train_time:9767ms step_avg:34.39ms step:285/1555 train_time:9798ms step_avg:34.38ms step:286/1555 train_time:9836ms step_avg:34.39ms step:287/1555 train_time:9867ms step_avg:34.38ms step:288/1555 train_time:9905ms step_avg:34.39ms step:289/1555 train_time:9936ms step_avg:34.38ms step:290/1555 train_time:9973ms step_avg:34.39ms step:291/1555 train_time:10004ms step_avg:34.38ms step:292/1555 train_time:10042ms step_avg:34.39ms step:293/1555 train_time:10073ms step_avg:34.38ms step:294/1555 train_time:10110ms step_avg:34.39ms step:295/1555 train_time:10141ms step_avg:34.38ms step:296/1555 train_time:10179ms step_avg:34.39ms step:297/1555 train_time:10209ms step_avg:34.38ms step:298/1555 train_time:10247ms step_avg:34.39ms step:299/1555 train_time:10278ms step_avg:34.37ms step:300/1555 train_time:10315ms step_avg:34.38ms step:301/1555 train_time:10346ms step_avg:34.37ms step:302/1555 train_time:10383ms step_avg:34.38ms step:303/1555 train_time:10414ms step_avg:34.37ms step:304/1555 train_time:10452ms step_avg:34.38ms step:305/1555 train_time:10483ms step_avg:34.37ms step:306/1555 train_time:10520ms step_avg:34.38ms step:307/1555 train_time:10551ms step_avg:34.37ms step:308/1555 train_time:10588ms step_avg:34.38ms step:309/1555 train_time:10620ms step_avg:34.37ms step:310/1555 train_time:10657ms step_avg:34.38ms step:311/1555 train_time:10689ms 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:10826ms 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:10931ms step_avg:34.37ms 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:11032ms step_avg:34.37ms step:322/1555 train_time:11069ms 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.37ms step:327/1555 train_time:11237ms step_avg:34.36ms step:328/1555 train_time:11274ms 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:11548ms step_avg:34.37ms step:337/1555 train_time:11580ms step_avg:34.36ms step:338/1555 train_time:11617ms step_avg:34.37ms step:339/1555 train_time:11649ms 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:11823ms 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:11960ms step_avg:34.37ms step:349/1555 train_time:11991ms step_avg:34.36ms step:350/1555 train_time:12029ms step_avg:34.37ms step:351/1555 train_time:12060ms step_avg:34.36ms step:352/1555 train_time:12097ms step_avg:34.37ms step:353/1555 train_time:12128ms step_avg:34.36ms step:354/1555 train_time:12166ms step_avg:34.37ms step:355/1555 train_time:12197ms step_avg:34.36ms step:356/1555 train_time:12235ms step_avg:34.37ms step:357/1555 train_time:12266ms step_avg:34.36ms step:358/1555 train_time:12304ms step_avg:34.37ms step:359/1555 train_time:12334ms step_avg:34.36ms step:360/1555 train_time:12371ms step_avg:34.37ms step:361/1555 train_time:12403ms step_avg:34.36ms step:362/1555 train_time:12440ms step_avg:34.36ms step:363/1555 train_time:12471ms step_avg:34.36ms step:364/1555 train_time:12508ms step_avg:34.36ms step:365/1555 train_time:12539ms step_avg:34.35ms step:366/1555 train_time:12577ms step_avg:34.36ms step:367/1555 train_time:12608ms step_avg:34.35ms step:368/1555 train_time:12645ms step_avg:34.36ms step:369/1555 train_time:12676ms step_avg:34.35ms step:370/1555 train_time:12713ms step_avg:34.36ms step:371/1555 train_time:12744ms step_avg:34.35ms step:372/1555 train_time:12782ms step_avg:34.36ms step:373/1555 train_time:12813ms step_avg:34.35ms step:374/1555 train_time:12850ms step_avg:34.36ms step:375/1555 train_time:12881ms 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:13055ms step_avg:34.35ms step:381/1555 train_time:13086ms step_avg:34.35ms step:382/1555 train_time:13124ms step_avg:34.36ms step:383/1555 train_time:13155ms step_avg:34.35ms step:384/1555 train_time:13192ms step_avg:34.35ms step:385/1555 train_time:13223ms step_avg:34.35ms step:386/1555 train_time:13261ms step_avg:34.36ms step:387/1555 train_time:13292ms step_avg:34.35ms step:388/1555 train_time:13330ms step_avg:34.35ms step:389/1555 train_time:13361ms step_avg:34.35ms step:390/1555 train_time:13398ms step_avg:34.35ms step:391/1555 train_time:13429ms 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:13633ms step_avg:34.34ms step:398/1555 train_time:13670ms step_avg:34.35ms step:399/1555 train_time:13702ms 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:13876ms step_avg:34.35ms step:405/1555 train_time:13907ms step_avg:34.34ms step:406/1555 train_time:13944ms step_avg:34.35ms step:407/1555 train_time:13975ms step_avg:34.34ms step:408/1555 train_time:14012ms step_avg:34.34ms step:409/1555 train_time:14043ms 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:14149ms step_avg:34.34ms step:413/1555 train_time:14180ms step_avg:34.33ms step:414/1555 train_time:14218ms step_avg:34.34ms step:415/1555 train_time:14249ms step_avg:34.33ms step:416/1555 train_time:14287ms step_avg:34.34ms step:417/1555 train_time:14318ms step_avg:34.34ms step:418/1555 train_time:14355ms step_avg:34.34ms step:419/1555 train_time:14386ms step_avg:34.33ms step:420/1555 train_time:14424ms step_avg:34.34ms step:421/1555 train_time:14454ms step_avg:34.33ms step:422/1555 train_time:14492ms step_avg:34.34ms step:423/1555 train_time:14523ms step_avg:34.33ms step:424/1555 train_time:14560ms step_avg:34.34ms step:425/1555 train_time:14591ms step_avg:34.33ms step:426/1555 train_time:14628ms step_avg:34.34ms step:427/1555 train_time:14659ms step_avg:34.33ms step:428/1555 train_time:14697ms step_avg:34.34ms step:429/1555 train_time:14728ms step_avg:34.33ms step:430/1555 train_time:14765ms step_avg:34.34ms step:431/1555 train_time:14796ms step_avg:34.33ms step:432/1555 train_time:14833ms step_avg:34.34ms step:433/1555 train_time:14864ms step_avg:34.33ms step:434/1555 train_time:14901ms step_avg:34.34ms step:435/1555 train_time:14932ms step_avg:34.33ms step:436/1555 train_time:14970ms step_avg:34.33ms step:437/1555 train_time:15001ms step_avg:34.33ms step:438/1555 train_time:15039ms step_avg:34.34ms step:439/1555 train_time:15070ms step_avg:34.33ms step:440/1555 train_time:15107ms step_avg:34.33ms step:441/1555 train_time:15138ms step_avg:34.33ms step:442/1555 train_time:15176ms step_avg:34.33ms step:443/1555 train_time:15207ms step_avg:34.33ms step:444/1555 train_time:15245ms step_avg:34.34ms step:445/1555 train_time:15276ms step_avg:34.33ms step:446/1555 train_time:15313ms step_avg:34.33ms step:447/1555 train_time:15344ms step_avg:34.33ms step:448/1555 train_time:15382ms step_avg:34.33ms step:449/1555 train_time:15413ms step_avg:34.33ms step:450/1555 train_time:15450ms step_avg:34.33ms step:451/1555 train_time:15481ms step_avg:34.33ms step:452/1555 train_time:15518ms step_avg:34.33ms step:453/1555 train_time:15549ms step_avg:34.32ms step:454/1555 train_time:15586ms step_avg:34.33ms step:455/1555 train_time:15618ms step_avg:34.32ms step:456/1555 train_time:15655ms step_avg:34.33ms step:457/1555 train_time:15686ms step_avg:34.32ms step:458/1555 train_time:15724ms step_avg:34.33ms step:459/1555 train_time:15754ms step_avg:34.32ms step:460/1555 train_time:15792ms step_avg:34.33ms step:461/1555 train_time:15822ms step_avg:34.32ms step:462/1555 train_time:15860ms step_avg:34.33ms step:463/1555 train_time:15891ms step_avg:34.32ms step:464/1555 train_time:15928ms step_avg:34.33ms step:465/1555 train_time:15959ms step_avg:34.32ms step:466/1555 train_time:15996ms step_avg:34.33ms step:467/1555 train_time:16027ms step_avg:34.32ms step:468/1555 train_time:16064ms step_avg:34.33ms step:469/1555 train_time:16095ms step_avg:34.32ms step:470/1555 train_time:16133ms step_avg:34.32ms step:471/1555 train_time:16163ms step_avg:34.32ms step:472/1555 train_time:16201ms step_avg:34.32ms step:473/1555 train_time:16232ms step_avg:34.32ms step:474/1555 train_time:16269ms step_avg:34.32ms step:475/1555 train_time:16300ms step_avg:34.32ms step:476/1555 train_time:16338ms step_avg:34.32ms step:477/1555 train_time:16369ms step_avg:34.32ms step:478/1555 train_time:16406ms step_avg:34.32ms step:479/1555 train_time:16438ms step_avg:34.32ms step:480/1555 train_time:16475ms step_avg:34.32ms step:481/1555 train_time:16506ms step_avg:34.32ms step:482/1555 train_time:16543ms step_avg:34.32ms step:483/1555 train_time:16574ms step_avg:34.32ms step:484/1555 train_time:16612ms step_avg:34.32ms step:485/1555 train_time:16643ms step_avg:34.31ms step:486/1555 train_time:16680ms step_avg:34.32ms step:487/1555 train_time:16711ms step_avg:34.31ms step:488/1555 train_time:16748ms step_avg:34.32ms step:489/1555 train_time:16779ms step_avg:34.31ms step:490/1555 train_time:16817ms step_avg:34.32ms step:491/1555 train_time:16848ms step_avg:34.31ms step:492/1555 train_time:16886ms step_avg:34.32ms step:493/1555 train_time:16916ms step_avg:34.31ms step:494/1555 train_time:16953ms step_avg:34.32ms step:495/1555 train_time:16984ms step_avg:34.31ms step:496/1555 train_time:17022ms step_avg:34.32ms step:497/1555 train_time:17052ms step_avg:34.31ms step:498/1555 train_time:17090ms step_avg:34.32ms step:499/1555 train_time:17121ms step_avg:34.31ms step:500/1555 train_time:17159ms step_avg:34.32ms step:500/1555 val_loss:4.2305 train_time:17208ms step_avg:34.42ms step:501/1555 train_time:17227ms step_avg:34.39ms step:502/1555 train_time:17247ms step_avg:34.36ms step:503/1555 train_time:17265ms step_avg:34.32ms step:504/1555 train_time:17297ms step_avg:34.32ms step:505/1555 train_time:17329ms step_avg:34.32ms step:506/1555 train_time:17372ms step_avg:34.33ms step:507/1555 train_time:17427ms step_avg:34.37ms step:508/1555 train_time:17491ms step_avg:34.43ms step:509/1555 train_time:17549ms step_avg:34.48ms step:510/1555 train_time:17612ms step_avg:34.53ms step:511/1555 train_time:17670ms step_avg:34.58ms step:512/1555 train_time:17733ms step_avg:34.63ms step:513/1555 train_time:17790ms step_avg:34.68ms step:514/1555 train_time:17853ms step_avg:34.73ms step:515/1555 train_time:17910ms step_avg:34.78ms step:516/1555 train_time:17972ms step_avg:34.83ms step:517/1555 train_time:18029ms step_avg:34.87ms step:518/1555 train_time:18093ms step_avg:34.93ms step:519/1555 train_time:18151ms step_avg:34.97ms step:520/1555 train_time:18217ms step_avg:35.03ms step:521/1555 train_time:18276ms step_avg:35.08ms step:522/1555 train_time:18342ms step_avg:35.14ms step:523/1555 train_time:18399ms step_avg:35.18ms step:524/1555 train_time:18464ms step_avg:35.24ms step:525/1555 train_time:18523ms step_avg:35.28ms step:526/1555 train_time:18588ms step_avg:35.34ms step:527/1555 train_time:18646ms step_avg:35.38ms step:528/1555 train_time:18709ms step_avg:35.43ms step:529/1555 train_time:18766ms step_avg:35.48ms step:530/1555 train_time:18830ms step_avg:35.53ms step:531/1555 train_time:18888ms step_avg:35.57ms step:532/1555 train_time:18951ms step_avg:35.62ms step:533/1555 train_time:19008ms step_avg:35.66ms step:534/1555 train_time:19071ms step_avg:35.71ms step:535/1555 train_time:19130ms step_avg:35.76ms step:536/1555 train_time:19195ms step_avg:35.81ms step:537/1555 train_time:19254ms step_avg:35.85ms step:538/1555 train_time:19318ms step_avg:35.91ms step:539/1555 train_time:19375ms step_avg:35.95ms step:540/1555 train_time:19439ms step_avg:36.00ms step:541/1555 train_time:19497ms step_avg:36.04ms step:542/1555 train_time:19561ms step_avg:36.09ms step:543/1555 train_time:19619ms step_avg:36.13ms step:544/1555 train_time:19685ms step_avg:36.19ms step:545/1555 train_time:19744ms step_avg:36.23ms step:546/1555 train_time:19808ms step_avg:36.28ms step:547/1555 train_time:19866ms step_avg:36.32ms step:548/1555 train_time:19931ms step_avg:36.37ms step:549/1555 train_time:19988ms step_avg:36.41ms step:550/1555 train_time:20051ms step_avg:36.46ms step:551/1555 train_time:20109ms step_avg:36.50ms step:552/1555 train_time:20174ms step_avg:36.55ms step:553/1555 train_time:20232ms step_avg:36.59ms step:554/1555 train_time:20296ms step_avg:36.64ms step:555/1555 train_time:20353ms step_avg:36.67ms step:556/1555 train_time:20418ms step_avg:36.72ms step:557/1555 train_time:20475ms step_avg:36.76ms step:558/1555 train_time:20539ms step_avg:36.81ms step:559/1555 train_time:20597ms step_avg:36.85ms step:560/1555 train_time:20661ms step_avg:36.89ms step:561/1555 train_time:20719ms step_avg:36.93ms step:562/1555 train_time:20783ms step_avg:36.98ms step:563/1555 train_time:20842ms step_avg:37.02ms step:564/1555 train_time:20906ms step_avg:37.07ms step:565/1555 train_time:20965ms step_avg:37.11ms step:566/1555 train_time:21030ms step_avg:37.16ms step:567/1555 train_time:21088ms step_avg:37.19ms step:568/1555 train_time:21152ms step_avg:37.24ms step:569/1555 train_time:21209ms step_avg:37.27ms step:570/1555 train_time:21273ms step_avg:37.32ms step:571/1555 train_time:21332ms step_avg:37.36ms step:572/1555 train_time:21396ms step_avg:37.41ms step:573/1555 train_time:21453ms step_avg:37.44ms step:574/1555 train_time:21517ms step_avg:37.49ms step:575/1555 train_time:21574ms step_avg:37.52ms step:576/1555 train_time:21639ms step_avg:37.57ms step:577/1555 train_time:21696ms step_avg:37.60ms step:578/1555 train_time:21760ms step_avg:37.65ms step:579/1555 train_time:21817ms step_avg:37.68ms step:580/1555 train_time:21882ms step_avg:37.73ms step:581/1555 train_time:21941ms step_avg:37.76ms step:582/1555 train_time:22006ms step_avg:37.81ms step:583/1555 train_time:22063ms step_avg:37.84ms step:584/1555 train_time:22128ms step_avg:37.89ms step:585/1555 train_time:22186ms step_avg:37.92ms step:586/1555 train_time:22252ms step_avg:37.97ms step:587/1555 train_time:22309ms step_avg:38.00ms step:588/1555 train_time:22373ms step_avg:38.05ms step:589/1555 train_time:22431ms step_avg:38.08ms step:590/1555 train_time:22495ms step_avg:38.13ms step:591/1555 train_time:22552ms step_avg:38.16ms step:592/1555 train_time:22616ms step_avg:38.20ms step:593/1555 train_time:22673ms step_avg:38.24ms step:594/1555 train_time:22737ms step_avg:38.28ms step:595/1555 train_time:22795ms step_avg:38.31ms step:596/1555 train_time:22859ms step_avg:38.35ms step:597/1555 train_time:22916ms step_avg:38.38ms step:598/1555 train_time:22980ms step_avg:38.43ms step:599/1555 train_time:23037ms step_avg:38.46ms step:600/1555 train_time:23102ms step_avg:38.50ms step:601/1555 train_time:23159ms step_avg:38.53ms step:602/1555 train_time:23225ms step_avg:38.58ms step:603/1555 train_time:23284ms step_avg:38.61ms step:604/1555 train_time:23348ms step_avg:38.66ms step:605/1555 train_time:23407ms step_avg:38.69ms step:606/1555 train_time:23471ms step_avg:38.73ms step:607/1555 train_time:23530ms step_avg:38.76ms step:608/1555 train_time:23595ms step_avg:38.81ms step:609/1555 train_time:23652ms step_avg:38.84ms step:610/1555 train_time:23715ms step_avg:38.88ms step:611/1555 train_time:23774ms step_avg:38.91ms step:612/1555 train_time:23837ms step_avg:38.95ms step:613/1555 train_time:23895ms step_avg:38.98ms step:614/1555 train_time:23957ms step_avg:39.02ms step:615/1555 train_time:24015ms step_avg:39.05ms step:616/1555 train_time:24080ms step_avg:39.09ms step:617/1555 train_time:24138ms step_avg:39.12ms step:618/1555 train_time:24203ms step_avg:39.16ms step:619/1555 train_time:24260ms step_avg:39.19ms step:620/1555 train_time:24325ms step_avg:39.23ms step:621/1555 train_time:24384ms step_avg:39.27ms step:622/1555 train_time:24449ms step_avg:39.31ms step:623/1555 train_time:24507ms step_avg:39.34ms step:624/1555 train_time:24572ms step_avg:39.38ms step:625/1555 train_time:24630ms step_avg:39.41ms step:626/1555 train_time:24694ms step_avg:39.45ms step:627/1555 train_time:24752ms step_avg:39.48ms step:628/1555 train_time:24816ms step_avg:39.52ms step:629/1555 train_time:24873ms step_avg:39.54ms step:630/1555 train_time:24937ms step_avg:39.58ms step:631/1555 train_time:24994ms step_avg:39.61ms step:632/1555 train_time:25057ms step_avg:39.65ms step:633/1555 train_time:25114ms step_avg:39.67ms step:634/1555 train_time:25179ms step_avg:39.71ms step:635/1555 train_time:25236ms step_avg:39.74ms step:636/1555 train_time:25301ms step_avg:39.78ms step:637/1555 train_time:25359ms step_avg:39.81ms step:638/1555 train_time:25423ms step_avg:39.85ms step:639/1555 train_time:25482ms step_avg:39.88ms step:640/1555 train_time:25547ms step_avg:39.92ms step:641/1555 train_time:25606ms step_avg:39.95ms step:642/1555 train_time:25670ms step_avg:39.98ms step:643/1555 train_time:25728ms step_avg:40.01ms step:644/1555 train_time:25793ms step_avg:40.05ms step:645/1555 train_time:25850ms step_avg:40.08ms step:646/1555 train_time:25914ms step_avg:40.11ms step:647/1555 train_time:25972ms step_avg:40.14ms step:648/1555 train_time:26036ms step_avg:40.18ms step:649/1555 train_time:26093ms step_avg:40.21ms step:650/1555 train_time:26157ms step_avg:40.24ms step:651/1555 train_time:26213ms step_avg:40.27ms step:652/1555 train_time:26278ms step_avg:40.30ms step:653/1555 train_time:26335ms step_avg:40.33ms step:654/1555 train_time:26399ms step_avg:40.37ms step:655/1555 train_time:26456ms step_avg:40.39ms step:656/1555 train_time:26521ms step_avg:40.43ms step:657/1555 train_time:26578ms step_avg:40.45ms step:658/1555 train_time:26643ms step_avg:40.49ms step:659/1555 train_time:26702ms step_avg:40.52ms step:660/1555 train_time:26767ms step_avg:40.56ms step:661/1555 train_time:26826ms step_avg:40.58ms step:662/1555 train_time:26891ms step_avg:40.62ms step:663/1555 train_time:26948ms step_avg:40.65ms step:664/1555 train_time:27013ms step_avg:40.68ms step:665/1555 train_time:27071ms step_avg:40.71ms step:666/1555 train_time:27135ms step_avg:40.74ms step:667/1555 train_time:27193ms step_avg:40.77ms step:668/1555 train_time:27256ms step_avg:40.80ms step:669/1555 train_time:27313ms step_avg:40.83ms step:670/1555 train_time:27376ms step_avg:40.86ms step:671/1555 train_time:27434ms step_avg:40.88ms step:672/1555 train_time:27498ms step_avg:40.92ms step:673/1555 train_time:27555ms step_avg:40.94ms step:674/1555 train_time:27619ms step_avg:40.98ms step:675/1555 train_time:27676ms step_avg:41.00ms step:676/1555 train_time:27742ms step_avg:41.04ms step:677/1555 train_time:27799ms step_avg:41.06ms step:678/1555 train_time:27865ms step_avg:41.10ms step:679/1555 train_time:27924ms step_avg:41.12ms step:680/1555 train_time:27989ms step_avg:41.16ms step:681/1555 train_time:28047ms step_avg:41.19ms step:682/1555 train_time:28112ms step_avg:41.22ms step:683/1555 train_time:28170ms step_avg:41.24ms step:684/1555 train_time:28234ms step_avg:41.28ms step:685/1555 train_time:28291ms step_avg:41.30ms step:686/1555 train_time:28354ms step_avg:41.33ms step:687/1555 train_time:28412ms step_avg:41.36ms step:688/1555 train_time:28475ms step_avg:41.39ms step:689/1555 train_time:28532ms step_avg:41.41ms step:690/1555 train_time:28596ms step_avg:41.44ms step:691/1555 train_time:28654ms step_avg:41.47ms step:692/1555 train_time:28718ms step_avg:41.50ms step:693/1555 train_time:28775ms step_avg:41.52ms step:694/1555 train_time:28840ms step_avg:41.56ms step:695/1555 train_time:28898ms step_avg:41.58ms step:696/1555 train_time:28962ms step_avg:41.61ms step:697/1555 train_time:29020ms step_avg:41.64ms step:698/1555 train_time:29085ms step_avg:41.67ms step:699/1555 train_time:29144ms step_avg:41.69ms step:700/1555 train_time:29209ms step_avg:41.73ms step:701/1555 train_time:29267ms step_avg:41.75ms step:702/1555 train_time:29331ms step_avg:41.78ms step:703/1555 train_time:29389ms step_avg:41.81ms step:704/1555 train_time:29453ms step_avg:41.84ms step:705/1555 train_time:29511ms step_avg:41.86ms step:706/1555 train_time:29575ms step_avg:41.89ms step:707/1555 train_time:29633ms step_avg:41.91ms step:708/1555 train_time:29696ms step_avg:41.94ms step:709/1555 train_time:29753ms step_avg:41.97ms step:710/1555 train_time:29817ms step_avg:42.00ms step:711/1555 train_time:29875ms step_avg:42.02ms step:712/1555 train_time:29939ms step_avg:42.05ms step:713/1555 train_time:29997ms step_avg:42.07ms step:714/1555 train_time:30062ms step_avg:42.10ms step:715/1555 train_time:30119ms step_avg:42.12ms step:716/1555 train_time:30184ms step_avg:42.16ms step:717/1555 train_time:30243ms step_avg:42.18ms step:718/1555 train_time:30307ms step_avg:42.21ms step:719/1555 train_time:30365ms step_avg:42.23ms step:720/1555 train_time:30430ms step_avg:42.26ms step:721/1555 train_time:30489ms step_avg:42.29ms step:722/1555 train_time:30552ms step_avg:42.32ms step:723/1555 train_time:30610ms step_avg:42.34ms step:724/1555 train_time:30674ms step_avg:42.37ms step:725/1555 train_time:30732ms step_avg:42.39ms step:726/1555 train_time:30796ms step_avg:42.42ms step:727/1555 train_time:30852ms step_avg:42.44ms step:728/1555 train_time:30916ms step_avg:42.47ms step:729/1555 train_time:30974ms step_avg:42.49ms step:730/1555 train_time:31038ms step_avg:42.52ms step:731/1555 train_time:31094ms step_avg:42.54ms step:732/1555 train_time:31159ms step_avg:42.57ms step:733/1555 train_time:31216ms step_avg:42.59ms step:734/1555 train_time:31281ms step_avg:42.62ms step:735/1555 train_time:31338ms step_avg:42.64ms step:736/1555 train_time:31403ms step_avg:42.67ms step:737/1555 train_time:31462ms step_avg:42.69ms step:738/1555 train_time:31528ms step_avg:42.72ms step:739/1555 train_time:31585ms step_avg:42.74ms step:740/1555 train_time:31650ms step_avg:42.77ms step:741/1555 train_time:31708ms step_avg:42.79ms step:742/1555 train_time:31772ms step_avg:42.82ms step:743/1555 train_time:31831ms step_avg:42.84ms step:744/1555 train_time:31894ms step_avg:42.87ms step:745/1555 train_time:31952ms step_avg:42.89ms step:746/1555 train_time:32016ms step_avg:42.92ms step:747/1555 train_time:32074ms step_avg:42.94ms step:748/1555 train_time:32138ms step_avg:42.97ms step:749/1555 train_time:32195ms step_avg:42.98ms step:750/1555 train_time:32260ms step_avg:43.01ms step:750/1555 val_loss:3.8780 train_time:32341ms step_avg:43.12ms step:751/1555 train_time:32361ms step_avg:43.09ms step:752/1555 train_time:32381ms step_avg:43.06ms step:753/1555 train_time:32438ms step_avg:43.08ms step:754/1555 train_time:32506ms step_avg:43.11ms step:755/1555 train_time:32563ms step_avg:43.13ms step:756/1555 train_time:32627ms step_avg:43.16ms step:757/1555 train_time:32685ms step_avg:43.18ms step:758/1555 train_time:32749ms step_avg:43.20ms step:759/1555 train_time:32806ms step_avg:43.22ms step:760/1555 train_time:32870ms step_avg:43.25ms step:761/1555 train_time:32926ms step_avg:43.27ms step:762/1555 train_time:32991ms step_avg:43.30ms step:763/1555 train_time:33048ms step_avg:43.31ms step:764/1555 train_time:33111ms step_avg:43.34ms step:765/1555 train_time:33169ms step_avg:43.36ms step:766/1555 train_time:33234ms step_avg:43.39ms step:767/1555 train_time:33293ms step_avg:43.41ms step:768/1555 train_time:33358ms step_avg:43.43ms step:769/1555 train_time:33417ms step_avg:43.45ms step:770/1555 train_time:33482ms step_avg:43.48ms step:771/1555 train_time:33540ms step_avg:43.50ms step:772/1555 train_time:33605ms step_avg:43.53ms step:773/1555 train_time:33662ms step_avg:43.55ms step:774/1555 train_time:33725ms step_avg:43.57ms step:775/1555 train_time:33782ms step_avg:43.59ms step:776/1555 train_time:33845ms step_avg:43.61ms step:777/1555 train_time:33902ms step_avg:43.63ms step:778/1555 train_time:33965ms step_avg:43.66ms step:779/1555 train_time:34022ms step_avg:43.67ms step:780/1555 train_time:34086ms step_avg:43.70ms step:781/1555 train_time:34142ms step_avg:43.72ms step:782/1555 train_time:34207ms step_avg:43.74ms step:783/1555 train_time:34264ms step_avg:43.76ms step:784/1555 train_time:34330ms step_avg:43.79ms step:785/1555 train_time:34389ms step_avg:43.81ms step:786/1555 train_time:34455ms step_avg:43.84ms step:787/1555 train_time:34514ms step_avg:43.85ms step:788/1555 train_time:34578ms step_avg:43.88ms step:789/1555 train_time:34636ms step_avg:43.90ms step:790/1555 train_time:34701ms step_avg:43.92ms step:791/1555 train_time:34758ms step_avg:43.94ms step:792/1555 train_time:34822ms step_avg:43.97ms step:793/1555 train_time:34879ms step_avg:43.98ms step:794/1555 train_time:34942ms step_avg:44.01ms step:795/1555 train_time:35000ms step_avg:44.03ms step:796/1555 train_time:35063ms step_avg:44.05ms step:797/1555 train_time:35121ms step_avg:44.07ms step:798/1555 train_time:35185ms step_avg:44.09ms step:799/1555 train_time:35243ms step_avg:44.11ms step:800/1555 train_time:35307ms step_avg:44.13ms step:801/1555 train_time:35363ms step_avg:44.15ms step:802/1555 train_time:35428ms step_avg:44.17ms step:803/1555 train_time:35486ms step_avg:44.19ms step:804/1555 train_time:35551ms step_avg:44.22ms step:805/1555 train_time:35610ms step_avg:44.24ms step:806/1555 train_time:35675ms step_avg:44.26ms step:807/1555 train_time:35734ms step_avg:44.28ms step:808/1555 train_time:35798ms step_avg:44.30ms step:809/1555 train_time:35855ms step_avg:44.32ms step:810/1555 train_time:35920ms step_avg:44.35ms step:811/1555 train_time:35977ms step_avg:44.36ms step:812/1555 train_time:36041ms step_avg:44.39ms step:813/1555 train_time:36100ms step_avg:44.40ms step:814/1555 train_time:36162ms step_avg:44.43ms step:815/1555 train_time:36219ms step_avg:44.44ms step:816/1555 train_time:36284ms step_avg:44.47ms step:817/1555 train_time:36341ms step_avg:44.48ms step:818/1555 train_time:36405ms step_avg:44.51ms step:819/1555 train_time:36462ms step_avg:44.52ms step:820/1555 train_time:36527ms step_avg:44.55ms step:821/1555 train_time:36584ms step_avg:44.56ms step:822/1555 train_time:36649ms step_avg:44.59ms step:823/1555 train_time:36707ms step_avg:44.60ms step:824/1555 train_time:36771ms step_avg:44.63ms step:825/1555 train_time:36830ms step_avg:44.64ms step:826/1555 train_time:36895ms step_avg:44.67ms step:827/1555 train_time:36953ms step_avg:44.68ms step:828/1555 train_time:37017ms step_avg:44.71ms step:829/1555 train_time:37075ms step_avg:44.72ms step:830/1555 train_time:37139ms step_avg:44.75ms step:831/1555 train_time:37196ms step_avg:44.76ms step:832/1555 train_time:37261ms step_avg:44.78ms step:833/1555 train_time:37319ms step_avg:44.80ms step:834/1555 train_time:37383ms step_avg:44.82ms step:835/1555 train_time:37441ms step_avg:44.84ms step:836/1555 train_time:37505ms step_avg:44.86ms step:837/1555 train_time:37562ms step_avg:44.88ms step:838/1555 train_time:37625ms step_avg:44.90ms step:839/1555 train_time:37683ms step_avg:44.91ms step:840/1555 train_time:37747ms step_avg:44.94ms step:841/1555 train_time:37804ms step_avg:44.95ms step:842/1555 train_time:37868ms step_avg:44.97ms step:843/1555 train_time:37926ms step_avg:44.99ms step:844/1555 train_time:37991ms step_avg:45.01ms step:845/1555 train_time:38049ms step_avg:45.03ms step:846/1555 train_time:38114ms step_avg:45.05ms step:847/1555 train_time:38172ms step_avg:45.07ms step:848/1555 train_time:38236ms step_avg:45.09ms step:849/1555 train_time:38295ms step_avg:45.11ms step:850/1555 train_time:38360ms step_avg:45.13ms step:851/1555 train_time:38418ms step_avg:45.14ms step:852/1555 train_time:38482ms step_avg:45.17ms step:853/1555 train_time:38539ms step_avg:45.18ms step:854/1555 train_time:38604ms step_avg:45.20ms step:855/1555 train_time:38662ms step_avg:45.22ms step:856/1555 train_time:38725ms step_avg:45.24ms step:857/1555 train_time:38783ms step_avg:45.25ms step:858/1555 train_time:38847ms step_avg:45.28ms step:859/1555 train_time:38904ms step_avg:45.29ms step:860/1555 train_time:38968ms step_avg:45.31ms step:861/1555 train_time:39025ms step_avg:45.33ms step:862/1555 train_time:39089ms step_avg:45.35ms step:863/1555 train_time:39147ms step_avg:45.36ms step:864/1555 train_time:39212ms step_avg:45.38ms step:865/1555 train_time:39270ms step_avg:45.40ms step:866/1555 train_time:39336ms step_avg:45.42ms step:867/1555 train_time:39393ms step_avg:45.44ms step:868/1555 train_time:39458ms step_avg:45.46ms step:869/1555 train_time:39516ms step_avg:45.47ms step:870/1555 train_time:39580ms step_avg:45.49ms step:871/1555 train_time:39639ms step_avg:45.51ms step:872/1555 train_time:39703ms step_avg:45.53ms step:873/1555 train_time:39760ms step_avg:45.54ms step:874/1555 train_time:39823ms step_avg:45.56ms step:875/1555 train_time:39880ms step_avg:45.58ms step:876/1555 train_time:39945ms step_avg:45.60ms step:877/1555 train_time:40001ms step_avg:45.61ms step:878/1555 train_time:40065ms step_avg:45.63ms step:879/1555 train_time:40122ms step_avg:45.65ms step:880/1555 train_time:40187ms step_avg:45.67ms step:881/1555 train_time:40245ms step_avg:45.68ms step:882/1555 train_time:40310ms step_avg:45.70ms step:883/1555 train_time:40367ms step_avg:45.72ms step:884/1555 train_time:40432ms step_avg:45.74ms step:885/1555 train_time:40489ms step_avg:45.75ms step:886/1555 train_time:40555ms step_avg:45.77ms step:887/1555 train_time:40613ms step_avg:45.79ms step:888/1555 train_time:40678ms step_avg:45.81ms step:889/1555 train_time:40735ms step_avg:45.82ms step:890/1555 train_time:40800ms step_avg:45.84ms step:891/1555 train_time:40858ms step_avg:45.86ms step:892/1555 train_time:40921ms step_avg:45.88ms step:893/1555 train_time:40979ms step_avg:45.89ms step:894/1555 train_time:41044ms step_avg:45.91ms step:895/1555 train_time:41101ms step_avg:45.92ms step:896/1555 train_time:41164ms step_avg:45.94ms step:897/1555 train_time:41222ms step_avg:45.95ms step:898/1555 train_time:41287ms step_avg:45.98ms step:899/1555 train_time:41343ms step_avg:45.99ms step:900/1555 train_time:41407ms step_avg:46.01ms step:901/1555 train_time:41463ms step_avg:46.02ms step:902/1555 train_time:41529ms step_avg:46.04ms step:903/1555 train_time:41586ms step_avg:46.05ms step:904/1555 train_time:41650ms step_avg:46.07ms step:905/1555 train_time:41709ms step_avg:46.09ms step:906/1555 train_time:41773ms step_avg:46.11ms step:907/1555 train_time:41831ms step_avg:46.12ms step:908/1555 train_time:41896ms step_avg:46.14ms step:909/1555 train_time:41954ms step_avg:46.15ms step:910/1555 train_time:42018ms step_avg:46.17ms step:911/1555 train_time:42076ms step_avg:46.19ms step:912/1555 train_time:42140ms step_avg:46.21ms step:913/1555 train_time:42199ms step_avg:46.22ms step:914/1555 train_time:42262ms step_avg:46.24ms step:915/1555 train_time:42320ms step_avg:46.25ms step:916/1555 train_time:42384ms step_avg:46.27ms step:917/1555 train_time:42441ms step_avg:46.28ms step:918/1555 train_time:42505ms step_avg:46.30ms step:919/1555 train_time:42563ms step_avg:46.31ms step:920/1555 train_time:42627ms step_avg:46.33ms step:921/1555 train_time:42684ms step_avg:46.35ms step:922/1555 train_time:42749ms step_avg:46.37ms step:923/1555 train_time:42807ms step_avg:46.38ms step:924/1555 train_time:42872ms step_avg:46.40ms step:925/1555 train_time:42931ms step_avg:46.41ms step:926/1555 train_time:42995ms step_avg:46.43ms step:927/1555 train_time:43053ms step_avg:46.44ms step:928/1555 train_time:43118ms step_avg:46.46ms step:929/1555 train_time:43176ms step_avg:46.48ms step:930/1555 train_time:43241ms step_avg:46.50ms step:931/1555 train_time:43299ms step_avg:46.51ms step:932/1555 train_time:43363ms step_avg:46.53ms step:933/1555 train_time:43421ms step_avg:46.54ms step:934/1555 train_time:43484ms step_avg:46.56ms step:935/1555 train_time:43542ms step_avg:46.57ms step:936/1555 train_time:43605ms step_avg:46.59ms step:937/1555 train_time:43663ms step_avg:46.60ms step:938/1555 train_time:43727ms step_avg:46.62ms step:939/1555 train_time:43783ms step_avg:46.63ms step:940/1555 train_time:43847ms step_avg:46.65ms step:941/1555 train_time:43905ms step_avg:46.66ms step:942/1555 train_time:43969ms step_avg:46.68ms step:943/1555 train_time:44027ms step_avg:46.69ms step:944/1555 train_time:44092ms step_avg:46.71ms step:945/1555 train_time:44151ms step_avg:46.72ms step:946/1555 train_time:44216ms step_avg:46.74ms step:947/1555 train_time:44273ms step_avg:46.75ms step:948/1555 train_time:44338ms step_avg:46.77ms step:949/1555 train_time:44396ms step_avg:46.78ms step:950/1555 train_time:44460ms step_avg:46.80ms step:951/1555 train_time:44518ms step_avg:46.81ms step:952/1555 train_time:44581ms step_avg:46.83ms step:953/1555 train_time:44640ms step_avg:46.84ms step:954/1555 train_time:44704ms step_avg:46.86ms step:955/1555 train_time:44761ms step_avg:46.87ms step:956/1555 train_time:44825ms step_avg:46.89ms step:957/1555 train_time:44882ms step_avg:46.90ms step:958/1555 train_time:44946ms step_avg:46.92ms step:959/1555 train_time:45004ms step_avg:46.93ms step:960/1555 train_time:45069ms step_avg:46.95ms step:961/1555 train_time:45126ms step_avg:46.96ms step:962/1555 train_time:45190ms step_avg:46.97ms step:963/1555 train_time:45248ms step_avg:46.99ms step:964/1555 train_time:45313ms step_avg:47.01ms step:965/1555 train_time:45371ms step_avg:47.02ms step:966/1555 train_time:45436ms step_avg:47.03ms step:967/1555 train_time:45495ms step_avg:47.05ms step:968/1555 train_time:45559ms step_avg:47.06ms step:969/1555 train_time:45617ms step_avg:47.08ms step:970/1555 train_time:45681ms step_avg:47.09ms step:971/1555 train_time:45738ms step_avg:47.10ms step:972/1555 train_time:45802ms step_avg:47.12ms step:973/1555 train_time:45860ms step_avg:47.13ms step:974/1555 train_time:45923ms step_avg:47.15ms step:975/1555 train_time:45981ms step_avg:47.16ms step:976/1555 train_time:46045ms step_avg:47.18ms step:977/1555 train_time:46102ms step_avg:47.19ms step:978/1555 train_time:46166ms step_avg:47.20ms step:979/1555 train_time:46223ms step_avg:47.21ms step:980/1555 train_time:46288ms step_avg:47.23ms step:981/1555 train_time:46344ms step_avg:47.24ms step:982/1555 train_time:46409ms step_avg:47.26ms step:983/1555 train_time:46467ms step_avg:47.27ms step:984/1555 train_time:46532ms step_avg:47.29ms step:985/1555 train_time:46590ms step_avg:47.30ms step:986/1555 train_time:46655ms step_avg:47.32ms step:987/1555 train_time:46713ms step_avg:47.33ms step:988/1555 train_time:46778ms step_avg:47.35ms step:989/1555 train_time:46836ms step_avg:47.36ms step:990/1555 train_time:46901ms step_avg:47.38ms step:991/1555 train_time:46959ms step_avg:47.39ms step:992/1555 train_time:47022ms step_avg:47.40ms step:993/1555 train_time:47080ms step_avg:47.41ms step:994/1555 train_time:47144ms step_avg:47.43ms step:995/1555 train_time:47201ms step_avg:47.44ms step:996/1555 train_time:47264ms step_avg:47.45ms step:997/1555 train_time:47321ms step_avg:47.46ms step:998/1555 train_time:47386ms step_avg:47.48ms step:999/1555 train_time:47442ms step_avg:47.49ms step:1000/1555 train_time:47507ms step_avg:47.51ms step:1000/1555 val_loss:3.5716 train_time:47590ms step_avg:47.59ms step:1001/1555 train_time:47610ms step_avg:47.56ms step:1002/1555 train_time:47631ms step_avg:47.54ms step:1003/1555 train_time:47691ms step_avg:47.55ms step:1004/1555 train_time:47759ms step_avg:47.57ms step:1005/1555 train_time:47819ms step_avg:47.58ms step:1006/1555 train_time:47884ms step_avg:47.60ms step:1007/1555 train_time:47940ms step_avg:47.61ms step:1008/1555 train_time:48003ms step_avg:47.62ms step:1009/1555 train_time:48060ms step_avg:47.63ms step:1010/1555 train_time:48123ms step_avg:47.65ms step:1011/1555 train_time:48183ms step_avg:47.66ms step:1012/1555 train_time:48268ms step_avg:47.70ms step:1013/1555 train_time:48351ms step_avg:47.73ms step:1014/1555 train_time:48442ms step_avg:47.77ms step:1015/1555 train_time:48524ms step_avg:47.81ms step:1016/1555 train_time:48615ms step_avg:47.85ms step:1017/1555 train_time:48702ms step_avg:47.89ms step:1018/1555 train_time:48792ms step_avg:47.93ms step:1019/1555 train_time:48879ms step_avg:47.97ms step:1020/1555 train_time:48969ms step_avg:48.01ms step:1021/1555 train_time:49055ms step_avg:48.05ms step:1022/1555 train_time:49145ms step_avg:48.09ms step:1023/1555 train_time:49227ms step_avg:48.12ms step:1024/1555 train_time:49316ms step_avg:48.16ms step:1025/1555 train_time:49400ms step_avg:48.20ms step:1026/1555 train_time:49488ms step_avg:48.23ms step:1027/1555 train_time:49573ms step_avg:48.27ms step:1028/1555 train_time:49664ms step_avg:48.31ms step:1029/1555 train_time:49748ms step_avg:48.35ms step:1030/1555 train_time:49841ms step_avg:48.39ms step:1031/1555 train_time:49925ms step_avg:48.42ms step:1032/1555 train_time:50015ms step_avg:48.46ms step:1033/1555 train_time:50100ms step_avg:48.50ms step:1034/1555 train_time:50189ms step_avg:48.54ms step:1035/1555 train_time:50273ms step_avg:48.57ms step:1036/1555 train_time:50363ms step_avg:48.61ms step:1037/1555 train_time:50445ms step_avg:48.64ms step:1038/1555 train_time:50534ms step_avg:48.68ms step:1039/1555 train_time:50619ms step_avg:48.72ms step:1040/1555 train_time:50710ms step_avg:48.76ms step:1041/1555 train_time:50795ms step_avg:48.79ms step:1042/1555 train_time:50886ms step_avg:48.83ms step:1043/1555 train_time:50970ms step_avg:48.87ms step:1044/1555 train_time:51061ms step_avg:48.91ms step:1045/1555 train_time:51144ms step_avg:48.94ms step:1046/1555 train_time:51234ms step_avg:48.98ms step:1047/1555 train_time:51317ms step_avg:49.01ms step:1048/1555 train_time:51407ms step_avg:49.05ms step:1049/1555 train_time:51490ms step_avg:49.08ms step:1050/1555 train_time:51581ms step_avg:49.12ms step:1051/1555 train_time:51665ms step_avg:49.16ms step:1052/1555 train_time:51757ms step_avg:49.20ms step:1053/1555 train_time:51841ms step_avg:49.23ms step:1054/1555 train_time:51930ms step_avg:49.27ms step:1055/1555 train_time:52016ms step_avg:49.30ms step:1056/1555 train_time:52107ms step_avg:49.34ms step:1057/1555 train_time:52190ms step_avg:49.38ms step:1058/1555 train_time:52279ms step_avg:49.41ms step:1059/1555 train_time:52362ms step_avg:49.45ms step:1060/1555 train_time:52452ms step_avg:49.48ms step:1061/1555 train_time:52536ms step_avg:49.52ms step:1062/1555 train_time:52626ms step_avg:49.55ms step:1063/1555 train_time:52710ms step_avg:49.59ms step:1064/1555 train_time:52800ms step_avg:49.62ms step:1065/1555 train_time:52885ms step_avg:49.66ms step:1066/1555 train_time:52975ms step_avg:49.70ms step:1067/1555 train_time:53060ms step_avg:49.73ms step:1068/1555 train_time:53149ms step_avg:49.77ms step:1069/1555 train_time:53235ms step_avg:49.80ms step:1070/1555 train_time:53324ms step_avg:49.84ms step:1071/1555 train_time:53407ms step_avg:49.87ms step:1072/1555 train_time:53497ms step_avg:49.90ms step:1073/1555 train_time:53581ms step_avg:49.94ms step:1074/1555 train_time:53670ms step_avg:49.97ms step:1075/1555 train_time:53755ms step_avg:50.00ms step:1076/1555 train_time:53845ms step_avg:50.04ms step:1077/1555 train_time:53928ms step_avg:50.07ms step:1078/1555 train_time:54019ms step_avg:50.11ms step:1079/1555 train_time:54103ms step_avg:50.14ms step:1080/1555 train_time:54193ms step_avg:50.18ms step:1081/1555 train_time:54277ms step_avg:50.21ms step:1082/1555 train_time:54367ms step_avg:50.25ms step:1083/1555 train_time:54450ms step_avg:50.28ms step:1084/1555 train_time:54540ms step_avg:50.31ms step:1085/1555 train_time:54623ms step_avg:50.34ms step:1086/1555 train_time:54713ms step_avg:50.38ms step:1087/1555 train_time:54797ms step_avg:50.41ms step:1088/1555 train_time:54888ms step_avg:50.45ms step:1089/1555 train_time:54972ms step_avg:50.48ms step:1090/1555 train_time:55063ms step_avg:50.52ms step:1091/1555 train_time:55147ms step_avg:50.55ms step:1092/1555 train_time:55235ms step_avg:50.58ms step:1093/1555 train_time:55319ms step_avg:50.61ms step:1094/1555 train_time:55409ms step_avg:50.65ms step:1095/1555 train_time:55493ms step_avg:50.68ms step:1096/1555 train_time:55588ms step_avg:50.72ms step:1097/1555 train_time:55667ms step_avg:50.74ms step:1098/1555 train_time:55757ms step_avg:50.78ms step:1099/1555 train_time:55843ms step_avg:50.81ms step:1100/1555 train_time:55932ms step_avg:50.85ms step:1101/1555 train_time:56016ms step_avg:50.88ms step:1102/1555 train_time:56107ms step_avg:50.91ms step:1103/1555 train_time:56191ms step_avg:50.94ms step:1104/1555 train_time:56280ms step_avg:50.98ms step:1105/1555 train_time:56364ms step_avg:51.01ms step:1106/1555 train_time:56454ms step_avg:51.04ms step:1107/1555 train_time:56537ms step_avg:51.07ms step:1108/1555 train_time:56627ms step_avg:51.11ms step:1109/1555 train_time:56710ms step_avg:51.14ms step:1110/1555 train_time:56802ms step_avg:51.17ms step:1111/1555 train_time:56885ms step_avg:51.20ms step:1112/1555 train_time:56975ms step_avg:51.24ms step:1113/1555 train_time:57059ms step_avg:51.27ms step:1114/1555 train_time:57149ms step_avg:51.30ms step:1115/1555 train_time:57234ms step_avg:51.33ms step:1116/1555 train_time:57323ms step_avg:51.37ms step:1117/1555 train_time:57407ms step_avg:51.39ms step:1118/1555 train_time:57497ms step_avg:51.43ms step:1119/1555 train_time:57582ms step_avg:51.46ms step:1120/1555 train_time:57672ms step_avg:51.49ms step:1121/1555 train_time:57757ms step_avg:51.52ms step:1122/1555 train_time:57847ms step_avg:51.56ms step:1123/1555 train_time:57930ms step_avg:51.58ms step:1124/1555 train_time:58021ms step_avg:51.62ms step:1125/1555 train_time:58105ms step_avg:51.65ms step:1126/1555 train_time:58194ms step_avg:51.68ms step:1127/1555 train_time:58279ms step_avg:51.71ms step:1128/1555 train_time:58368ms step_avg:51.74ms step:1129/1555 train_time:58452ms step_avg:51.77ms step:1130/1555 train_time:58542ms step_avg:51.81ms step:1131/1555 train_time:58625ms step_avg:51.83ms step:1132/1555 train_time:58714ms step_avg:51.87ms step:1133/1555 train_time:58798ms step_avg:51.90ms step:1134/1555 train_time:58890ms step_avg:51.93ms step:1135/1555 train_time:58974ms step_avg:51.96ms step:1136/1555 train_time:59065ms step_avg:51.99ms step:1137/1555 train_time:59148ms step_avg:52.02ms step:1138/1555 train_time:59237ms step_avg:52.05ms step:1139/1555 train_time:59322ms step_avg:52.08ms step:1140/1555 train_time:59411ms step_avg:52.12ms step:1141/1555 train_time:59497ms step_avg:52.14ms step:1142/1555 train_time:59586ms step_avg:52.18ms step:1143/1555 train_time:59669ms step_avg:52.20ms step:1144/1555 train_time:59761ms step_avg:52.24ms step:1145/1555 train_time:59845ms step_avg:52.27ms step:1146/1555 train_time:59934ms step_avg:52.30ms step:1147/1555 train_time:60018ms step_avg:52.33ms step:1148/1555 train_time:60108ms step_avg:52.36ms step:1149/1555 train_time:60192ms step_avg:52.39ms step:1150/1555 train_time:60282ms step_avg:52.42ms step:1151/1555 train_time:60365ms step_avg:52.45ms step:1152/1555 train_time:60455ms step_avg:52.48ms step:1153/1555 train_time:60539ms step_avg:52.51ms step:1154/1555 train_time:60629ms step_avg:52.54ms step:1155/1555 train_time:60713ms step_avg:52.57ms step:1156/1555 train_time:60804ms step_avg:52.60ms step:1157/1555 train_time:60887ms step_avg:52.62ms step:1158/1555 train_time:60976ms step_avg:52.66ms step:1159/1555 train_time:61061ms step_avg:52.68ms step:1160/1555 train_time:61150ms step_avg:52.72ms step:1161/1555 train_time:61235ms step_avg:52.74ms step:1162/1555 train_time:61325ms step_avg:52.78ms step:1163/1555 train_time:61409ms step_avg:52.80ms step:1164/1555 train_time:61499ms step_avg:52.83ms step:1165/1555 train_time:61584ms step_avg:52.86ms step:1166/1555 train_time:61673ms step_avg:52.89ms step:1167/1555 train_time:61758ms step_avg:52.92ms step:1168/1555 train_time:61849ms step_avg:52.95ms step:1169/1555 train_time:61932ms step_avg:52.98ms step:1170/1555 train_time:62023ms step_avg:53.01ms step:1171/1555 train_time:62107ms step_avg:53.04ms step:1172/1555 train_time:62197ms step_avg:53.07ms step:1173/1555 train_time:62281ms step_avg:53.10ms step:1174/1555 train_time:62370ms step_avg:53.13ms step:1175/1555 train_time:62454ms step_avg:53.15ms step:1176/1555 train_time:62545ms step_avg:53.18ms step:1177/1555 train_time:62628ms step_avg:53.21ms step:1178/1555 train_time:62718ms step_avg:53.24ms step:1179/1555 train_time:62803ms step_avg:53.27ms step:1180/1555 train_time:62891ms step_avg:53.30ms step:1181/1555 train_time:62976ms step_avg:53.32ms step:1182/1555 train_time:63067ms step_avg:53.36ms step:1183/1555 train_time:63150ms step_avg:53.38ms step:1184/1555 train_time:63240ms step_avg:53.41ms step:1185/1555 train_time:63324ms step_avg:53.44ms step:1186/1555 train_time:63414ms step_avg:53.47ms step:1187/1555 train_time:63498ms step_avg:53.49ms step:1188/1555 train_time:63588ms step_avg:53.53ms step:1189/1555 train_time:63671ms step_avg:53.55ms step:1190/1555 train_time:63762ms step_avg:53.58ms step:1191/1555 train_time:63845ms step_avg:53.61ms step:1192/1555 train_time:63934ms step_avg:53.64ms step:1193/1555 train_time:64020ms step_avg:53.66ms step:1194/1555 train_time:64109ms step_avg:53.69ms step:1195/1555 train_time:64194ms step_avg:53.72ms step:1196/1555 train_time:64284ms step_avg:53.75ms step:1197/1555 train_time:64368ms step_avg:53.77ms step:1198/1555 train_time:64457ms step_avg:53.80ms step:1199/1555 train_time:64542ms step_avg:53.83ms step:1200/1555 train_time:64631ms step_avg:53.86ms step:1201/1555 train_time:64715ms step_avg:53.88ms step:1202/1555 train_time:64806ms step_avg:53.92ms step:1203/1555 train_time:64889ms step_avg:53.94ms step:1204/1555 train_time:64981ms step_avg:53.97ms step:1205/1555 train_time:65064ms step_avg:53.99ms step:1206/1555 train_time:65153ms step_avg:54.02ms step:1207/1555 train_time:65237ms step_avg:54.05ms step:1208/1555 train_time:65327ms step_avg:54.08ms step:1209/1555 train_time:65411ms step_avg:54.10ms step:1210/1555 train_time:65502ms step_avg:54.13ms step:1211/1555 train_time:65585ms step_avg:54.16ms step:1212/1555 train_time:65674ms step_avg:54.19ms step:1213/1555 train_time:65758ms step_avg:54.21ms step:1214/1555 train_time:65849ms step_avg:54.24ms step:1215/1555 train_time:65934ms step_avg:54.27ms step:1216/1555 train_time:66024ms step_avg:54.30ms step:1217/1555 train_time:66107ms step_avg:54.32ms step:1218/1555 train_time:66198ms step_avg:54.35ms step:1219/1555 train_time:66281ms step_avg:54.37ms step:1220/1555 train_time:66371ms step_avg:54.40ms step:1221/1555 train_time:66456ms step_avg:54.43ms step:1222/1555 train_time:66547ms step_avg:54.46ms step:1223/1555 train_time:66629ms step_avg:54.48ms step:1224/1555 train_time:66720ms step_avg:54.51ms step:1225/1555 train_time:66805ms step_avg:54.53ms step:1226/1555 train_time:66894ms step_avg:54.56ms step:1227/1555 train_time:66978ms step_avg:54.59ms step:1228/1555 train_time:67067ms step_avg:54.62ms step:1229/1555 train_time:67151ms step_avg:54.64ms step:1230/1555 train_time:67242ms step_avg:54.67ms step:1231/1555 train_time:67325ms step_avg:54.69ms step:1232/1555 train_time:67416ms step_avg:54.72ms step:1233/1555 train_time:67501ms step_avg:54.75ms step:1234/1555 train_time:67590ms step_avg:54.77ms step:1235/1555 train_time:67676ms step_avg:54.80ms step:1236/1555 train_time:67766ms step_avg:54.83ms step:1237/1555 train_time:67849ms step_avg:54.85ms step:1238/1555 train_time:67938ms step_avg:54.88ms step:1239/1555 train_time:68023ms step_avg:54.90ms step:1240/1555 train_time:68112ms step_avg:54.93ms step:1241/1555 train_time:68196ms step_avg:54.95ms step:1242/1555 train_time:68287ms step_avg:54.98ms step:1243/1555 train_time:68371ms step_avg:55.00ms step:1244/1555 train_time:68461ms step_avg:55.03ms step:1245/1555 train_time:68545ms step_avg:55.06ms step:1246/1555 train_time:68634ms step_avg:55.08ms step:1247/1555 train_time:68718ms step_avg:55.11ms step:1248/1555 train_time:68807ms step_avg:55.13ms step:1249/1555 train_time:68892ms step_avg:55.16ms step:1250/1555 train_time:68982ms step_avg:55.19ms step:1250/1555 val_loss:3.3982 train_time:69096ms step_avg:55.28ms step:1251/1555 train_time:69116ms step_avg:55.25ms step:1252/1555 train_time:69158ms step_avg:55.24ms step:1253/1555 train_time:69246ms step_avg:55.26ms step:1254/1555 train_time:69339ms step_avg:55.29ms step:1255/1555 train_time:69423ms step_avg:55.32ms step:1256/1555 train_time:69512ms step_avg:55.34ms step:1257/1555 train_time:69595ms step_avg:55.37ms step:1258/1555 train_time:69683ms step_avg:55.39ms step:1259/1555 train_time:69767ms step_avg:55.41ms step:1260/1555 train_time:69856ms step_avg:55.44ms step:1261/1555 train_time:69939ms step_avg:55.46ms step:1262/1555 train_time:70029ms step_avg:55.49ms step:1263/1555 train_time:70114ms step_avg:55.51ms step:1264/1555 train_time:70207ms step_avg:55.54ms step:1265/1555 train_time:70294ms step_avg:55.57ms step:1266/1555 train_time:70387ms step_avg:55.60ms step:1267/1555 train_time:70472ms step_avg:55.62ms step:1268/1555 train_time:70561ms step_avg:55.65ms step:1269/1555 train_time:70645ms step_avg:55.67ms step:1270/1555 train_time:70734ms step_avg:55.70ms step:1271/1555 train_time:70816ms step_avg:55.72ms step:1272/1555 train_time:70906ms step_avg:55.74ms step:1273/1555 train_time:70990ms step_avg:55.77ms step:1274/1555 train_time:71080ms step_avg:55.79ms step:1275/1555 train_time:71167ms step_avg:55.82ms step:1276/1555 train_time:71259ms step_avg:55.85ms step:1277/1555 train_time:71345ms step_avg:55.87ms step:1278/1555 train_time:71435ms step_avg:55.90ms step:1279/1555 train_time:71520ms step_avg:55.92ms step:1280/1555 train_time:71610ms step_avg:55.95ms step:1281/1555 train_time:71693ms step_avg:55.97ms step:1282/1555 train_time:71782ms step_avg:55.99ms step:1283/1555 train_time:71865ms step_avg:56.01ms step:1284/1555 train_time:71954ms step_avg:56.04ms step:1285/1555 train_time:72037ms step_avg:56.06ms step:1286/1555 train_time:72128ms step_avg:56.09ms step:1287/1555 train_time:72213ms step_avg:56.11ms step:1288/1555 train_time:72305ms step_avg:56.14ms step:1289/1555 train_time:72390ms step_avg:56.16ms step:1290/1555 train_time:72480ms step_avg:56.19ms step:1291/1555 train_time:72564ms step_avg:56.21ms step:1292/1555 train_time:72655ms step_avg:56.23ms step:1293/1555 train_time:72737ms step_avg:56.25ms step:1294/1555 train_time:72828ms step_avg:56.28ms step:1295/1555 train_time:72911ms step_avg:56.30ms step:1296/1555 train_time:73000ms step_avg:56.33ms step:1297/1555 train_time:73084ms step_avg:56.35ms step:1298/1555 train_time:73175ms step_avg:56.38ms step:1299/1555 train_time:73260ms step_avg:56.40ms step:1300/1555 train_time:73351ms step_avg:56.42ms step:1301/1555 train_time:73435ms step_avg:56.44ms step:1302/1555 train_time:73524ms step_avg:56.47ms step:1303/1555 train_time:73609ms step_avg:56.49ms step:1304/1555 train_time:73697ms step_avg:56.52ms step:1305/1555 train_time:73781ms step_avg:56.54ms step:1306/1555 train_time:73871ms step_avg:56.56ms step:1307/1555 train_time:73954ms step_avg:56.58ms step:1308/1555 train_time:74044ms step_avg:56.61ms step:1309/1555 train_time:74128ms step_avg:56.63ms step:1310/1555 train_time:74217ms step_avg:56.65ms step:1311/1555 train_time:74302ms step_avg:56.68ms step:1312/1555 train_time:74394ms step_avg:56.70ms step:1313/1555 train_time:74477ms step_avg:56.72ms step:1314/1555 train_time:74567ms step_avg:56.75ms step:1315/1555 train_time:74651ms step_avg:56.77ms step:1316/1555 train_time:74741ms step_avg:56.79ms step:1317/1555 train_time:74825ms step_avg:56.81ms step:1318/1555 train_time:74914ms step_avg:56.84ms step:1319/1555 train_time:74998ms step_avg:56.86ms step:1320/1555 train_time:75088ms step_avg:56.88ms step:1321/1555 train_time:75172ms step_avg:56.91ms step:1322/1555 train_time:75261ms step_avg:56.93ms step:1323/1555 train_time:75346ms step_avg:56.95ms step:1324/1555 train_time:75437ms step_avg:56.98ms step:1325/1555 train_time:75520ms step_avg:57.00ms step:1326/1555 train_time:75610ms step_avg:57.02ms step:1327/1555 train_time:75694ms step_avg:57.04ms step:1328/1555 train_time:75783ms step_avg:57.07ms step:1329/1555 train_time:75868ms step_avg:57.09ms step:1330/1555 train_time:75957ms step_avg:57.11ms step:1331/1555 train_time:76040ms step_avg:57.13ms step:1332/1555 train_time:76132ms step_avg:57.16ms step:1333/1555 train_time:76215ms step_avg:57.18ms step:1334/1555 train_time:76305ms step_avg:57.20ms step:1335/1555 train_time:76391ms step_avg:57.22ms step:1336/1555 train_time:76480ms step_avg:57.25ms step:1337/1555 train_time:76565ms step_avg:57.27ms step:1338/1555 train_time:76655ms step_avg:57.29ms step:1339/1555 train_time:76738ms step_avg:57.31ms step:1340/1555 train_time:76829ms step_avg:57.34ms step:1341/1555 train_time:76912ms step_avg:57.35ms step:1342/1555 train_time:77002ms step_avg:57.38ms step:1343/1555 train_time:77087ms step_avg:57.40ms step:1344/1555 train_time:77176ms step_avg:57.42ms step:1345/1555 train_time:77260ms step_avg:57.44ms step:1346/1555 train_time:77352ms step_avg:57.47ms step:1347/1555 train_time:77436ms step_avg:57.49ms step:1348/1555 train_time:77525ms step_avg:57.51ms step:1349/1555 train_time:77610ms step_avg:57.53ms step:1350/1555 train_time:77698ms step_avg:57.55ms step:1351/1555 train_time:77783ms step_avg:57.57ms step:1352/1555 train_time:77875ms step_avg:57.60ms step:1353/1555 train_time:77958ms step_avg:57.62ms step:1354/1555 train_time:78050ms step_avg:57.64ms step:1355/1555 train_time:78132ms step_avg:57.66ms step:1356/1555 train_time:78222ms step_avg:57.69ms step:1357/1555 train_time:78308ms step_avg:57.71ms step:1358/1555 train_time:78397ms step_avg:57.73ms step:1359/1555 train_time:78481ms step_avg:57.75ms step:1360/1555 train_time:78572ms step_avg:57.77ms step:1361/1555 train_time:78655ms step_avg:57.79ms step:1362/1555 train_time:78745ms step_avg:57.82ms step:1363/1555 train_time:78829ms step_avg:57.84ms step:1364/1555 train_time:78919ms step_avg:57.86ms step:1365/1555 train_time:79004ms step_avg:57.88ms step:1366/1555 train_time:79093ms step_avg:57.90ms step:1367/1555 train_time:79176ms step_avg:57.92ms step:1368/1555 train_time:79266ms step_avg:57.94ms step:1369/1555 train_time:79350ms step_avg:57.96ms step:1370/1555 train_time:79440ms step_avg:57.99ms step:1371/1555 train_time:79525ms step_avg:58.00ms step:1372/1555 train_time:79615ms step_avg:58.03ms step:1373/1555 train_time:79697ms step_avg:58.05ms step:1374/1555 train_time:79788ms step_avg:58.07ms step:1375/1555 train_time:79872ms step_avg:58.09ms step:1376/1555 train_time:79962ms step_avg:58.11ms step:1377/1555 train_time:80047ms step_avg:58.13ms step:1378/1555 train_time:80136ms step_avg:58.15ms step:1379/1555 train_time:80220ms step_avg:58.17ms step:1380/1555 train_time:80310ms step_avg:58.20ms step:1381/1555 train_time:80394ms step_avg:58.21ms step:1382/1555 train_time:80484ms step_avg:58.24ms step:1383/1555 train_time:80569ms step_avg:58.26ms step:1384/1555 train_time:80658ms step_avg:58.28ms step:1385/1555 train_time:80743ms step_avg:58.30ms step:1386/1555 train_time:80834ms step_avg:58.32ms step:1387/1555 train_time:80918ms step_avg:58.34ms step:1388/1555 train_time:81010ms step_avg:58.36ms step:1389/1555 train_time:81093ms step_avg:58.38ms step:1390/1555 train_time:81183ms step_avg:58.41ms step:1391/1555 train_time:81267ms step_avg:58.42ms step:1392/1555 train_time:81357ms step_avg:58.45ms step:1393/1555 train_time:81441ms step_avg:58.46ms step:1394/1555 train_time:81532ms step_avg:58.49ms step:1395/1555 train_time:81617ms step_avg:58.51ms step:1396/1555 train_time:81707ms step_avg:58.53ms step:1397/1555 train_time:81791ms step_avg:58.55ms step:1398/1555 train_time:81881ms step_avg:58.57ms step:1399/1555 train_time:81966ms step_avg:58.59ms step:1400/1555 train_time:82056ms step_avg:58.61ms step:1401/1555 train_time:82141ms step_avg:58.63ms step:1402/1555 train_time:82230ms step_avg:58.65ms step:1403/1555 train_time:82315ms step_avg:58.67ms step:1404/1555 train_time:82405ms step_avg:58.69ms step:1405/1555 train_time:82490ms step_avg:58.71ms step:1406/1555 train_time:82579ms step_avg:58.73ms step:1407/1555 train_time:82664ms step_avg:58.75ms step:1408/1555 train_time:82753ms step_avg:58.77ms step:1409/1555 train_time:82837ms step_avg:58.79ms step:1410/1555 train_time:82927ms step_avg:58.81ms step:1411/1555 train_time:83011ms step_avg:58.83ms step:1412/1555 train_time:83100ms step_avg:58.85ms step:1413/1555 train_time:83185ms step_avg:58.87ms step:1414/1555 train_time:83276ms step_avg:58.89ms step:1415/1555 train_time:83359ms step_avg:58.91ms step:1416/1555 train_time:83450ms step_avg:58.93ms step:1417/1555 train_time:83533ms step_avg:58.95ms step:1418/1555 train_time:83622ms step_avg:58.97ms step:1419/1555 train_time:83706ms step_avg:58.99ms step:1420/1555 train_time:83797ms step_avg:59.01ms step:1421/1555 train_time:83881ms step_avg:59.03ms step:1422/1555 train_time:83972ms step_avg:59.05ms step:1423/1555 train_time:84054ms step_avg:59.07ms step:1424/1555 train_time:84144ms step_avg:59.09ms step:1425/1555 train_time:84230ms step_avg:59.11ms step:1426/1555 train_time:84320ms step_avg:59.13ms step:1427/1555 train_time:84404ms step_avg:59.15ms step:1428/1555 train_time:84494ms step_avg:59.17ms step:1429/1555 train_time:84578ms step_avg:59.19ms step:1430/1555 train_time:84668ms step_avg:59.21ms step:1431/1555 train_time:84753ms step_avg:59.23ms step:1432/1555 train_time:84843ms step_avg:59.25ms step:1433/1555 train_time:84927ms step_avg:59.27ms step:1434/1555 train_time:85016ms step_avg:59.29ms step:1435/1555 train_time:85102ms step_avg:59.30ms step:1436/1555 train_time:85191ms step_avg:59.33ms step:1437/1555 train_time:85275ms step_avg:59.34ms step:1438/1555 train_time:85365ms step_avg:59.36ms step:1439/1555 train_time:85450ms step_avg:59.38ms step:1440/1555 train_time:85539ms step_avg:59.40ms step:1441/1555 train_time:85624ms step_avg:59.42ms step:1442/1555 train_time:85714ms step_avg:59.44ms step:1443/1555 train_time:85798ms step_avg:59.46ms step:1444/1555 train_time:85889ms step_avg:59.48ms step:1445/1555 train_time:85973ms step_avg:59.50ms step:1446/1555 train_time:86062ms step_avg:59.52ms step:1447/1555 train_time:86146ms step_avg:59.53ms step:1448/1555 train_time:86236ms step_avg:59.56ms step:1449/1555 train_time:86322ms step_avg:59.57ms step:1450/1555 train_time:86411ms step_avg:59.59ms step:1451/1555 train_time:86495ms step_avg:59.61ms step:1452/1555 train_time:86585ms step_avg:59.63ms step:1453/1555 train_time:86670ms step_avg:59.65ms step:1454/1555 train_time:86759ms step_avg:59.67ms step:1455/1555 train_time:86844ms step_avg:59.69ms step:1456/1555 train_time:86934ms step_avg:59.71ms step:1457/1555 train_time:87018ms step_avg:59.72ms step:1458/1555 train_time:87107ms step_avg:59.74ms step:1459/1555 train_time:87191ms step_avg:59.76ms step:1460/1555 train_time:87281ms step_avg:59.78ms step:1461/1555 train_time:87365ms step_avg:59.80ms step:1462/1555 train_time:87455ms step_avg:59.82ms step:1463/1555 train_time:87540ms step_avg:59.84ms step:1464/1555 train_time:87630ms step_avg:59.86ms step:1465/1555 train_time:87713ms step_avg:59.87ms step:1466/1555 train_time:87802ms step_avg:59.89ms step:1467/1555 train_time:87887ms step_avg:59.91ms step:1468/1555 train_time:87977ms step_avg:59.93ms step:1469/1555 train_time:88062ms step_avg:59.95ms step:1470/1555 train_time:88152ms step_avg:59.97ms step:1471/1555 train_time:88236ms step_avg:59.98ms step:1472/1555 train_time:88325ms step_avg:60.00ms step:1473/1555 train_time:88409ms step_avg:60.02ms step:1474/1555 train_time:88499ms step_avg:60.04ms step:1475/1555 train_time:88584ms step_avg:60.06ms step:1476/1555 train_time:88675ms step_avg:60.08ms step:1477/1555 train_time:88758ms step_avg:60.09ms step:1478/1555 train_time:88849ms step_avg:60.11ms step:1479/1555 train_time:88932ms step_avg:60.13ms step:1480/1555 train_time:89022ms step_avg:60.15ms step:1481/1555 train_time:89107ms step_avg:60.17ms step:1482/1555 train_time:89197ms step_avg:60.19ms step:1483/1555 train_time:89281ms step_avg:60.20ms step:1484/1555 train_time:89371ms step_avg:60.22ms step:1485/1555 train_time:89454ms step_avg:60.24ms step:1486/1555 train_time:89544ms step_avg:60.26ms step:1487/1555 train_time:89628ms step_avg:60.27ms step:1488/1555 train_time:89718ms step_avg:60.29ms step:1489/1555 train_time:89803ms step_avg:60.31ms step:1490/1555 train_time:89893ms step_avg:60.33ms step:1491/1555 train_time:89977ms step_avg:60.35ms step:1492/1555 train_time:90067ms step_avg:60.37ms step:1493/1555 train_time:90151ms step_avg:60.38ms step:1494/1555 train_time:90240ms step_avg:60.40ms step:1495/1555 train_time:90324ms step_avg:60.42ms step:1496/1555 train_time:90414ms step_avg:60.44ms step:1497/1555 train_time:90498ms step_avg:60.45ms step:1498/1555 train_time:90589ms step_avg:60.47ms step:1499/1555 train_time:90672ms step_avg:60.49ms step:1500/1555 train_time:90762ms step_avg:60.51ms step:1500/1555 val_loss:3.2946 train_time:90878ms step_avg:60.59ms step:1501/1555 train_time:90898ms step_avg:60.56ms step:1502/1555 train_time:90940ms step_avg:60.55ms step:1503/1555 train_time:91028ms step_avg:60.56ms step:1504/1555 train_time:91121ms step_avg:60.59ms step:1505/1555 train_time:91206ms step_avg:60.60ms step:1506/1555 train_time:91295ms step_avg:60.62ms step:1507/1555 train_time:91379ms step_avg:60.64ms step:1508/1555 train_time:91468ms step_avg:60.66ms step:1509/1555 train_time:91551ms step_avg:60.67ms step:1510/1555 train_time:91641ms step_avg:60.69ms step:1511/1555 train_time:91723ms step_avg:60.70ms step:1512/1555 train_time:91812ms step_avg:60.72ms step:1513/1555 train_time:91897ms step_avg:60.74ms step:1514/1555 train_time:91990ms step_avg:60.76ms step:1515/1555 train_time:92077ms step_avg:60.78ms step:1516/1555 train_time:92172ms step_avg:60.80ms step:1517/1555 train_time:92258ms step_avg:60.82ms step:1518/1555 train_time:92347ms step_avg:60.83ms step:1519/1555 train_time:92431ms step_avg:60.85ms step:1520/1555 train_time:92521ms step_avg:60.87ms step:1521/1555 train_time:92604ms step_avg:60.88ms step:1522/1555 train_time:92692ms step_avg:60.90ms step:1523/1555 train_time:92776ms step_avg:60.92ms step:1524/1555 train_time:92868ms step_avg:60.94ms step:1525/1555 train_time:92953ms step_avg:60.95ms step:1526/1555 train_time:93046ms step_avg:60.97ms step:1527/1555 train_time:93131ms step_avg:60.99ms step:1528/1555 train_time:93222ms step_avg:61.01ms step:1529/1555 train_time:93306ms step_avg:61.02ms step:1530/1555 train_time:93396ms step_avg:61.04ms step:1531/1555 train_time:93480ms step_avg:61.06ms step:1532/1555 train_time:93572ms step_avg:61.08ms step:1533/1555 train_time:93655ms step_avg:61.09ms step:1534/1555 train_time:93746ms step_avg:61.11ms step:1535/1555 train_time:93830ms step_avg:61.13ms step:1536/1555 train_time:93920ms step_avg:61.15ms step:1537/1555 train_time:94005ms step_avg:61.16ms step:1538/1555 train_time:94098ms step_avg:61.18ms step:1539/1555 train_time:94185ms step_avg:61.20ms step:1540/1555 train_time:94275ms step_avg:61.22ms step:1541/1555 train_time:94359ms step_avg:61.23ms step:1542/1555 train_time:94449ms step_avg:61.25ms step:1543/1555 train_time:94533ms step_avg:61.27ms step:1544/1555 train_time:94623ms step_avg:61.28ms step:1545/1555 train_time:94707ms step_avg:61.30ms step:1546/1555 train_time:94797ms step_avg:61.32ms step:1547/1555 train_time:94882ms step_avg:61.33ms step:1548/1555 train_time:94973ms step_avg:61.35ms step:1549/1555 train_time:95059ms step_avg:61.37ms step:1550/1555 train_time:95150ms step_avg:61.39ms step:1551/1555 train_time:95235ms step_avg:61.40ms step:1552/1555 train_time:95326ms step_avg:61.42ms step:1553/1555 train_time:95410ms step_avg:61.44ms step:1554/1555 train_time:95500ms step_avg:61.45ms step:1555/1555 train_time:95585ms step_avg:61.47ms step:1555/1555 val_loss:3.2787 train_time:95698ms step_avg:61.54ms peak memory allocated: 31630 MiB reserved: 46618 MiB