61 lines
3.4 KiB
Diff
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":
|