213 lines
6.9 KiB
Diff
213 lines
6.9 KiB
Diff
|
|
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
|