# Copyright (c) 2026 Lumina Moon and Contributors. # Licensed under the Apache License, Version 2.0 (see LICENSE for details). """ LuminaV: We Were Too Broke for AdamW So We Trapped Gradients in a Hyperbolic Straitjacket and Hired a Traffic Cop to Slap Them Official Repo : https://huggingface.co/cloverx-id/LuminaV-Optimizer-Paper Software DOI : 10.57967/hf/10365 Paper PDF : https://huggingface.co/cloverx-id/XoneLM-1.0-Paper/blob/main/LuminaV.pdf Paper DOI : 10.57967/hf/10270 """ from __future__ import annotations import contextlib import logging import math from typing import Callable, Dict, List, Optional, Set, Tuple, Union import torch from torch import Tensor from torch.optim import Optimizer __all__ = ["LuminaV"] __version__ = "1.3.0" __author__ = "Silver Moon (cloverxion)" __organization__ = "Lumina Moon" __license__ = "Apache-2.0" logger = logging.getLogger("LuminaV") HAS_TRITON = False try: import triton import triton.language as tl HAS_TRITON = True except ImportError: HAS_TRITON = False def _get_device_context(dev: torch.device): if dev.type == "cuda" and torch.cuda.is_available(): return torch.cuda.device(dev) elif hasattr(torch, "xpu") and dev.type == "xpu" and torch.xpu.is_available(): return torch.xpu.device(dev) elif hasattr(torch, "accelerator") and hasattr(torch.accelerator, "device"): return torch.accelerator.device(dev) return contextlib.nullcontext() if HAS_TRITON: @triton.jit def _prng(seed, idx): x = idx * 1103515245 + seed + 12345 x = x ^ (x >> 16) return x @triton.jit def _store_param(p_ptr, offsets, p_val, mask, seed, dtype_mode: tl.constexpr, use_sr: tl.constexpr): if use_sr: if dtype_mode == 1: p_int = p_val.to(tl.int32, bitcast=True) noise = _prng(seed, offsets) & 0xFFFF p_int = (p_int + noise) & ~0xFFFF p_out = p_int.to(tl.float32, bitcast=True).to(tl.bfloat16) tl.store(p_ptr + offsets, p_out, mask=mask) return elif dtype_mode == 2: p_int = p_val.to(tl.int32, bitcast=True) noise = _prng(seed, offsets) & 0x1FFF p_int = (p_int + noise) & ~0x1FFF p_out = p_int.to(tl.float32, bitcast=True).to(tl.float16) tl.store(p_ptr + offsets, p_out, mask=mask) return if dtype_mode == 1: tl.store(p_ptr + offsets, p_val.to(tl.bfloat16), mask=mask) elif dtype_mode == 2: tl.store(p_ptr + offsets, p_val.to(tl.float16), mask=mask) else: tl.store(p_ptr + offsets, p_val, mask=mask) @triton.jit def _store_state_buffer( ptr, offsets, val, mask, seed, state_dtype_mode: tl.constexpr, use_sr: tl.constexpr ): if use_sr: if state_dtype_mode == 1: p_int = val.to(tl.int32, bitcast=True) noise = _prng(seed, offsets) & 0xFFFF p_int = (p_int + noise) & ~0xFFFF p_out = p_int.to(tl.float32, bitcast=True).to(tl.bfloat16) tl.store(ptr + offsets, p_out, mask=mask) return elif state_dtype_mode == 2: p_int = val.to(tl.int32, bitcast=True) noise = _prng(seed, offsets) & 0x1FFF p_int = (p_int + noise) & ~0x1FFF p_out = p_int.to(tl.float32, bitcast=True).to(tl.float16) tl.store(ptr + offsets, p_out, mask=mask) return if state_dtype_mode == 1: tl.store(ptr + offsets, val.to(tl.bfloat16), mask=mask) elif state_dtype_mode == 2: tl.store(ptr + offsets, val.to(tl.float16), mask=mask) else: tl.store(ptr + offsets, val, mask=mask) _HAS_TL_MATH_TANH = hasattr(tl, "math") and hasattr(tl.math, "tanh") _HAS_TL_ROOT_TANH = hasattr(tl, "tanh") if _HAS_TL_MATH_TANH: @triton.jit def _triton_tanh_fast(x): return tl.math.tanh(x) elif _HAS_TL_ROOT_TANH: @triton.jit def _triton_tanh_fast(x): return tl.tanh(x) else: @triton.jit def _triton_tanh_fast(x): return (2.0 / (1.0 + tl.exp(-2.0 * x))) - 1.0 @triton.jit def _lumina_v2_fused_single_pass_kernel( p_ptr, master_p_ptr, grad_ptr, exp_avg_ptr, exp_avg_sq_ptr, scratch_ptr, n_elements, beta1, beta2, c1, c2, lr, weight_decay, bound_mode: tl.constexpr, bound_ratio, seed, dtype_mode: tl.constexpr, state_dtype_mode: tl.constexpr, use_sr: tl.constexpr, cautious: tl.constexpr, has_master: tl.constexpr, BLOCK_SIZE: tl.constexpr ): pid = tl.program_id(axis=0) offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < n_elements seed_m = seed seed_p = (seed + 0x55555555) & 0x7FFFFFFF mask_sum_ptr = scratch_ptr + 0 u_sq_sum_ptr = scratch_ptr + 1 p_sq_sum_ptr = scratch_ptr + 2 cautious_scale_ptr = scratch_ptr + 3 bound_scale_ptr = scratch_ptr + 4 cautious_scale = tl.load(cautious_scale_ptr) bound_scale = tl.load(bound_scale_ptr) if has_master: p = tl.load(master_p_ptr + offsets, mask=mask, other=0.0) else: p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32) p_init = p g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32) is_finite_g = (g == g) & (g < 1e38) & (g > -1e38) g = tl.where(is_finite_g, g, 0.0) m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32) v = tl.load(exp_avg_sq_ptr + offsets, mask=mask, other=0.0).to(tl.float32) m_new = beta1 * m + (1.0 - beta1) * g nes_m = beta1 * m_new + (1.0 - beta1) * g diff = g - m_new v_new = beta2 * v + (1.0 - beta2) * (diff * diff) if dtype_mode == 2 or state_dtype_mode == 2: v_new = tl.minimum(v_new, 65000.0) sigma = tl.sqrt(tl.maximum(v_new, 0.0)) * c1 + c2 u = _triton_tanh_fast(nes_m / tl.maximum(sigma, 1e-12)) m_mask = tl.where((u * g) > 0.0, 1.0, 0.0) _store_state_buffer(exp_avg_ptr, offsets, m_new, mask, seed_m, state_dtype_mode, use_sr) _store_state_buffer(exp_avg_sq_ptr, offsets, v_new, mask, seed_m, state_dtype_mode, use_sr=False) if cautious: delta_theta = (u * m_mask) * cautious_scale else: delta_theta = u step_val = lr * delta_theta if bound_mode == 1: step_val = step_val * bound_scale elif bound_mode == 2: B_i = bound_ratio * tl.sqrt(p_init * p_init + 1e-4) step_val = B_i * _triton_tanh_fast(step_val / B_i) if dtype_mode == 2: b_safe = tl.maximum(tl.abs(p_init) * 0.01, 1.5e-4) step_val = b_safe * _triton_tanh_fast(step_val / b_safe) if weight_decay != 0.0: p = p * (1.0 - lr * weight_decay) p_updated = p - step_val if has_master: tl.store(master_p_ptr + offsets, p_updated, mask=mask) _store_param(p_ptr, offsets, p_updated, mask, seed_p, dtype_mode, use_sr=False) else: _store_param(p_ptr, offsets, p_updated, mask, seed_p, dtype_mode, use_sr) if cautious: block_sum = tl.sum(tl.where(mask, m_mask, 0.0), axis=0) tl.atomic_add(mask_sum_ptr, block_sum) if bound_mode == 1: if cautious: u_eff = u * m_mask else: u_eff = u block_u_sq = tl.sum(tl.where(mask, u_eff * u_eff, 0.0), axis=0) tl.atomic_add(u_sq_sum_ptr, block_u_sq) block_p_sq = tl.sum(tl.where(mask, p_init * p_init, 0.0), axis=0) tl.atomic_add(p_sq_sum_ptr, block_p_sq) @triton.jit def _lumina_v2_post_step_ema_kernel( scratch_ptr, n_elements, lr, clamp_min, ema_decay, step, bound_mode: tl.constexpr, bound_ratio, cautious: tl.constexpr ): if tl.program_id(axis=0) == 0: mask_sum_ptr = scratch_ptr + 0 u_sq_sum_ptr = scratch_ptr + 1 p_sq_sum_ptr = scratch_ptr + 2 cautious_scale_ptr = scratch_ptr + 3 bound_scale_ptr = scratch_ptr + 4 if cautious: mask_sum = tl.load(mask_sum_ptr) m_bar = tl.minimum(tl.maximum(mask_sum / n_elements, clamp_min), 1.0) target_cautious_scale = 1.0 / m_bar old_cautious = tl.load(cautious_scale_ptr) new_cautious = ema_decay * old_cautious + (1.0 - ema_decay) * target_cautious_scale tl.store(cautious_scale_ptr, new_cautious) if bound_mode == 1: u_sq_sum = tl.load(u_sq_sum_ptr) p_sq_sum = tl.load(p_sq_sum_ptr) if cautious: c_scale = tl.load(cautious_scale_ptr) else: c_scale = 1.0 norm_p = tl.sqrt(tl.maximum(p_sq_sum, 0.0)) norm_delta = tl.sqrt(tl.maximum(u_sq_sum, 0.0)) * c_scale R = bound_ratio * tl.maximum(norm_p, 1.0) r = (lr * norm_delta) / R target_scale = tl.where(r < 1e-4, 1.0, _triton_tanh_fast(r) / tl.maximum(r, 1e-8)) old_bound = tl.load(bound_scale_ptr) new_bound = ema_decay * old_bound + (1.0 - ema_decay) * target_scale tl.store(bound_scale_ptr, new_bound) tl.store(mask_sum_ptr, 0.0) tl.store(u_sq_sum_ptr, 0.0) tl.store(p_sq_sum_ptr, 0.0) # ARCHITECTURAL RATIONALE: REDUCTION STRATEGY & ALTERNATIVES REJECTION # # Global reductions (mask_sum, u_sq_sum, p_sq_sum) across Pass 1 adopt # a hybrid Intra-Block Tree Reduction + Inter-Block Hardware Atomic pattern: # 1. Each CTA (thread block) performs parallel register-level reduction # via `tl.sum(..., axis=0)` across local tile memory. # 2. A single `tl.atomic_add` per CTA commits the block partial sum to # a dedicated L2-resident GPU scratch buffer (`scratch_ptr`). # 3. Pass 2 immediately reads the finalized scalars directly from the # scratch buffer within the same CUDA stream. # # EXPLICIT REJECTION OF ALTERNATIVE ARCHITECTURES: # # 1. Rejection of Host-Side Synchronization (.item() / PCIe Roundtrips): # Exporting reductions back to the CPU via `.item()` introduces blocking # CUDA stream fences and severe PCIe latency penalties (~10-25 us per # tensor). Across deep 30+ layer Transformer architectures, this creates # hundreds of pipeline bubbles per step and fatally shatters CUDA Graph # capture (`torch.accelerator.Graph`) as well as TorchInductor AOT tracing. # # 2. Rejection of Multi-Tier Hierarchical Grid Reductions (Cascaded Passes): # Traditional hierarchical tree reductions (CTA reduction -> intermediate # grid buffer -> secondary reduction kernel launch) require dispatching # auxiliary micro-kernels simply to resolve three scalars. The cumulative # kernel launch overhead and driver latency far outweigh any hypothetical # atomic throughput gain. # For a 10M-parameter tensor with BLOCK_SIZE=2048, only ~4,882 atomic # additions are issued across the entire grid. This volume resides well # within the bandwidth budget of modern GPU L2 cache atomic units, # resolving in sub-microsecond time with zero extra dispatch bubbles. @triton.jit def _lumina_v2_pass1_kernel( p_ptr, master_p_ptr, grad_ptr, exp_avg_ptr, exp_avg_sq_ptr, scratch_ptr, n_elements, beta1, beta2, c1, c2, seed, bound_mode: tl.constexpr, dtype_mode: tl.constexpr, state_dtype_mode: tl.constexpr, use_sr: tl.constexpr, cautious: tl.constexpr, has_master: tl.constexpr, BLOCK_SIZE: tl.constexpr ): pid = tl.program_id(axis=0) offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < n_elements seed_m = seed mask_sum_ptr = scratch_ptr + 0 u_sq_sum_ptr = scratch_ptr + 1 p_sq_sum_ptr = scratch_ptr + 2 g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32) is_finite_g = (g == g) & (g < 1e38) & (g > -1e38) g = tl.where(is_finite_g, g, 0.0) m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32) v = tl.load(exp_avg_sq_ptr + offsets, mask=mask, other=0.0).to(tl.float32) m_new = beta1 * m + (1.0 - beta1) * g nes_m = beta1 * m_new + (1.0 - beta1) * g diff = g - m_new v_new = beta2 * v + (1.0 - beta2) * (diff * diff) if dtype_mode == 2 or state_dtype_mode == 2: v_new = tl.minimum(v_new, 65000.0) sigma = tl.sqrt(tl.maximum(v_new, 0.0)) * c1 + c2 u = _triton_tanh_fast(nes_m / tl.maximum(sigma, 1e-12)) m_mask = tl.where((u * g) > 0.0, 1.0, 0.0) _store_state_buffer(exp_avg_ptr, offsets, m_new, mask, seed_m, state_dtype_mode, use_sr) _store_state_buffer(exp_avg_sq_ptr, offsets, v_new, mask, seed_m, state_dtype_mode, use_sr=False) block_sum = tl.sum(tl.where(mask, m_mask, 0.0), axis=0) tl.atomic_add(mask_sum_ptr, block_sum) if bound_mode == 1: if cautious: u_eff = u * m_mask else: u_eff = u block_u_sq = tl.sum(tl.where(mask, u_eff * u_eff, 0.0), axis=0) tl.atomic_add(u_sq_sum_ptr, block_u_sq) if has_master: p = tl.load(master_p_ptr + offsets, mask=mask, other=0.0) else: p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32) block_p_sq = tl.sum(tl.where(mask, p * p, 0.0), axis=0) tl.atomic_add(p_sq_sum_ptr, block_p_sq) @triton.jit def _lumina_v2_pass2_kernel( p_ptr, master_p_ptr, grad_ptr, exp_avg_ptr, exp_avg_sq_ptr, scratch_ptr, n_elements, beta1, c1, c2, lr, weight_decay, clamp_min, bound_mode: tl.constexpr, bound_ratio, seed, dtype_mode: tl.constexpr, use_sr: tl.constexpr, cautious: tl.constexpr, has_master: tl.constexpr, BLOCK_SIZE: tl.constexpr ): pid = tl.program_id(axis=0) offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < n_elements seed_p = (seed + 0x55555555) & 0x7FFFFFFF cautious_scale = 1.0 if cautious: mask_sum = tl.load(scratch_ptr + 0) m_bar = tl.minimum(tl.maximum(mask_sum / n_elements, clamp_min), 1.0) cautious_scale = 1.0 / m_bar bound_scale = 1.0 if bound_mode == 1: u_sq_sum = tl.load(scratch_ptr + 1) p_sq_sum = tl.load(scratch_ptr + 2) norm_p = tl.sqrt(tl.maximum(p_sq_sum, 0.0)) norm_delta = tl.sqrt(tl.maximum(u_sq_sum, 0.0)) * cautious_scale R = bound_ratio * tl.maximum(norm_p, 1.0) r = (lr * norm_delta) / R bound_scale = tl.where(r < 1e-4, 1.0, _triton_tanh_fast(r) / tl.maximum(r, 1e-8)) if has_master: p = tl.load(master_p_ptr + offsets, mask=mask, other=0.0) else: p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32) p_init = p g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32) is_finite_g = (g == g) & (g < 1e38) & (g > -1e38) g = tl.where(is_finite_g, g, 0.0) m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32) v = tl.load(exp_avg_sq_ptr + offsets, mask=mask, other=0.0).to(tl.float32) nes_m = beta1 * m + (1.0 - beta1) * g sigma = tl.sqrt(tl.maximum(v, 0.0)) * c1 + c2 u = _triton_tanh_fast(nes_m / tl.maximum(sigma, 1e-12)) m_mask = tl.where((u * g) > 0.0, 1.0, 0.0) if cautious: delta_theta = (u * m_mask) * cautious_scale else: delta_theta = u step_val = lr * delta_theta if bound_mode == 1: step_val = step_val * bound_scale elif bound_mode == 2: B_i = bound_ratio * tl.sqrt(p_init * p_init + 1e-4) step_val = B_i * _triton_tanh_fast(step_val / B_i) if dtype_mode == 2: b_safe = tl.maximum(tl.abs(p_init) * 0.01, 1.5e-4) step_val = b_safe * _triton_tanh_fast(step_val / b_safe) if weight_decay != 0.0: p = p * (1.0 - lr * weight_decay) p_updated = p - step_val if has_master: tl.store(master_p_ptr + offsets, p_updated, mask=mask) _store_param(p_ptr, offsets, p_updated, mask, seed_p, dtype_mode, use_sr=False) else: _store_param(p_ptr, offsets, p_updated, mask, seed_p, dtype_mode, use_sr) if tl.program_id(axis=0) == 0: tl.store(scratch_ptr + 3, cautious_scale) tl.store(scratch_ptr + 4, bound_scale) @triton.jit def _lumina_v2_single_pass_kernel( p_ptr, master_p_ptr, grad_ptr, exp_avg_ptr, exp_avg_sq_ptr, n_elements, beta1, beta2, c1, c2, lr, weight_decay, bound_mode: tl.constexpr, bound_ratio, seed, dtype_mode: tl.constexpr, state_dtype_mode: tl.constexpr, use_sr: tl.constexpr, has_master: tl.constexpr, BLOCK_SIZE: tl.constexpr ): pid = tl.program_id(axis=0) offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < n_elements seed_m = seed seed_p = (seed + 0x55555555) & 0x7FFFFFFF if has_master: p = tl.load(master_p_ptr + offsets, mask=mask, other=0.0) else: p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32) p_init = p g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32) is_finite_g = (g == g) & (g < 1e38) & (g > -1e38) g = tl.where(is_finite_g, g, 0.0) m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32) v = tl.load(exp_avg_sq_ptr + offsets, mask=mask, other=0.0).to(tl.float32) m_new = beta1 * m + (1.0 - beta1) * g nes_m = beta1 * m_new + (1.0 - beta1) * g diff = g - m_new v_new = beta2 * v + (1.0 - beta2) * (diff * diff) if dtype_mode == 2 or state_dtype_mode == 2: v_new = tl.minimum(v_new, 65000.0) sigma = tl.sqrt(tl.maximum(v_new, 0.0)) * c1 + c2 u = _triton_tanh_fast(nes_m / tl.maximum(sigma, 1e-12)) _store_state_buffer(exp_avg_ptr, offsets, m_new, mask, seed_m, state_dtype_mode, use_sr) _store_state_buffer(exp_avg_sq_ptr, offsets, v_new, mask, seed_m, state_dtype_mode, use_sr=False) step_val = lr * u if bound_mode == 2: B_i = bound_ratio * tl.sqrt(p_init * p_init + 1e-4) step_val = B_i * _triton_tanh_fast(step_val / B_i) if dtype_mode == 2: b_safe = tl.maximum(tl.abs(p_init) * 0.01, 1.5e-4) step_val = b_safe * _triton_tanh_fast(step_val / b_safe) if weight_decay != 0.0: p = p * (1.0 - lr * weight_decay) p_updated = p - step_val if has_master: tl.store(master_p_ptr + offsets, p_updated, mask=mask) _store_param(p_ptr, offsets, p_updated, mask, seed_p, dtype_mode, use_sr=False) else: _store_param(p_ptr, offsets, p_updated, mask, seed_p, dtype_mode, use_sr) @triton.jit def _lumina_v1_pass1_kernel( p_ptr, master_p_ptr, grad_ptr, exp_avg_ptr, scratch_ptr, n_elements, beta1, seed, bound_mode: tl.constexpr, state_dtype_mode: tl.constexpr, use_sr: tl.constexpr, has_master: tl.constexpr, BLOCK_SIZE: tl.constexpr ): pid = tl.program_id(axis=0) offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < n_elements seed_m = seed rms_sum_ptr = scratch_ptr + 0 p_sq_sum_ptr = scratch_ptr + 3 g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32) is_finite_g = (g == g) & (g < 1e38) & (g > -1e38) g = tl.where(is_finite_g, g, 0.0) m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32) m_new = beta1 * m + (1.0 - beta1) * g nes_m = beta1 * m_new + (1.0 - beta1) * g _store_state_buffer(exp_avg_ptr, offsets, m_new, mask, seed_m, state_dtype_mode, use_sr) sq_val = nes_m * nes_m block_sq_sum = tl.sum(tl.where(mask, sq_val, 0.0), axis=0) tl.atomic_add(rms_sum_ptr, block_sq_sum) if bound_mode == 1: if has_master: p = tl.load(master_p_ptr + offsets, mask=mask, other=0.0) else: p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32) block_p_sq = tl.sum(tl.where(mask, p * p, 0.0), axis=0) tl.atomic_add(p_sq_sum_ptr, block_p_sq) @triton.jit def _lumina_v1_pass2_kernel( grad_ptr, exp_avg_ptr, scratch_ptr, n_elements, beta1, tau, eps, c2, alpha_ss, cautious: tl.constexpr, BLOCK_SIZE: tl.constexpr ): pid = tl.program_id(axis=0) offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < n_elements rms_sum_ptr = scratch_ptr + 0 mask_sum_ptr = scratch_ptr + 1 u_sq_sum_ptr = scratch_ptr + 2 sq_sum = tl.load(rms_sum_ptr) rms = tl.sqrt(sq_sum / n_elements + eps) sigma = tau * rms + c2 g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32) is_finite_g = (g == g) & (g < 1e38) & (g > -1e38) g = tl.where(is_finite_g, g, 0.0) m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32) nes_m = beta1 * m + (1.0 - beta1) * g z = nes_m / tl.maximum(sigma, 1e-12) z_soft = z / (1.0 + alpha_ss * tl.abs(z)) u = _triton_tanh_fast(z_soft) m_mask = tl.where((u * g) > 0.0, 1.0, 0.0) block_sum = tl.sum(tl.where(mask, m_mask, 0.0), axis=0) tl.atomic_add(mask_sum_ptr, block_sum) if cautious: u_eff = u * m_mask else: u_eff = u block_u_sq = tl.sum(tl.where(mask, u_eff * u_eff, 0.0), axis=0) tl.atomic_add(u_sq_sum_ptr, block_u_sq) @triton.jit def _lumina_v1_pass3_update_kernel( p_ptr, master_p_ptr, grad_ptr, exp_avg_ptr, scratch_ptr, n_elements, beta1, tau, eps, c2, alpha_ss, lr, weight_decay, clamp_min, bound_mode: tl.constexpr, bound_ratio, seed, dtype_mode: tl.constexpr, use_sr: tl.constexpr, cautious: tl.constexpr, has_master: tl.constexpr, BLOCK_SIZE: tl.constexpr ): pid = tl.program_id(axis=0) offsets = pid * BLOCK_SIZE + tl.arange(0, BLOCK_SIZE) mask = offsets < n_elements seed_p = (seed + 0x55555555) & 0x7FFFFFFF rms_sum_ptr = scratch_ptr + 0 mask_sum_ptr = scratch_ptr + 1 u_sq_sum_ptr = scratch_ptr + 2 p_sq_sum_ptr = scratch_ptr + 3 sq_sum = tl.load(rms_sum_ptr) rms = tl.sqrt(sq_sum / n_elements + eps) sigma = tau * rms + c2 m_sum = tl.load(mask_sum_ptr) m_bar = tl.minimum(tl.maximum(m_sum / n_elements, clamp_min), 1.0) if has_master: p = tl.load(master_p_ptr + offsets, mask=mask, other=0.0) else: p = tl.load(p_ptr + offsets, mask=mask, other=0.0).to(tl.float32) p_init = p g = tl.load(grad_ptr + offsets, mask=mask, other=0.0).to(tl.float32) is_finite_g = (g == g) & (g < 1e38) & (g > -1e38) g = tl.where(is_finite_g, g, 0.0) m = tl.load(exp_avg_ptr + offsets, mask=mask, other=0.0).to(tl.float32) nes_m = beta1 * m + (1.0 - beta1) * g z = nes_m / tl.maximum(sigma, 1e-12) z_soft = z / (1.0 + alpha_ss * tl.abs(z)) u = _triton_tanh_fast(z_soft) m_mask = tl.where((u * g) > 0.0, 1.0, 0.0) if cautious: delta_theta = (u * m_mask) / m_bar else: delta_theta = u step_val = lr * delta_theta if bound_mode == 1: u_sq_sum = tl.load(u_sq_sum_ptr) p_sq_sum = tl.load(p_sq_sum_ptr) norm_p = tl.sqrt(tl.maximum(p_sq_sum, 0.0)) if cautious: norm_delta = tl.sqrt(tl.maximum(u_sq_sum, 0.0)) / m_bar else: norm_delta = tl.sqrt(tl.maximum(u_sq_sum, 0.0)) R = bound_ratio * tl.maximum(norm_p, 1.0) r = (lr * norm_delta) / R scale = tl.where(r < 1e-4, 1.0, _triton_tanh_fast(r) / tl.maximum(r, 1e-8)) step_val = step_val * scale elif bound_mode == 2: B_i = bound_ratio * tl.sqrt(p_init * p_init + 1e-4) step_val = B_i * _triton_tanh_fast(step_val / B_i) if dtype_mode == 2: b_safe = tl.maximum(tl.abs(p_init) * 0.01, 1.5e-4) step_val = b_safe * _triton_tanh_fast(step_val / b_safe) if weight_decay != 0.0: p = p * (1.0 - lr * weight_decay) p_updated = p - step_val if has_master: tl.store(master_p_ptr + offsets, p_updated, mask=mask) _store_param(p_ptr, offsets, p_updated, mask, seed_p, dtype_mode, use_sr=False) else: _store_param(p_ptr, offsets, p_updated, mask, seed_p, dtype_mode, use_sr) def _apply_bound( update: Tensor, p: Tensor, lr: float, bound: bool, bound_type: str, bound_ratio: float, target_dtype: Optional[torch.dtype] = None, ) -> Tensor: if lr == 0.0: return update effective_dtype = target_dtype if target_dtype is not None else p.dtype u_float = torch.nan_to_num(update.float(), nan=0.0, posinf=1.0, neginf=-1.0) step_val = lr * u_float p_ref = p.float() if bound and bound_ratio > 0.0: if bound_type == "radial": norm_u = torch.linalg.vector_norm(u_float) norm_p = torch.linalg.vector_norm(p_ref) R = bound_ratio * torch.clamp_min(norm_p, 1.0) r = (lr * norm_u) / R scale = torch.where(r < 1e-4, torch.ones_like(r), torch.tanh(r) / torch.clamp_min(r, 1e-8)) step_val = step_val * scale elif bound_type == "coordinate": B = bound_ratio * torch.sqrt(p_ref.square() + 1e-4) step_val = B * torch.tanh(step_val / B) if effective_dtype == torch.float16: b_safe = torch.clamp_min(torch.abs(p_ref) * 0.01, 1.5e-4) step_val = b_safe * torch.tanh(step_val / b_safe) return step_val / lr def _sr_update( p: Tensor, update: Tensor, lr: float, weight_decay: float, use_sr: bool, generator: Optional[torch.Generator] = None, ): update_f32 = torch.nan_to_num(update.float(), nan=0.0, posinf=1.0, neginf=-1.0) if p.dtype not in (torch.bfloat16, torch.float16): if weight_decay != 0.0: p.mul_(1.0 - lr * weight_decay) p.add_(update_f32, alpha=-lr) return p_fp32 = p.float() if weight_decay != 0.0: p_fp32.mul_(1.0 - lr * weight_decay) p_fp32.add_(update_f32, alpha=-lr) if use_sr: int_view = p_fp32.contiguous().view(torch.int32) n_elem = int_view.numel() mask_bits = 0xFFFF if p.dtype == torch.bfloat16 else 0x1FFF shift_range = 1 << 16 if p.dtype == torch.bfloat16 else 1 << 13 chunk_size = 1_048_576 flat_view = int_view.view(-1) if n_elem <= chunk_size: try: noise = torch.empty(n_elem, dtype=torch.int32, device=p.device).random_( 0, shift_range, generator=generator ) except RuntimeError: noise = torch.empty(n_elem, dtype=torch.int32, device="cpu").random_( 0, shift_range, generator=None ).to(device=p.device, non_blocking=True) flat_view.add_(noise).bitwise_and_(~mask_bits) else: buf_device = p.device try: chunk_noise = torch.empty(chunk_size, dtype=torch.int32, device=buf_device) except RuntimeError: buf_device = torch.device("cpu") chunk_noise = torch.empty(chunk_size, dtype=torch.int32, device=buf_device) for offset in range(0, n_elem, chunk_size): curr_len = min(chunk_size, n_elem - offset) sub_slice = flat_view[offset : offset + curr_len] sub_noise = chunk_noise[:curr_len] if buf_device.type == "cpu": sub_noise.random_(0, shift_range, generator=None) sub_slice.add_(sub_noise.to(device=p.device, non_blocking=True)).bitwise_and_(~mask_bits) else: sub_noise.random_(0, shift_range, generator=generator) sub_slice.add_(sub_noise).bitwise_and_(~mask_bits) p.copy_(int_view.view(torch.float32).to(dtype=p.dtype)) else: p.copy_(p_fp32.to(dtype=p.dtype)) class LuminaV(Optimizer): def __init__( self, params, lr: float = 8e-4, betas: Tuple[float, float] = (0.9, 0.999), eps: float = 1e-8, weight_decay: float = 8e-2, tau: float = 0.8, alpha_ss: float = 0.5, cautious: bool = True, cautious_clamp_min: float = 0.5, buffer: int = 2, stochastic_rounding: bool = True, bound: bool = True, bound_type: str = "radial", bound_ratio: float = 0.03, master_weights: Union[bool, str] = False, execution: str = "auto", fused_single_pass: bool = False, ema_decay: float = 0.8, seed: int = 1337, ): if lr < 0.0: raise ValueError(f"Invalid learning rate: {lr}") if not 0.0 <= betas[0] < 1.0: raise ValueError(f"Invalid beta1 parameter: {betas[0]}") if not 0.0 <= betas[1] < 1.0: raise ValueError(f"Invalid beta2 parameter: {betas[1]}") if eps <= 0.0: raise ValueError(f"Invalid epsilon value: {eps}") if weight_decay < 0.0: raise ValueError(f"Invalid weight_decay value: {weight_decay}") if tau <= 0.0: raise ValueError(f"Invalid tau parameter: {tau}") if not 0.0 < cautious_clamp_min <= 1.0: raise ValueError(f"Invalid cautious_clamp_min: {cautious_clamp_min}") if buffer not in (1, 2): raise ValueError(f"Buffer count must be 1 (Single) or 2 (Dual), got: {buffer}") if bound_type not in ("radial", "coordinate"): raise ValueError(f"Invalid bound_type: {bound_type}. Must be 'radial' or 'coordinate'") if bound_ratio <= 0.0: raise ValueError(f"Invalid bound_ratio: {bound_ratio}. Must be > 0.0") if not isinstance(seed, int) or seed < 0: raise ValueError(f"Invalid seed parameter: {seed}. Must be a non-negative integer") if isinstance(master_weights, bool): master_mode = "full" if master_weights else "none" elif isinstance(master_weights, str): m_str = master_weights.lower().strip() if m_str in ("false", "none", "off"): master_mode = "none" elif m_str in ("semi", "hybrid", "half"): master_mode = "semi" elif m_str in ("full", "true", "fp32", "masterweight"): master_mode = "full" else: raise ValueError(f"Invalid master_weights option: {master_weights}. Must be False, 'semi', or 'full'") else: raise TypeError(f"master_weights must be bool or str, got: {type(master_weights).__name__}") defaults = dict( lr=lr, betas=betas, eps=eps, weight_decay=weight_decay, tau=tau, alpha_ss=alpha_ss, cautious=cautious, cautious_clamp_min=cautious_clamp_min, buffer=buffer, stochastic_rounding=stochastic_rounding, bound=bound, bound_type=bound_type, bound_ratio=bound_ratio, master_weights=master_mode, execution=execution, fused_single_pass=fused_single_pass, ema_decay=ema_decay, seed=seed ) super().__init__(params, defaults) self.seed = seed self._warned_non_contiguous = False self._generators: Dict[Tuple[str, int, int], torch.Generator] = {} def _get_generator(self, device: torch.device, seed: Optional[int] = None) -> torch.Generator: effective_seed = self.seed if seed is None else seed dev_type = device.type dev_idx = device.index if device.index is not None else 0 key = (dev_type, dev_idx, effective_seed) if key not in self._generators: try: gen = torch.Generator(device=device) except (RuntimeError, ValueError, TypeError): gen = torch.Generator(device="cpu") gen.manual_seed(effective_seed) self._generators[key] = gen return self._generators[key] @torch.no_grad() def step(self, closure: Optional[Callable[[], float]] = None) -> Optional[float]: loss = None if closure is not None: with torch.enable_grad(): loss = closure() for group in self.param_groups: params_with_grad = [] grads = [] exp_avgs = [] exp_avg_sqs = [] master_params = [] steps = [] scratches = [] buf_count = group["buffer"] master_mode = group["master_weights"] for p in group["params"]: if p.grad is None: continue if p.grad.is_sparse: raise RuntimeError("LuminaV does not support sparse gradients.") if p.numel() == 0: continue state = self.state[p] if len(state) == 0: state["step"] = 0 if master_mode != "none": state["master_param"] = p.detach().clone().float().contiguous() state_dtype = torch.float32 if master_mode == "full" else p.dtype state["exp_avg"] = torch.zeros_like(p, dtype=state_dtype, memory_format=torch.contiguous_format) if buf_count == 2: state["exp_avg_sq"] = torch.zeros_like(p, dtype=state_dtype, memory_format=torch.contiguous_format) scratch = torch.zeros((5,), device=p.device, dtype=torch.float32) if buf_count == 2: scratch[3] = 1.0 scratch[4] = 1.0 state["scratch"] = scratch state["buf_count"] = buf_count else: if master_mode != "none": if "master_param" not in state: state["master_param"] = p.detach().clone().float().contiguous() elif state["master_param"].device != p.device: state["master_param"] = state["master_param"].to(device=p.device, non_blocking=True) if "scratch" not in state: scratch = torch.zeros((5,), device=p.device, dtype=torch.float32) if buf_count == 2: scratch[3] = 1.0 scratch[4] = 1.0 state["scratch"] = scratch state["buf_count"] = buf_count elif state["scratch"].device != p.device: state["scratch"] = state["scratch"].to(device=p.device, non_blocking=True) if buf_count == 2: state["scratch"][3] = 1.0 state["scratch"][4] = 1.0 state["buf_count"] = buf_count elif state.get("buf_count") != buf_count: state["scratch"].zero_() if buf_count == 2: state["scratch"][3] = 1.0 state["scratch"][4] = 1.0 state["buf_count"] = buf_count if "exp_avg" in state and state["exp_avg"].device != p.device: state["exp_avg"] = state["exp_avg"].to(device=p.device, non_blocking=True) if buf_count == 2: if "exp_avg_sq" not in state: state_dtype = torch.float32 if master_mode == "full" else p.dtype state["exp_avg_sq"] = torch.zeros_like(p, dtype=state_dtype, memory_format=torch.contiguous_format) elif state["exp_avg_sq"].device != p.device: state["exp_avg_sq"] = state["exp_avg_sq"].to(device=p.device, non_blocking=True) params_with_grad.append(p) grads.append(p.grad) exp_avgs.append(state["exp_avg"]) scratches.append(state["scratch"]) if buf_count == 2: exp_avg_sqs.append(state["exp_avg_sq"]) if master_mode != "none": master_params.append(state["master_param"]) state["step"] += 1 curr_step = state["step"] step_val = int(curr_step.item() if isinstance(curr_step, torch.Tensor) else curr_step) steps.append(step_val) if not params_with_grad: continue lr = group["lr"] beta1, beta2 = group["betas"] eps = group["eps"] weight_decay = group["weight_decay"] tau = group["tau"] alpha_ss = group["alpha_ss"] cautious = group["cautious"] clamp_min = group["cautious_clamp_min"] use_sr = group["stochastic_rounding"] bound = group.get("bound", True) bound_type = group.get("bound_type", "radial") bound_ratio = group.get("bound_ratio", 0.03) exec_mode = group["execution"] fused_single_pass = group.get("fused_single_pass", False) ema_decay = group.get("ema_decay", 0.8) seed = group.get("seed", self.seed) all_accelerator = all( (p.is_cuda or (hasattr(p, "is_xpu") and p.is_xpu)) for p in params_with_grad ) deterministic_mode = False if hasattr(torch, "are_deterministic_algorithms_enabled"): deterministic_mode = torch.are_deterministic_algorithms_enabled() if exec_mode == "auto": if deterministic_mode: exec_mode = "foreach" elif HAS_TRITON and all_accelerator: exec_mode = "triton" elif hasattr(torch, "_foreach_mul_"): exec_mode = "foreach" else: exec_mode = "single" if exec_mode == "triton" and not (HAS_TRITON and all_accelerator): exec_mode = "foreach" has_master = master_mode != "none" m_list = master_params if has_master else None if exec_mode == "triton" and HAS_TRITON and all_accelerator: completed_indices: Set[int] = set() try: self._dispatch_triton( params_with_grad, grads, exp_avgs, exp_avg_sqs, m_list, steps, scratches, lr, beta1, beta2, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, buf_count, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed, completed_indices=completed_indices ) except Exception as e: logger.warning(f"Triton dispatch failed ({e}), falling back to foreach for uncompleted parameters.") uncompleted = [i for i in range(len(params_with_grad)) if i not in completed_indices] if uncompleted: for idx_u in uncompleted: scratches[idx_u].zero_() if buf_count == 2: scratches[idx_u][3] = 1.0 scratches[idx_u][4] = 1.0 sub_p = [params_with_grad[i] for i in uncompleted] sub_g = [grads[i] for i in uncompleted] sub_m = [exp_avgs[i] for i in uncompleted] sub_v = [exp_avg_sqs[i] for i in uncompleted] if buf_count == 2 else None sub_mp = [m_list[i] for i in uncompleted] if m_list is not None else None sub_st = [steps[i] for i in uncompleted] sub_sc = [scratches[i] for i in uncompleted] self._grouped_foreach_step( sub_p, sub_g, sub_m, sub_v, sub_mp, sub_st, sub_sc, lr, beta1, beta2, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, buf_count, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed ) elif exec_mode == "foreach": self._grouped_foreach_step( params_with_grad, grads, exp_avgs, exp_avg_sqs, m_list, steps, scratches, lr, beta1, beta2, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, buf_count, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed ) else: if buf_count == 2: self._single_step_v2( params_with_grad, grads, exp_avgs, exp_avg_sqs, m_list, steps, scratches, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed ) else: self._single_step_v1( params_with_grad, grads, exp_avgs, m_list, steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, seed ) return loss def _dispatch_triton( self, params, grads, exp_avgs, exp_avg_sqs, master_params, steps, scratches, lr, beta1, beta2, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, buf_count, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed: int = 1337, completed_indices: Optional[Set[int]] = None ): device_groups: Dict[torch.device, List[int]] = {} for idx, p in enumerate(params): dev = p.device if dev not in device_groups: device_groups[dev] = [] device_groups[dev].append(idx) for dev, indices in device_groups.items(): sub_p = [params[i] for i in indices] sub_g = [grads[i] for i in indices] sub_m = [exp_avgs[i] for i in indices] sub_st = [steps[i] for i in indices] sub_sc = [scratches[i] for i in indices] sub_mp = [master_params[i] for i in indices] if master_params is not None else None with _get_device_context(dev): if buf_count == 2: sub_v = [exp_avg_sqs[i] for i in indices] self._triton_step_v2( sub_p, sub_g, sub_m, sub_v, sub_mp, sub_st, sub_sc, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed, orig_indices=indices, completed_indices=completed_indices ) else: self._triton_step_v1( sub_p, sub_g, sub_m, sub_mp, sub_st, sub_sc, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, seed, orig_indices=indices, completed_indices=completed_indices ) def _grouped_foreach_step( self, params, grads, exp_avgs, exp_avg_sqs, master_params, steps, scratches, lr, beta1, beta2, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, buf_count, use_sr, bound, bound_type, bound_ratio, fused_single_pass=False, ema_decay=0.8, seed: int = 1337 ): groups: Dict[Tuple[torch.device, torch.dtype], List[int]] = {} for idx, p in enumerate(params): key = (p.device, p.dtype) if key not in groups: groups[key] = [] groups[key].append(idx) for (dev, dt), indices in groups.items(): sub_params = [params[i] for i in indices] sub_grads = [grads[i] for i in indices] sub_exp_avgs = [exp_avgs[i] for i in indices] sub_steps = [steps[i] for i in indices] sub_mp = [master_params[i] for i in indices] if master_params is not None else None sub_sc = [scratches[i] for i in indices] if scratches is not None else None with _get_device_context(dev): any_non_contig = any(not p.is_contiguous() for p in sub_params) or any(not g.is_contiguous() for g in sub_grads) if any_non_contig: if buf_count == 2: sub_exp_avg_sqs = [exp_avg_sqs[i] for i in indices] self._single_step_v2( sub_params, sub_grads, sub_exp_avgs, sub_exp_avg_sqs, sub_mp, sub_steps, sub_sc, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed ) else: self._single_step_v1( sub_params, sub_grads, sub_exp_avgs, sub_mp, sub_steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, seed ) else: if buf_count == 2: sub_exp_avg_sqs = [exp_avg_sqs[i] for i in indices] self._foreach_step_v2( sub_params, sub_grads, sub_exp_avgs, sub_exp_avg_sqs, sub_mp, sub_steps, sub_sc, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed ) else: self._foreach_step_v1( sub_params, sub_grads, sub_exp_avgs, sub_mp, sub_steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, seed ) def _triton_step_v2( self, params, grads, exp_avgs, exp_avg_sqs, master_params, steps, scratches, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound=True, bound_type="radial", bound_ratio=0.03, fused_single_pass=False, ema_decay=0.8, seed: int = 1337, orig_indices: Optional[List[int]] = None, completed_indices: Optional[Set[int]] = None ): BLOCK_SIZE = 1024 bound_mode = 0 if bound: bound_mode = 1 if bound_type == "radial" else 2 has_master = master_params is not None base_step = int(steps[0].item() if isinstance(steps[0], torch.Tensor) else steps[0]) bc1_nes_base = 1.0 - (beta1 ** (base_step + 1)) bc2_base = 1.0 - (beta2 ** base_step) c1_base = (bc1_nes_base * tau) / math.sqrt(bc2_base) c2_default = eps * bc1_nes_base * tau c2_fp16 = max(eps, 1e-4) * bc1_nes_base * tau for i in range(len(params)): p, grad, exp_avg, exp_avg_sq, step_raw, scratch_param = ( params[i], grads[i], exp_avgs[i], exp_avg_sqs[i], steps[i], scratches[i] ) step = int(step_raw.item() if isinstance(step_raw, torch.Tensor) else step_raw) pass1_executed = False try: is_orig_contig = p.is_contiguous() is_grad_contig = grad.is_contiguous() is_m_contig = exp_avg.is_contiguous() is_v_contig = exp_avg_sq.is_contiguous() p_contig = p if is_orig_contig else p.contiguous() grad_contig = grad if is_grad_contig else grad.contiguous() exp_avg_contig = exp_avg if is_m_contig else exp_avg.contiguous() exp_avg_sq_contig = exp_avg_sq if is_v_contig else exp_avg_sq.contiguous() if has_master: mp = master_params[i] is_mp_contig = mp.is_contiguous() mp_contig = mp if is_mp_contig else mp.contiguous() else: is_mp_contig = True mp_contig = p_contig if not (is_orig_contig and is_grad_contig and is_m_contig and is_v_contig and is_mp_contig): if not self._warned_non_contiguous: logger.warning( "LuminaV detected non-contiguous tensor buffer(s). An explicit contiguous copy " "was created for Triton kernel execution. To eliminate temporary VRAM allocations, " "ensure parameters and optimizer states are contiguous." ) self._warned_non_contiguous = True if step == base_step: c1 = c1_base c2 = c2_fp16 if p.dtype == torch.float16 else c2_default else: effective_eps = max(eps, 1e-4) if p.dtype == torch.float16 else eps bc1_nes = 1.0 - (beta1 ** (step + 1)) bc2 = 1.0 - (beta2 ** step) c1 = (bc1_nes * tau) / math.sqrt(bc2) c2 = effective_eps * bc1_nes * tau if p.dtype == torch.bfloat16: dtype_mode = 1 elif p.dtype == torch.float16: dtype_mode = 2 else: dtype_mode = 0 if exp_avg.dtype == torch.bfloat16: state_dtype_mode = 1 elif exp_avg.dtype == torch.float16: state_dtype_mode = 2 else: state_dtype_mode = 0 kernel_seed = (seed + step * 0x9E3779B9 + i * 10007) & 0x7FFFFFFF n_elements = p.numel() if n_elements >= 2_097_152: BLOCK_SIZE = 2048 num_warps = 8 else: BLOCK_SIZE = 1024 num_warps = 4 grid = (triton.cdiv(n_elements, BLOCK_SIZE),) if fused_single_pass and step > 1: _lumina_v2_fused_single_pass_kernel[grid]( p_contig, mp_contig, grad_contig, exp_avg_contig, exp_avg_sq_contig, scratch_param, n_elements, beta1, beta2, c1, c2, lr, weight_decay, bound_mode, bound_ratio, kernel_seed, dtype_mode, state_dtype_mode, use_sr, cautious, has_master, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps ) _lumina_v2_post_step_ema_kernel[(1,)]( scratch_param, n_elements, lr, clamp_min, ema_decay, step, bound_mode, bound_ratio, cautious, num_warps=1 ) elif not cautious and bound_mode != 1: _lumina_v2_single_pass_kernel[grid]( p_contig, mp_contig, grad_contig, exp_avg_contig, exp_avg_sq_contig, n_elements, beta1, beta2, c1, c2, lr, weight_decay, bound_mode, bound_ratio, kernel_seed, dtype_mode, state_dtype_mode, use_sr, has_master, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps ) else: scratch_param[:3].zero_() _lumina_v2_pass1_kernel[grid]( p_contig, mp_contig, grad_contig, exp_avg_contig, exp_avg_sq_contig, scratch_param, n_elements, beta1, beta2, c1, c2, kernel_seed, bound_mode, dtype_mode, state_dtype_mode, use_sr, cautious, has_master, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps ) pass1_executed = True _lumina_v2_pass2_kernel[grid]( p_contig, mp_contig, grad_contig, exp_avg_contig, exp_avg_sq_contig, scratch_param, n_elements, beta1, c1, c2, lr, weight_decay, clamp_min, bound_mode, bound_ratio, kernel_seed, dtype_mode, use_sr, cautious, has_master, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps ) if not is_orig_contig: p.copy_(p_contig) if not is_m_contig: exp_avg.copy_(exp_avg_contig) if not is_v_contig: exp_avg_sq.copy_(exp_avg_sq_contig) if has_master and not is_mp_contig: master_params[i].copy_(mp_contig) if completed_indices is not None and orig_indices is not None: completed_indices.add(orig_indices[i]) except Exception as e: logger.warning(f"Triton kernel failed for param {i}, falling back to single step: {e}") scratch_param.zero_() scratch_param[3] = 1.0 scratch_param[4] = 1.0 if pass1_executed: if not is_m_contig: exp_avg.copy_(exp_avg_contig) if not is_v_contig: exp_avg_sq.copy_(exp_avg_sq_contig) self._single_step_v2_apply_param( [p], [grad], [exp_avg], [exp_avg_sq], [master_params[i]] if has_master else None, [step], [scratch_param], lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed ) else: self._single_step_v2( [p], [grad], [exp_avg], [exp_avg_sq], [master_params[i]] if has_master else None, [step], [scratch_param], lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, fused_single_pass, ema_decay, seed ) if completed_indices is not None and orig_indices is not None: completed_indices.add(orig_indices[i]) def _triton_step_v1( self, params, grads, exp_avgs, master_params, steps, scratches, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound=True, bound_type="radial", bound_ratio=0.03, seed: int = 1337, orig_indices: Optional[List[int]] = None, completed_indices: Optional[Set[int]] = None ): BLOCK_SIZE = 1024 bound_mode = 0 if bound: bound_mode = 1 if bound_type == "radial" else 2 has_master = master_params is not None base_step = int(steps[0].item() if isinstance(steps[0], torch.Tensor) else steps[0]) bc1_nes_base = 1.0 - (beta1 ** (base_step + 1)) c2_default = eps * bc1_nes_base * tau c2_fp16 = max(eps, 1e-4) * bc1_nes_base * tau for i in range(len(params)): p, grad, exp_avg, step_raw, scratch_param = ( params[i], grads[i], exp_avgs[i], steps[i], scratches[i] ) step = int(step_raw.item() if isinstance(step_raw, torch.Tensor) else step_raw) pass1_executed = False try: is_orig_contig = p.is_contiguous() is_grad_contig = grad.is_contiguous() is_m_contig = exp_avg.is_contiguous() p_contig = p if is_orig_contig else p.contiguous() grad_contig = grad if is_grad_contig else grad.contiguous() exp_avg_contig = exp_avg if is_m_contig else exp_avg.contiguous() if has_master: mp = master_params[i] is_mp_contig = mp.is_contiguous() mp_contig = mp if is_mp_contig else mp.contiguous() else: is_mp_contig = True mp_contig = p_contig if not (is_orig_contig and is_grad_contig and is_m_contig and is_mp_contig): if not self._warned_non_contiguous: logger.warning( "LuminaV detected non-contiguous tensor buffer(s). An explicit contiguous copy " "was created for Triton kernel execution. To eliminate temporary VRAM allocations, " "ensure parameters and optimizer states are contiguous." ) self._warned_non_contiguous = True effective_eps = max(eps, 1e-4) if p.dtype == torch.float16 else eps if step == base_step: c2 = c2_fp16 if p.dtype == torch.float16 else c2_default else: bc1_nes = 1.0 - (beta1 ** (step + 1)) c2 = effective_eps * bc1_nes * tau if p.dtype == torch.bfloat16: dtype_mode = 1 elif p.dtype == torch.float16: dtype_mode = 2 else: dtype_mode = 0 if exp_avg.dtype == torch.bfloat16: state_dtype_mode = 1 elif exp_avg.dtype == torch.float16: state_dtype_mode = 2 else: state_dtype_mode = 0 kernel_seed = (seed + step * 0x9E3779B9 + i * 10007) & 0x7FFFFFFF n_elements = p.numel() if n_elements >= 2_097_152: BLOCK_SIZE = 2048 num_warps = 8 else: BLOCK_SIZE = 1024 num_warps = 4 grid = (triton.cdiv(n_elements, BLOCK_SIZE),) scratch_param[:4].zero_() _lumina_v1_pass1_kernel[grid]( p_contig, mp_contig, grad_contig, exp_avg_contig, scratch_param, n_elements, beta1, kernel_seed, bound_mode, state_dtype_mode, use_sr, has_master, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps ) pass1_executed = True _lumina_v1_pass2_kernel[grid]( grad_contig, exp_avg_contig, scratch_param, n_elements, beta1, tau, effective_eps, c2, alpha_ss, cautious, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps ) _lumina_v1_pass3_update_kernel[grid]( p_contig, mp_contig, grad_contig, exp_avg_contig, scratch_param, n_elements, beta1, tau, effective_eps, c2, alpha_ss, lr, weight_decay, clamp_min, bound_mode, bound_ratio, kernel_seed, dtype_mode, use_sr, cautious, has_master, BLOCK_SIZE=BLOCK_SIZE, num_warps=num_warps ) if not is_orig_contig: p.copy_(p_contig) if not is_m_contig: exp_avg.copy_(exp_avg_contig) if has_master and not is_mp_contig: master_params[i].copy_(mp_contig) if completed_indices is not None and orig_indices is not None: completed_indices.add(orig_indices[i]) except Exception as e: logger.warning(f"Triton kernel failed for param {i}, falling back to single step: {e}") scratch_param.zero_() if pass1_executed: if not is_m_contig: exp_avg.copy_(exp_avg_contig) self._single_step_v1_apply_param( [p], [grad], [exp_avg], [master_params[i]] if has_master else None, [step], lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, seed ) else: self._single_step_v1( [p], [grad], [exp_avg], [master_params[i]] if has_master else None, [step], lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound, bound_type, bound_ratio, seed ) if completed_indices is not None and orig_indices is not None: completed_indices.add(orig_indices[i]) def _single_step_v2( self, params, grads, exp_avgs, exp_avg_sqs, master_params, steps, scratches, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound=True, bound_type="radial", bound_ratio=0.03, fused_single_pass=False, ema_decay=0.8, seed: int = 1337 ): base_step = int(steps[0].item() if isinstance(steps[0], torch.Tensor) else steps[0]) bc1_nes_base = 1.0 - (beta1 ** (base_step + 1)) bc2_base = 1.0 - (beta2 ** base_step) c1_base = (bc1_nes_base * tau) / math.sqrt(bc2_base) c2_default = eps * bc1_nes_base * tau c2_fp16 = max(eps, 1e-4) * bc1_nes_base * tau for i in range(len(params)): p, grad, exp_avg, exp_avg_sq, step_raw = params[i], grads[i], exp_avgs[i], exp_avg_sqs[i], steps[i] step = int(step_raw.item() if isinstance(step_raw, torch.Tensor) else step_raw) if step == base_step: c1 = c1_base c2 = c2_fp16 if p.dtype == torch.float16 else c2_default else: effective_eps = max(eps, 1e-4) if p.dtype == torch.float16 else eps bc1_nes = 1.0 - (beta1 ** (step + 1)) bc2 = 1.0 - (beta2 ** step) c1 = (bc1_nes * tau) / math.sqrt(bc2) c2 = effective_eps * bc1_nes * tau grad_clean = torch.nan_to_num(grad, nan=0.0, posinf=0.0, neginf=0.0) grad_f = grad_clean.to(dtype=exp_avg.dtype) exp_avg.mul_(beta1).add_(grad_f, alpha=1.0 - beta1) nes_m = torch.mul(exp_avg, beta1).add_(grad_f, alpha=1.0 - beta1) grad_diff = grad_f - exp_avg if exp_avg_sq.dtype in (torch.float16, torch.bfloat16): diff_f32 = grad_diff.float() v_f32 = exp_avg_sq.float().mul_(beta2).addcmul_(diff_f32, diff_f32, value=1.0 - beta2) if exp_avg_sq.dtype == torch.float16: v_f32 = v_f32.clamp_(0.0, 65000.0) exp_avg_sq.copy_(v_f32.to(dtype=exp_avg_sq.dtype)) else: exp_avg_sq.mul_(beta2).addcmul_(grad_diff, grad_diff, value=1.0 - beta2) sigma = exp_avg_sq.float().sqrt().mul_(c1).add_(c2) u = torch.tanh(nes_m.float() / sigma.clamp_min(1e-12)) ref_p = master_params[i] if master_params is not None else p sc = scratches[i] if scratches is not None else None if fused_single_pass and sc is not None and step > 1: cautious_scale = sc[3] bound_scale = sc[4] if cautious: m_mask = (u * grad_clean.float() > 0).float() delta_theta = u * m_mask * cautious_scale else: m_mask = None delta_theta = u step_val = lr * delta_theta.float() if bound and bound_ratio > 0.0: if bound_type == "radial": step_val = step_val * bound_scale elif bound_type == "coordinate": p_ref = ref_p.float() B = bound_ratio * torch.sqrt(p_ref.square() + 1e-4) step_val = B * torch.tanh(step_val / B) if p.dtype == torch.float16: p_ref = ref_p.float() b_safe = torch.clamp_min(torch.abs(p_ref) * 0.01, 1.5e-4) step_val = b_safe * torch.tanh(step_val / b_safe) if cautious and m_mask is not None: m_bar = m_mask.mean().clamp_(min=clamp_min, max=1.0).to(torch.float32) target_cautious = 1.0 / m_bar new_cautious = ema_decay * cautious_scale + (1.0 - ema_decay) * target_cautious sc[3].copy_(new_cautious) if bound and bound_type == "radial" and bound_ratio > 0.0: c_eff = sc[3] if cautious else torch.tensor(1.0, device=p.device, dtype=torch.float32) u_eff = (u * m_mask) if (cautious and m_mask is not None) else u norm_delta = torch.linalg.vector_norm(u_eff.float()) * c_eff norm_p = torch.linalg.vector_norm(ref_p.float()) R = bound_ratio * torch.clamp_min(norm_p, 1.0) r = (lr * norm_delta) / R target_bound = torch.where(r < 1e-4, torch.ones_like(r), torch.tanh(r) / torch.clamp_min(r, 1e-8)) new_bound = ema_decay * bound_scale + (1.0 - ema_decay) * target_bound sc[4].copy_(new_bound) update = delta_theta if lr == 0.0 else (step_val / lr) else: if cautious: mask = (u * grad_clean.float() > 0).float() m_bar = mask.mean().clamp_(min=clamp_min, max=1.0) cautious_scale = 1.0 / m_bar update = u.mul(mask).mul(cautious_scale) if sc is not None: sc[3].copy_(cautious_scale.to(torch.float32)) else: update = u if bound and bound_type == "radial" and bound_ratio > 0.0: norm_u = torch.linalg.vector_norm(update.float()) norm_p = torch.linalg.vector_norm(ref_p.float()) R = bound_ratio * torch.clamp_min(norm_p, 1.0) r = (lr * norm_u) / R b_scale = torch.where(r < 1e-4, torch.ones_like(r), torch.tanh(r) / torch.clamp_min(r, 1e-8)) if sc is not None: sc[4].copy_(b_scale.to(torch.float32)) if bound or p.dtype == torch.float16: update = _apply_bound(update, ref_p, lr, bound, bound_type, bound_ratio, target_dtype=p.dtype) if master_params is not None: mp = master_params[i] if weight_decay != 0.0: mp.mul_(1.0 - lr * weight_decay) mp.add_(update, alpha=-lr) p.copy_(mp.to(dtype=p.dtype)) else: _sr_update(p, update, lr, weight_decay, use_sr, generator=self._get_generator(p.device, seed=seed)) def _single_step_v2_apply_param( self, params, grads, exp_avgs, exp_avg_sqs, master_params, steps, scratches, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound=True, bound_type="radial", bound_ratio=0.03, fused_single_pass=False, ema_decay=0.8, seed: int = 1337 ): base_step = int(steps[0].item() if isinstance(steps[0], torch.Tensor) else steps[0]) bc1_nes_base = 1.0 - (beta1 ** (base_step + 1)) bc2_base = 1.0 - (beta2 ** base_step) c1_base = (bc1_nes_base * tau) / math.sqrt(bc2_base) c2_default = eps * bc1_nes_base * tau c2_fp16 = max(eps, 1e-4) * bc1_nes_base * tau for i in range(len(params)): p, grad, exp_avg, exp_avg_sq, step_raw = params[i], grads[i], exp_avgs[i], exp_avg_sqs[i], steps[i] step = int(step_raw.item() if isinstance(step_raw, torch.Tensor) else step_raw) if step == base_step: c1 = c1_base c2 = c2_fp16 if p.dtype == torch.float16 else c2_default else: effective_eps = max(eps, 1e-4) if p.dtype == torch.float16 else eps bc1_nes = 1.0 - (beta1 ** (step + 1)) bc2 = 1.0 - (beta2 ** step) c1 = (bc1_nes * tau) / math.sqrt(bc2) c2 = effective_eps * bc1_nes * tau grad_clean = torch.nan_to_num(grad, nan=0.0, posinf=0.0, neginf=0.0) grad_f = grad_clean.to(dtype=exp_avg.dtype) nes_m = torch.mul(exp_avg, beta1).add_(grad_f, alpha=1.0 - beta1) sigma = exp_avg_sq.float().sqrt().mul_(c1).add_(c2) update = torch.tanh(nes_m.float() / sigma.clamp_min(1e-12)) ref_p = master_params[i] if master_params is not None else p sc = scratches[i] if scratches is not None else None if fused_single_pass and sc is not None and step > 1: cautious_scale = sc[3] bound_scale = sc[4] if cautious: m_mask = (update * grad_clean.float() > 0).float() delta_theta = update * m_mask * cautious_scale else: m_mask = None delta_theta = update step_val = lr * delta_theta.float() if bound and bound_ratio > 0.0: if bound_type == "radial": step_val = step_val * bound_scale elif bound_type == "coordinate": p_ref = ref_p.float() B = bound_ratio * torch.sqrt(p_ref.square() + 1e-4) step_val = B * torch.tanh(step_val / B) if p.dtype == torch.float16: p_ref = ref_p.float() b_safe = torch.clamp_min(torch.abs(p_ref) * 0.01, 1.5e-4) step_val = b_safe * torch.tanh(step_val / b_safe) if cautious and m_mask is not None: m_bar = m_mask.mean().clamp_(min=clamp_min, max=1.0).to(torch.float32) target_cautious = 1.0 / m_bar new_cautious = ema_decay * cautious_scale + (1.0 - ema_decay) * target_cautious sc[3].copy_(new_cautious) if bound and bound_type == "radial" and bound_ratio > 0.0: c_eff = sc[3] if cautious else torch.tensor(1.0, device=p.device, dtype=torch.float32) u_eff = (update * m_mask) if (cautious and m_mask is not None) else update norm_delta = torch.linalg.vector_norm(u_eff.float()) * c_eff norm_p = torch.linalg.vector_norm(ref_p.float()) R = bound_ratio * torch.clamp_min(norm_p, 1.0) r = (lr * norm_delta) / R target_bound = torch.where(r < 1e-4, torch.ones_like(r), torch.tanh(r) / torch.clamp_min(r, 1e-8)) new_bound = ema_decay * bound_scale + (1.0 - ema_decay) * target_bound sc[4].copy_(new_bound) update = delta_theta if lr == 0.0 else (step_val / lr) else: if cautious: mask = (update * grad_clean.float() > 0).float() m_bar = mask.mean().clamp_(min=clamp_min, max=1.0) cautious_scale = 1.0 / m_bar update = update.mul(mask).mul(cautious_scale) if sc is not None: sc[3].copy_(cautious_scale.to(torch.float32)) if bound and bound_type == "radial" and bound_ratio > 0.0: norm_u = torch.linalg.vector_norm(update.float()) norm_p = torch.linalg.vector_norm(ref_p.float()) R = bound_ratio * torch.clamp_min(norm_p, 1.0) r = (lr * norm_u) / R b_scale = torch.where(r < 1e-4, torch.ones_like(r), torch.tanh(r) / torch.clamp_min(r, 1e-8)) if sc is not None: sc[4].copy_(b_scale.to(torch.float32)) if bound or p.dtype == torch.float16: update = _apply_bound(update, ref_p, lr, bound, bound_type, bound_ratio, target_dtype=p.dtype) if master_params is not None: mp = master_params[i] if weight_decay != 0.0: mp.mul_(1.0 - lr * weight_decay) mp.add_(update, alpha=-lr) p.copy_(mp.to(dtype=p.dtype)) else: _sr_update(p, update, lr, weight_decay, use_sr, generator=self._get_generator(p.device, seed=seed)) def _single_step_v1( self, params, grads, exp_avgs, master_params, steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound=True, bound_type="radial", bound_ratio=0.03, seed: int = 1337 ): base_step = int(steps[0].item() if isinstance(steps[0], torch.Tensor) else steps[0]) bc1_nes_base = 1.0 - (beta1 ** (base_step + 1)) c2_default = eps * bc1_nes_base * tau c2_fp16 = max(eps, 1e-4) * bc1_nes_base * tau for i in range(len(params)): p, grad, exp_avg, step_raw = params[i], grads[i], exp_avgs[i], steps[i] step = int(step_raw.item() if isinstance(step_raw, torch.Tensor) else step_raw) if step == base_step: c2 = c2_fp16 if p.dtype == torch.float16 else c2_default else: effective_eps = max(eps, 1e-4) if p.dtype == torch.float16 else eps bc1_nes = 1.0 - (beta1 ** (step + 1)) c2 = effective_eps * bc1_nes * tau grad_clean = torch.nan_to_num(grad, nan=0.0, posinf=0.0, neginf=0.0) grad_f = grad_clean.to(dtype=exp_avg.dtype) exp_avg.mul_(beta1).add_(grad_f, alpha=1.0 - beta1) nes_m = torch.mul(exp_avg, beta1).add_(grad_f, alpha=1.0 - beta1) effective_eps = max(eps, 1e-4) if p.dtype == torch.float16 else eps rms = torch.sqrt(nes_m.float().square().mean() + effective_eps) sigma = rms * tau + c2 z = nes_m.float() / sigma.clamp_min(1e-12) z_soft = z / (1.0 + alpha_ss * torch.abs(z)) update = torch.tanh(z_soft) if cautious: mask = (update * grad_clean.float() > 0).float() mask_scale = mask.mean().clamp_(min=clamp_min, max=1.0) update = update.mul(mask).div(mask_scale) ref_p = master_params[i] if master_params is not None else p if bound or p.dtype == torch.float16: update = _apply_bound(update, ref_p, lr, bound, bound_type, bound_ratio, target_dtype=p.dtype) if master_params is not None: mp = master_params[i] if weight_decay != 0.0: mp.mul_(1.0 - lr * weight_decay) mp.add_(update, alpha=-lr) p.copy_(mp.to(dtype=p.dtype)) else: _sr_update(p, update, lr, weight_decay, use_sr, generator=self._get_generator(p.device, seed=seed)) def _single_step_v1_apply_param( self, params, grads, exp_avgs, master_params, steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound=True, bound_type="radial", bound_ratio=0.03, seed: int = 1337 ): base_step = int(steps[0].item() if isinstance(steps[0], torch.Tensor) else steps[0]) bc1_nes_base = 1.0 - (beta1 ** (base_step + 1)) c2_default = eps * bc1_nes_base * tau c2_fp16 = max(eps, 1e-4) * bc1_nes_base * tau for i in range(len(params)): p, grad, exp_avg, step_raw = params[i], grads[i], exp_avgs[i], steps[i] step = int(step_raw.item() if isinstance(step_raw, torch.Tensor) else step_raw) if step == base_step: c2 = c2_fp16 if p.dtype == torch.float16 else c2_default else: effective_eps = max(eps, 1e-4) if p.dtype == torch.float16 else eps bc1_nes = 1.0 - (beta1 ** (step + 1)) c2 = effective_eps * bc1_nes * tau grad_clean = torch.nan_to_num(grad, nan=0.0, posinf=0.0, neginf=0.0) grad_f = grad_clean.to(dtype=exp_avg.dtype) nes_m = torch.mul(exp_avg, beta1).add_(grad_f, alpha=1.0 - beta1) effective_eps = max(eps, 1e-4) if p.dtype == torch.float16 else eps rms = torch.sqrt(nes_m.float().square().mean() + effective_eps) sigma = rms * tau + c2 z = nes_m.float() / sigma.clamp_min(1e-12) z_soft = z / (1.0 + alpha_ss * torch.abs(z)) update = torch.tanh(z_soft) if cautious: mask = (update * grad_clean.float() > 0).float() mask_scale = mask.mean().clamp_(min=clamp_min, max=1.0) update = update.mul(mask).div(mask_scale) ref_p = master_params[i] if master_params is not None else p if bound or p.dtype == torch.float16: update = _apply_bound(update, ref_p, lr, bound, bound_type, bound_ratio, target_dtype=p.dtype) if master_params is not None: mp = master_params[i] if weight_decay != 0.0: mp.mul_(1.0 - lr * weight_decay) mp.add_(update, alpha=-lr) p.copy_(master_params[i].to(dtype=p.dtype)) else: _sr_update(p, update, lr, weight_decay, use_sr, generator=self._get_generator(p.device, seed=seed)) def _foreach_step_v2( self, params, grads, exp_avgs, exp_avg_sqs, master_params, steps, scratches, lr, beta1, beta2, eps, weight_decay, tau, cautious, clamp_min, use_sr, bound=True, bound_type="radial", bound_ratio=0.03, fused_single_pass=False, ema_decay=0.8, seed: int = 1337 ): grads_clean = [torch.nan_to_num(g, nan=0.0, posinf=0.0, neginf=0.0) for g in grads] exp_avg_dtype = exp_avgs[0].dtype grad_dtype = grads_clean[0].dtype if exp_avg_dtype != grad_dtype: grads_for_m = [g.to(dtype=exp_avg_dtype) for g in grads_clean] else: grads_for_m = grads_clean torch._foreach_mul_(exp_avgs, beta1) torch._foreach_add_(exp_avgs, grads_for_m, alpha=1.0 - beta1) nes_m_list = torch._foreach_mul(exp_avgs, beta1) torch._foreach_add_(nes_m_list, grads_for_m, alpha=1.0 - beta1) grad_diff_list = torch._foreach_sub(grads_for_m, exp_avgs) is_low_prec = params[0].dtype in (torch.float16, torch.bfloat16) if is_low_prec and exp_avg_sqs[0].dtype in (torch.float16, torch.bfloat16): target_sq_dtype = exp_avg_sqs[0].dtype for i in range(len(params)): d_f32 = grad_diff_list[i].float() v_f32 = exp_avg_sqs[i].float().mul_(beta2).addcmul_(d_f32, d_f32, value=1.0 - beta2) if target_sq_dtype == torch.float16: v_f32 = v_f32.clamp_(0.0, 65000.0) exp_avg_sqs[i].copy_(v_f32.to(dtype=target_sq_dtype)) else: torch._foreach_mul_(exp_avg_sqs, beta2) torch._foreach_addcmul_(exp_avg_sqs, grad_diff_list, grad_diff_list, value=1.0 - beta2) if exp_avg_sqs[0].dtype == torch.float16: for v_tensor in exp_avg_sqs: v_tensor.clamp_(0.0, 65000.0) is_fp16 = params[0].dtype == torch.float16 effective_eps = max(eps, 1e-4) if is_fp16 else eps steps_int = [int(st.item() if isinstance(st, torch.Tensor) else st) for st in steps] first_step = steps_int[0] if all(st == first_step for st in steps_int): bc1_nes = 1.0 - (beta1 ** (first_step + 1)) bc2 = 1.0 - (beta2 ** first_step) c1_val = (bc1_nes * tau) / math.sqrt(bc2) c2_val = effective_eps * bc1_nes * tau c1_list = [c1_val] * len(steps_int) c2_list = [c2_val] * len(steps_int) else: bias_correction1_nes = [1.0 - (beta1 ** (st + 1)) for st in steps_int] bias_correction2 = [1.0 - (beta2 ** st) for st in steps_int] c1_list = [(bc1 * tau) / math.sqrt(bc2) for bc1, bc2 in zip(bias_correction1_nes, bias_correction2)] c2_list = [effective_eps * bc1 * tau for bc1 in bias_correction1_nes] updates = [] for i in range(len(params)): sigma = exp_avg_sqs[i].float().sqrt().mul_(c1_list[i]).add_(c2_list[i]) u = torch.tanh(nes_m_list[i].float() / sigma.clamp_min(1e-12)) ref_p = master_params[i] if master_params is not None else params[i] sc = scratches[i] if scratches is not None else None st = steps_int[i] if fused_single_pass and sc is not None and st > 1: cautious_scale = sc[3] bound_scale = sc[4] if cautious: m_mask = (u * grads_clean[i].float() > 0).float() delta_theta = u * m_mask * cautious_scale else: m_mask = None delta_theta = u step_val = lr * delta_theta.float() if bound and bound_ratio > 0.0: if bound_type == "radial": step_val = step_val * bound_scale elif bound_type == "coordinate": p_ref = ref_p.float() B = bound_ratio * torch.sqrt(p_ref.square() + 1e-4) step_val = B * torch.tanh(step_val / B) if params[i].dtype == torch.float16: p_ref = ref_p.float() b_safe = torch.clamp_min(torch.abs(p_ref) * 0.01, 1.5e-4) step_val = b_safe * torch.tanh(step_val / b_safe) if cautious and m_mask is not None: m_bar = m_mask.mean().clamp_(min=clamp_min, max=1.0).to(torch.float32) target_cautious = 1.0 / m_bar new_cautious = ema_decay * cautious_scale + (1.0 - ema_decay) * target_cautious sc[3].copy_(new_cautious) if bound and bound_type == "radial" and bound_ratio > 0.0: c_eff = sc[3] if cautious else torch.tensor(1.0, device=params[i].device, dtype=torch.float32) u_eff = (u * m_mask) if (cautious and m_mask is not None) else u norm_delta = torch.linalg.vector_norm(u_eff.float()) * c_eff norm_p = torch.linalg.vector_norm(ref_p.float()) R = bound_ratio * torch.clamp_min(norm_p, 1.0) r = (lr * norm_delta) / R target_bound = torch.where(r < 1e-4, torch.ones_like(r), torch.tanh(r) / torch.clamp_min(r, 1e-8)) new_bound = ema_decay * bound_scale + (1.0 - ema_decay) * target_bound sc[4].copy_(new_bound) u_final = delta_theta if lr == 0.0 else (step_val / lr) else: if cautious: mask = (u * grads_clean[i].float() > 0).float() m_bar = mask.mean().clamp_(min=clamp_min, max=1.0) cautious_scale = 1.0 / m_bar u = u.mul(mask).mul(cautious_scale) if sc is not None: sc[3].copy_(cautious_scale.to(torch.float32)) if bound and bound_type == "radial" and bound_ratio > 0.0: norm_u = torch.linalg.vector_norm(u.float()) norm_p = torch.linalg.vector_norm(ref_p.float()) R = bound_ratio * torch.clamp_min(norm_p, 1.0) r = (lr * norm_u) / R b_scale = torch.where(r < 1e-4, torch.ones_like(r), torch.tanh(r) / torch.clamp_min(r, 1e-8)) if sc is not None: sc[4].copy_(b_scale.to(torch.float32)) if bound or params[i].dtype == torch.float16: u = _apply_bound(u, ref_p, lr, bound, bound_type, bound_ratio, target_dtype=params[i].dtype) u_final = u updates.append(u_final) if master_params is not None: if weight_decay != 0.0: torch._foreach_mul_(master_params, 1.0 - lr * weight_decay) updates_f32 = [up.float() for up in updates] torch._foreach_add_(master_params, updates_f32, alpha=-lr) for i in range(len(params)): params[i].copy_(master_params[i].to(dtype=params[i].dtype)) else: if use_sr and params[0].dtype in (torch.bfloat16, torch.float16): gen = self._get_generator(params[0].device, seed=seed) for i in range(len(params)): _sr_update(params[i], updates[i], lr, weight_decay, use_sr=True, generator=gen) elif params[0].dtype in (torch.bfloat16, torch.float16): for i in range(len(params)): _sr_update(params[i], updates[i], lr, weight_decay, use_sr=False, generator=None) else: updates_cast = [up.to(dtype=params[0].dtype) for up in updates] if weight_decay != 0.0: torch._foreach_mul_(params, 1.0 - lr * weight_decay) torch._foreach_add_(params, updates_cast, alpha=-lr) def _foreach_step_v1( self, params, grads, exp_avgs, master_params, steps, lr, beta1, eps, weight_decay, tau, alpha_ss, cautious, clamp_min, use_sr, bound=True, bound_type="radial", bound_ratio=0.03, seed: int = 1337 ): grads_clean = [torch.nan_to_num(g, nan=0.0, posinf=0.0, neginf=0.0) for g in grads] exp_avg_dtype = exp_avgs[0].dtype grad_dtype = grads_clean[0].dtype if exp_avg_dtype != grad_dtype: grads_for_m = [g.to(dtype=exp_avg_dtype) for g in grads_clean] else: grads_for_m = grads_clean torch._foreach_mul_(exp_avgs, beta1) torch._foreach_add_(exp_avgs, grads_for_m, alpha=1.0 - beta1) nes_m_list = torch._foreach_mul(exp_avgs, beta1) torch._foreach_add_(nes_m_list, grads_for_m, alpha=1.0 - beta1) is_fp16 = params[0].dtype == torch.float16 effective_eps = max(eps, 1e-4) if is_fp16 else eps steps_int = [int(st.item() if isinstance(st, torch.Tensor) else st) for st in steps] first_step = steps_int[0] if all(st == first_step for st in steps_int): bc1_nes = 1.0 - (beta1 ** (first_step + 1)) c2_val = effective_eps * bc1_nes * tau c2_list = [c2_val] * len(steps_int) else: bias_correction1_nes = [1.0 - (beta1 ** (st + 1)) for st in steps_int] c2_list = [effective_eps * bc1 * tau for bc1 in bias_correction1_nes] updates = [] for i in range(len(params)): m = nes_m_list[i] rms = torch.sqrt(m.float().square().mean() + effective_eps) sigma = rms * tau + c2_list[i] z = m.float() / sigma.clamp_min(1e-12) z_soft = z / (1.0 + alpha_ss * torch.abs(z)) u = torch.tanh(z_soft) if cautious: mask = (u * grads_clean[i].float() > 0).float() scale = mask.mean().clamp_(min=clamp_min, max=1.0) u = u.mul(mask).div(scale) ref_p = master_params[i] if master_params is not None else params[i] if bound or params[i].dtype == torch.float16: u = _apply_bound(u, ref_p, lr, bound, bound_type, bound_ratio, target_dtype=params[i].dtype) updates.append(u) if master_params is not None: if weight_decay != 0.0: torch._foreach_mul_(master_params, 1.0 - lr * weight_decay) updates_f32 = [up.float() for up in updates] torch._foreach_add_(master_params, updates_f32, alpha=-lr) for i in range(len(params)): params[i].copy_(master_params[i].to(dtype=params[i].dtype)) else: if use_sr and params[0].dtype in (torch.bfloat16, torch.float16): gen = self._get_generator(params[0].device, seed=seed) for i in range(len(params)): _sr_update(params[i], updates[i], lr, weight_decay, use_sr=True, generator=gen) elif params[0].dtype in (torch.bfloat16, torch.float16): for i in range(len(params)): _sr_update(params[i], updates[i], lr, weight_decay, use_sr=False, generator=None) else: updates_cast = [up.to(dtype=params[0].dtype) for up in updates] if weight_decay != 0.0: torch._foreach_mul_(params, 1.0 - lr * weight_decay) torch._foreach_add_(params, updates_cast, alpha=-lr)