h3-blackwell-runtime/research/nvfp4_fused_variants/patches/0005-elementwise-support.patch

213 lines
6.9 KiB
Diff
Raw Permalink Normal View History

diff --git a/src/h3_blackwell_runtime/h3_fusion.py b/src/h3_blackwell_runtime/h3_fusion.py
new file mode 100644
index 0000000..0ee54db
--- /dev/null
+++ b/src/h3_blackwell_runtime/h3_fusion.py
@@ -0,0 +1,206 @@
+"""Triton fusion for H3's segmented BF16 modulation and residual gates."""
+
+from __future__ import annotations
+
+import torch
+import triton
+import triton.language as tl
+
+
+_TILE = 1024
+_segment_cache_key = None
+_segment_cache_value = None
+
+
+@triton.jit
+def _round_bf16_fp32(value):
+ """Apply round-to-nearest-even BF16 precision while retaining FP32."""
+ bits = value.to(tl.int32, bitcast=True)
+ bits = bits + 0x7FFF + ((bits >> 16) & 1)
+ return (bits & -65536).to(tl.float32, bitcast=True)
+
+
+@triton.jit
+def _modulate_kernel(
+ x_ptr,
+ scale_ptr,
+ shift_ptr,
+ row_index_ptr,
+ tokens,
+ width,
+ sx_t,
+ sx_d,
+ ss_t,
+ ss_d,
+ sh_t,
+ sh_d,
+ BLOCK: tl.constexpr,
+):
+ row = tl.program_id(0)
+ columns = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK)
+ valid = (row < tokens) & (columns < width)
+ table_row = tl.load(row_index_ptr + row)
+ x = tl.load(x_ptr + row * sx_t + columns * sx_d, mask=valid).to(tl.float32)
+ scale = tl.load(
+ scale_ptr + table_row * ss_t + columns * ss_d, mask=valid,
+ ).to(tl.float32)
+ shift = tl.load(
+ shift_ptr + table_row * sh_t + columns * sh_d, mask=valid,
+ ).to(tl.float32)
+
+ scale = _round_bf16_fp32(scale)
+ shift = _round_bf16_fp32(shift)
+ multiplied = _round_bf16_fp32(x * _round_bf16_fp32(1.0 + scale))
+ result = _round_bf16_fp32(multiplied + shift)
+ tl.store(x_ptr + row * sx_t + columns * sx_d, result.to(tl.bfloat16), mask=valid)
+
+
+@triton.jit
+def _gate_add_kernel(
+ residual_ptr,
+ gate_ptr,
+ update_ptr,
+ row_index_ptr,
+ tokens,
+ width,
+ sr_t,
+ sr_d,
+ sg_t,
+ sg_d,
+ su_t,
+ su_d,
+ BLOCK: tl.constexpr,
+):
+ row = tl.program_id(0)
+ columns = tl.program_id(1) * BLOCK + tl.arange(0, BLOCK)
+ valid = (row < tokens) & (columns < width)
+ table_row = tl.load(row_index_ptr + row)
+ residual = tl.load(
+ residual_ptr + row * sr_t + columns * sr_d, mask=valid,
+ ).to(tl.float32)
+ update = tl.load(
+ update_ptr + row * su_t + columns * su_d, mask=valid,
+ ).to(tl.float32)
+ gate = tl.load(
+ gate_ptr + table_row * sg_t + columns * sg_d, mask=valid,
+ ).to(tl.float32)
+
+ result = residual + update * _round_bf16_fp32(gate)
+ tl.store(
+ residual_ptr + row * sr_t + columns * sr_d,
+ result.to(tl.bfloat16),
+ mask=valid,
+ )
+
+
+def segment_index(
+ tokens: int,
+ segments: list[tuple[int, int, int]],
+ device: torch.device,
+) -> torch.Tensor:
+ """Return one cached row-to-modulation-table lookup for a packed layout."""
+ global _segment_cache_key, _segment_cache_value
+ normalized = tuple((int(start), int(stop), int(row)) for start, stop, row in segments)
+ key = (str(device), int(tokens), normalized)
+ if key == _segment_cache_key:
+ return _segment_cache_value
+ host = torch.empty(tokens, dtype=torch.int32)
+ cursor = 0
+ for start, stop, row in normalized:
+ if start != cursor or stop < start or stop > tokens or row < 0:
+ raise ValueError("segments must be ordered, contiguous, and in range")
+ host[start:stop] = row
+ cursor = stop
+ if cursor != tokens:
+ raise ValueError("segments must cover every packed token")
+ _segment_cache_key = key
+ _segment_cache_value = host.to(device=device)
+ return _segment_cache_value
+
+
+def _validate_activation(tensor: torch.Tensor, name: str) -> None:
+ if tensor.ndim != 2:
+ raise ValueError(f"{name} must have shape [tokens, hidden]")
+ if tensor.device.type != "cuda" or tensor.dtype != torch.bfloat16:
+ raise TypeError(f"{name} must be a CUDA BF16 tensor")
+ if tensor.stride(1) != 1:
+ raise ValueError(f"{name}'s hidden dimension must be contiguous")
+
+
+def _validate_table(table: torch.Tensor, activation: torch.Tensor, name: str) -> None:
+ if table.ndim != 2 or table.shape[1] != activation.shape[1]:
+ raise ValueError(f"{name} must have shape [rows, {activation.shape[1]}]")
+ if table.device != activation.device or table.dtype not in {torch.bfloat16, torch.float32}:
+ raise TypeError(f"{name} must be CUDA BF16/FP32 on the activation device")
+ if table.stride(1) != 1:
+ raise ValueError(f"{name}'s hidden dimension must be contiguous")
+
+
+def fused_modulate_(
+ activation: torch.Tensor,
+ shift: torch.Tensor,
+ scale: torch.Tensor,
+ row_index: torch.Tensor,
+) -> torch.Tensor:
+ """Apply segmented scale/shift in place with eager-equivalent BF16 rounding."""
+ _validate_activation(activation, "activation")
+ _validate_table(shift, activation, "shift")
+ _validate_table(scale, activation, "scale")
+ if row_index.shape != (activation.shape[0],) or row_index.dtype != torch.int32:
+ raise TypeError("row_index must be int32 with one entry per token")
+ if row_index.device != activation.device or not row_index.is_contiguous():
+ raise TypeError("row_index must be contiguous on the activation device")
+ grid = (activation.shape[0], triton.cdiv(activation.shape[1], _TILE))
+ _modulate_kernel[grid](
+ activation,
+ scale,
+ shift,
+ row_index,
+ activation.shape[0],
+ activation.shape[1],
+ activation.stride(0),
+ activation.stride(1),
+ scale.stride(0),
+ scale.stride(1),
+ shift.stride(0),
+ shift.stride(1),
+ BLOCK=_TILE,
+ num_warps=4,
+ )
+ return activation
+
+
+def fused_gate_add_(
+ residual: torch.Tensor,
+ update: torch.Tensor,
+ gate: torch.Tensor,
+ row_index: torch.Tensor,
+) -> torch.Tensor:
+ """Apply the segmented residual gate in place with addcmul-equivalent math."""
+ _validate_activation(residual, "residual")
+ _validate_activation(update, "update")
+ if residual.shape != update.shape or residual.device != update.device:
+ raise ValueError("residual and update must share shape and device")
+ _validate_table(gate, residual, "gate")
+ if row_index.shape != (residual.shape[0],) or row_index.dtype != torch.int32:
+ raise TypeError("row_index must be int32 with one entry per token")
+ if row_index.device != residual.device or not row_index.is_contiguous():
+ raise TypeError("row_index must be contiguous on the residual device")
+ grid = (residual.shape[0], triton.cdiv(residual.shape[1], _TILE))
+ _gate_add_kernel[grid](
+ residual,
+ gate,
+ update,
+ row_index,
+ residual.shape[0],
+ residual.shape[1],
+ residual.stride(0),
+ residual.stride(1),
+ gate.stride(0),
+ gate.stride(1),
+ update.stride(0),
+ update.stride(1),
+ BLOCK=_TILE,
+ num_warps=4,
+ )
+ return residual