87 lines
2.6 KiB
Diff
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
|