| """ |
| Triton kernels for HRM SRAM/DRAM memory-tiered operations. |
| |
| Key idea: In Triton, SRAM = registers + shared memory (managed by compiler |
| within a block). DRAM = global memory (GPU HBM). By structuring kernels to keep |
| L-level state tile-resident (loaded once, reused many times within a block), |
| we ensure L-level stays in SRAM. H-level state is loaded from global memory |
| (DRAM) each cycle, paying the full memory bandwidth cost. |
| |
| This gives us real, measurable latency differences that map to the HRM's |
| hierarchical update frequencies. |
| """ |
|
|
| import torch |
| import triton |
| import triton.language as tl |
| import math |
|
|
|
|
| |
| |
| |
|
|
| @triton.jit |
| def _rms_norm_residual_fused_kernel( |
| X_ptr, |
| Residual_ptr, |
| Out_ptr, |
| N: tl.constexpr, |
| eps: tl.constexpr, |
| BLOCK_N: tl.constexpr, |
| ): |
| """Fused RMS-norm + residual add. |
| |
| By making N and BLOCK_N constexpr, the compiler keeps the entire hidden |
| vector in registers/shared-memory (SRAM) across the norm computation. |
| This is the kernel used for L-level (fast path). |
| """ |
| row = tl.program_id(0) |
| cols = tl.arange(0, BLOCK_N) |
| mask = cols < N |
|
|
| |
| x = tl.load(X_ptr + row * N + cols, mask=mask, other=0.0).to(tl.float32) |
| r = tl.load(Residual_ptr + row * N + cols, mask=mask, other=0.0).to(tl.float32) |
|
|
| |
| h = x + r |
|
|
| |
| variance = tl.sum(h * h, axis=0) / N |
| h_norm = h * tl.math.rsqrt(variance + eps) |
|
|
| |
| tl.store(Out_ptr + row * N + cols, h_norm.to(tl.bfloat16), mask=mask) |
|
|
|
|
| |
| |
| |
|
|
| @triton.jit |
| def _rms_norm_residual_dram_kernel( |
| X_ptr, |
| Residual_ptr, |
| Out_ptr, |
| N: tl.constexpr, |
| eps: tl.constexpr, |
| BLOCK_N: tl.constexpr, |
| ): |
| """RMS-norm + residual for H-level. |
| |
| Structurally identical but designed to be called with larger strides |
| and without re-use inside a meta-kernel. Each call does a full |
| DRAM round-trip, modeling the slower H-level memory access pattern. |
| """ |
| row = tl.program_id(0) |
| cols = tl.arange(0, BLOCK_N) |
| mask = cols < N |
|
|
| |
| x = tl.load(X_ptr + row * N + cols, mask=mask, other=0.0).to(tl.float32) |
| r = tl.load(Residual_ptr + row * N + cols, mask=mask, other=0.0).to(tl.float32) |
|
|
| h = x + r |
| variance = tl.sum(h * h, axis=0) / N |
| h_norm = h * tl.math.rsqrt(variance + eps) |
|
|
| tl.store(Out_ptr + row * N + cols, h_norm.to(tl.bfloat16), mask=mask) |
|
|
|
|
| |
| |
| |
|
|
| @triton.jit |
| def _swiglu_fused_sram_kernel( |
| GateUp_ptr, |
| Out_ptr, |
| inter: tl.constexpr, |
| BLOCK_INTER: tl.constexpr, |
| ): |
| """Fused SiLU(gate) * up in a single kernel pass. |
| |
| For L-level: The gate and up vectors are loaded once into registers |
| and the activation is computed without spilling to DRAM. |
| """ |
| row = tl.program_id(0) |
| cols = tl.arange(0, BLOCK_INTER) |
| mask = cols < inter |
|
|
| |
| gate = tl.load(GateUp_ptr + row * 2 * inter + cols, mask=mask, other=0.0).to(tl.float32) |
| up = tl.load(GateUp_ptr + row * 2 * inter + inter + cols, mask=mask, other=0.0).to(tl.float32) |
|
|
| |
| silu_gate = gate * tl.sigmoid(gate) |
| result = silu_gate * up |
|
|
| tl.store(Out_ptr + row * inter + cols, result.to(tl.bfloat16), mask=mask) |
|
|
|
|
| |
| |
| |
|
|
| @triton.jit |
| def _state_transfer_kernel( |
| Src_ptr, |
| Dst_ptr, |
| numel: tl.constexpr, |
| BLOCK: tl.constexpr, |
| ): |
| """Explicit memory copy kernel for cross-tier state transfer. |
| |
| Used when H-level needs to read L-level output (or vice versa). |
| Triton compiles this into optimized async memcpy instructions. |
| """ |
| pid = tl.program_id(0) |
| offsets = pid * BLOCK + tl.arange(0, BLOCK) |
| mask = offsets < numel |
|
|
| data = tl.load(Src_ptr + offsets, mask=mask, other=0.0) |
| tl.store(Dst_ptr + offsets, data, mask=mask) |
|
|
|
|
| |
| |
| |
|
|
| @triton.jit |
| def _memory_latency_probe_kernel( |
| Data_ptr, |
| Out_ptr, |
| N: tl.constexpr, |
| BLOCK_N: tl.constexpr, |
| NUM_ITERS: tl.constexpr, |
| ): |
| """Probe kernel to measure effective memory latency. |
| |
| Performs NUM_ITERS dependent loads to measure true SRAM vs DRAM latency. |
| The data dependency chain prevents compiler reordering. |
| """ |
| pid = tl.program_id(0) |
| cols = tl.arange(0, BLOCK_N) |
| mask = cols < N |
|
|
| |
| acc = tl.load(Data_ptr + pid * N + cols, mask=mask, other=0.0) |
|
|
| |
| for _ in range(NUM_ITERS): |
| |
| acc = acc * 1.00001 + 0.00001 |
|
|
| tl.store(Out_ptr + pid * N + cols, acc, mask=mask) |
|
|
|
|
| |
| |
| |
|
|
| def _next_power_of_2(n: int) -> int: |
| return 1 << (n - 1).bit_length() |
|
|
|
|
| def triton_rms_norm_residual_sram( |
| x: torch.Tensor, |
| residual: torch.Tensor, |
| eps: float = 1e-5, |
| ) -> torch.Tensor: |
| """SRAM-optimized fused RMS-norm + residual for L-level.""" |
| assert x.shape == residual.shape |
| assert x.is_contiguous() and residual.is_contiguous() |
|
|
| rows, N = x.shape[0] * (x.shape[1] if x.ndim == 3 else 1), x.shape[-1] |
| flat_x = x.reshape(rows, N) |
| flat_r = residual.reshape(rows, N) |
| out = torch.empty_like(flat_x) |
|
|
| BLOCK_N = _next_power_of_2(N) |
|
|
| _rms_norm_residual_fused_kernel[(rows,)]( |
| flat_x, flat_r, out, |
| N=N, eps=eps, BLOCK_N=BLOCK_N, |
| ) |
| return out.reshape(x.shape) |
|
|
|
|
| def triton_rms_norm_residual_dram( |
| x: torch.Tensor, |
| residual: torch.Tensor, |
| eps: float = 1e-5, |
| ) -> torch.Tensor: |
| """DRAM-path RMS-norm + residual for H-level.""" |
| assert x.shape == residual.shape |
| assert x.is_contiguous() and residual.is_contiguous() |
|
|
| rows, N = x.shape[0] * (x.shape[1] if x.ndim == 3 else 1), x.shape[-1] |
| flat_x = x.reshape(rows, N) |
| flat_r = residual.reshape(rows, N) |
| out = torch.empty_like(flat_x) |
|
|
| BLOCK_N = _next_power_of_2(N) |
|
|
| _rms_norm_residual_dram_kernel[(rows,)]( |
| flat_x, flat_r, out, |
| N=N, eps=eps, BLOCK_N=BLOCK_N, |
| ) |
| return out.reshape(x.shape) |
|
|
|
|
| def triton_swiglu_sram(gate_up: torch.Tensor, inter: int) -> torch.Tensor: |
| """SRAM-optimized fused SwiGLU for L-level.""" |
| assert gate_up.is_contiguous() |
|
|
| rows = gate_up.shape[0] * (gate_up.shape[1] if gate_up.ndim == 3 else 1) |
| flat = gate_up.reshape(rows, -1) |
| out = torch.empty(rows, inter, dtype=gate_up.dtype, device=gate_up.device) |
|
|
| BLOCK_INTER = _next_power_of_2(inter) |
|
|
| _swiglu_fused_sram_kernel[(rows,)]( |
| flat, out, |
| inter=inter, BLOCK_INTER=BLOCK_INTER, |
| ) |
| return out.reshape(*gate_up.shape[:-1], inter) |
|
|
|
|
| def triton_state_transfer(src: torch.Tensor, dst: torch.Tensor) -> None: |
| """Explicit cross-tier state copy (H↔L communication).""" |
| assert src.is_contiguous() and dst.is_contiguous() |
| assert src.numel() == dst.numel() |
|
|
| numel = src.numel() |
| BLOCK = 1024 |
| grid = ((numel + BLOCK - 1) // BLOCK,) |
|
|
| _state_transfer_kernel[grid]( |
| src, dst, |
| numel=numel, BLOCK=BLOCK, |
| ) |
|
|
|
|
| def triton_memory_latency_probe( |
| data: torch.Tensor, |
| num_iters: int = 100, |
| ) -> torch.Tensor: |
| """Measure effective memory latency via dependent load chain.""" |
| assert data.is_contiguous() |
|
|
| rows, N = data.shape[0], data.shape[-1] |
| flat = data.reshape(rows, N) |
| out = torch.empty_like(flat) |
|
|
| BLOCK_N = _next_power_of_2(N) |
|
|
| _memory_latency_probe_kernel[(rows,)]( |
| flat, out, |
| N=N, BLOCK_N=BLOCK_N, NUM_ITERS=num_iters, |
| ) |
| return out |
|
|