diff --git a/src/h3_blackwell_runtime/sage2_entry.py b/src/h3_blackwell_runtime/sage2_entry.py new file mode 100644 index 0000000..215d4d1 --- /dev/null +++ b/src/h3_blackwell_runtime/sage2_entry.py @@ -0,0 +1,81 @@ +"""Exact fused entry preparation for the pinned SageAttention 2.2.0 path.""" + +from __future__ import annotations + +import math + +import torch + + +def prepare_qk( + q: torch.Tensor, + k: torch.Tensor, + rotation: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + epsilon: float, + *, + materialize_q: bool = False, +) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Prepare Q/K in place and emit Sage2's exact per-warp Q representation.""" + from .nvfp4_quant import _vortex_scale_extension + + return tuple( + _vortex_scale_extension().sage2_prepare_qk( + q, k, rotation, q_weight, k_weight, epsilon, materialize_q, + ) + ) + + +def prepare_v(v: torch.Tensor, *, scale_max: float = 2.25) -> tuple[torch.Tensor, torch.Tensor]: + """Emit Sage2's FP8 V tensor and per-channel scale without a BF16 transpose tensor.""" + from .nvfp4_quant import _vortex_scale_extension + + return tuple(_vortex_scale_extension().sage2_prepare_v(v, scale_max)) + + +def attention_nhd( + q: torch.Tensor, + k: torch.Tensor, + v: torch.Tensor, + rotation: torch.Tensor, + q_weight: torch.Tensor, + k_weight: torch.Tensor, + epsilon: float, +) -> torch.Tensor: + """Run the unchanged Sage2 mainloop after fused strided-NHD entry preparation.""" + import sageattention.core as sage_core + import sageattention.quant as sage_quant + + q_int8, q_scale, _ = prepare_qk( + q, k, rotation, q_weight, k_weight, epsilon, materialize_q=False, + ) + k_mean = k.mean(dim=1, keepdim=True) + k_int8 = torch.empty(k.shape, dtype=torch.int8, device=k.device) + k_scale = torch.empty( + (k.shape[0], k.shape[2], math.ceil(k.shape[1] / 64)), + dtype=torch.float32, + device=k.device, + ) + sage_quant._fused.quant_per_block_int8_fuse_sub_mean_cuda( + k, k_mean.squeeze(1), k_int8, k_scale, 64, 0, + ) + v_fp8, v_scale, _ = sage_core.per_channel_fp8( + v, tensor_layout="NHD", scale_max=2.25, smooth_v=False, + ) + output = torch.empty(q.shape, dtype=q.dtype, device=q.device) + sage_core.sm89_compile.qk_int8_sv_f8_accum_f16_fuse_v_scale_attn_inst_buf( + q_int8, + k_int8, + v_fp8, + output, + q_scale, + k_scale, + v_scale, + 0, + 0, + 2, + output.shape[-1] ** -0.5, + 0, + ) + return output