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