h3-blackwell-runtime/research/sage2_entry_fusion/patches/0001-sage2-entry-wrapper.patch
2026-08-25 20:30:22 +07:00

87 lines
2.6 KiB
Diff

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