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":