h3-blackwell-runtime/research/qkv_layout_variants/patches/0001-strided-nhd-dispatch.patch
2026-08-25 20:30:22 +07:00

61 lines
3.4 KiB
Diff

diff --git a/src/h3_blackwell_runtime/attention.py b/src/h3_blackwell_runtime/attention.py
index 7b7560d..a5372f8 100644
--- a/src/h3_blackwell_runtime/attention.py
+++ b/src/h3_blackwell_runtime/attention.py
@@ -169,6 +169,13 @@ def run_attention(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor, *, backend:
raise ValueError(f"Unsupported H3 attention backend: {backend}")
+def run_sage_attention_nhd(q: torch.Tensor, k: torch.Tensor, v: torch.Tensor) -> torch.Tensor:
+ """Run Sage2 directly on projection-strided NHD Q/K/V views."""
+ from sageattention import sageattn
+
+ return sageattn(q, k, v, is_causal=False, tensor_layout="NHD", smooth_k=False)
+
+
def apply_split_half_rope(x: torch.Tensor, rotation: torch.Tensor) -> torch.Tensor:
"""Apply H3's split-half rotary table to `[batch, sequence, heads, dim]`."""
rotated_width = rotation.shape[-3] * 2
@@ -231,8 +238,8 @@ class H3SageAttention(nn.Module):
@classmethod
def from_checkpoint(cls, checkpoint: H3Checkpoint, prefix: str, *, output_dtype=torch.bfloat16, backend: str = DEFAULT_ATTENTION_BACKEND):
return cls(
- checkpoint.nvfp4_linear(f"{prefix}.qkv_proj", output_dtype=output_dtype),
- checkpoint.nvfp4_linear(f"{prefix}.out_proj", output_dtype=output_dtype),
+ checkpoint.nvfp4_linear(f"{prefix}.qkv_proj", output_dtype=output_dtype, role="h3_attn_qkv"),
+ checkpoint.nvfp4_linear(f"{prefix}.out_proj", output_dtype=output_dtype, role="h3_attn_out"),
checkpoint.tensor(f"{prefix}.q_norm.weight", dtype=output_dtype),
checkpoint.tensor(f"{prefix}.k_norm.weight", dtype=output_dtype),
backend=backend,
@@ -244,6 +251,7 @@ class H3SageAttention(nn.Module):
rope_rotation: torch.Tensor,
sequence_parallel: "SequenceParallelContext | None" = None,
tensor_parallel: "SequenceParallelContext | None" = None,
+ modulation: tuple[torch.Tensor, torch.Tensor, torch.Tensor] | None = None,
) -> torch.Tensor:
if x.ndim != 2:
raise ValueError("H3 attention expects `[sequence, hidden]` input.")
@@ -253,13 +261,22 @@ class H3SageAttention(nn.Module):
return self._forward_tensor_parallel(x, rope_rotation, tensor_parallel)
sequence = x.shape[0]
inner = self.heads * self.head_dim
- qkv = self.qkv_proj(x)
+ qkv = self.qkv_proj.forward_modulated(x, *modulation) if modulation is not None else self.qkv_proj(x)
q, k, v = qkv.split(inner, dim=-1)
q = q.view(1, sequence, self.heads, self.head_dim)
k = k.view(1, sequence, self.heads, self.head_dim)
v = v.view(1, sequence, self.heads, self.head_dim)
q, k = rms_rope_split_half_(q, k, rope_rotation, self.q_norm_weight, self.k_norm_weight, self.eps)
+ if (
+ sequence_parallel is None
+ and self.backend == "sage2"
+ and os.getenv("H3_SAGE_QKV_LAYOUT", "hnd").lower() == "strided_nhd"
+ ):
+ out = run_sage_attention_nhd(q, k, v)
+ if not out.is_contiguous():
+ raise RuntimeError("Sage2 NHD output must be contiguous for zero-copy output projection")
+ return self.out_proj(out.reshape(sequence, inner))
if sequence_parallel is not None:
q, k, v = sequence_parallel.seq_to_heads(q, k, v)
if self.backend == "sol_attn":