h3-blackwell-runtime/src/h3_blackwell_runtime/block.py
2026-08-12 14:12:42 +07:00

99 lines
3.5 KiB
Python

"""Direct MiniMax H3 DiT block over the standalone Sage3 attention unit."""
import torch
from torch import nn
from .attention import H3SageAttention, rms_norm
from .checkpoint import H3Checkpoint
from .nvfp4 import Nvfp4Linear
def modulate_segments(
x: torch.Tensor,
shift: torch.Tensor,
scale: torch.Tensor,
segments: list[tuple[int, int, int]],
) -> torch.Tensor:
"""Apply H3's per-modality/per-timestep AdaLN parameters to contiguous rows."""
output = torch.empty_like(x)
for start, stop, row in segments:
output[start:stop] = x[start:stop] * (1 + scale[row].to(x.dtype)) + shift[row].to(x.dtype)
return output
def gate_segments(
residual: torch.Tensor,
update: torch.Tensor,
gate: torch.Tensor,
segments: list[tuple[int, int, int]],
) -> torch.Tensor:
"""Add a gated fresh sublayer output without altering the source residual."""
output = residual.clone()
for start, stop, row in segments:
output[start:stop].addcmul_(update[start:stop], gate[row].to(update.dtype))
return output
class H3SwiGLU(nn.Module):
def __init__(self, fc1: Nvfp4Linear, fc2: Nvfp4Linear):
super().__init__()
self.fc1 = fc1
self.fc2 = fc2
@classmethod
def from_checkpoint(cls, checkpoint: H3Checkpoint, prefix: str, *, output_dtype=torch.bfloat16):
return cls(
checkpoint.nvfp4_linear(f"{prefix}.fc1", output_dtype=output_dtype),
checkpoint.nvfp4_linear(f"{prefix}.fc2", output_dtype=output_dtype),
)
def forward(self, x: torch.Tensor) -> torch.Tensor:
gate, up = self.fc1(x).chunk(2, dim=-1)
return self.fc2(torch.nn.functional.silu(gate).mul_(up))
class H3DiTBlock(nn.Module):
"""One H3 transformer block with externally supplied AdaLN tensors."""
def __init__(
self,
norm1_weight: torch.Tensor,
norm2_weight: torch.Tensor,
attention: H3SageAttention,
mlp: H3SwiGLU,
*,
norm_eps: float = 1e-5,
):
super().__init__()
self.attention = attention
self.mlp = mlp
self.norm_eps = norm_eps
self.register_buffer("norm1_weight", norm1_weight, persistent=False)
self.register_buffer("norm2_weight", norm2_weight, persistent=False)
@classmethod
def from_checkpoint(cls, checkpoint: H3Checkpoint, index: int, *, output_dtype=torch.bfloat16, attention_backend: str = "sage2"):
prefix = f"blocks.{index}"
return cls(
checkpoint.tensor(f"{prefix}.norm1.weight", dtype=output_dtype),
checkpoint.tensor(f"{prefix}.norm2.weight", dtype=output_dtype),
H3SageAttention.from_checkpoint(checkpoint, f"{prefix}.attn", output_dtype=output_dtype, backend=attention_backend),
H3SwiGLU.from_checkpoint(checkpoint, f"{prefix}.mlp", output_dtype=output_dtype),
)
def forward(
self,
x: torch.Tensor,
rope_rotation: torch.Tensor,
shift_msa: torch.Tensor,
scale_msa: torch.Tensor,
gate_msa: torch.Tensor,
shift_mlp: torch.Tensor,
scale_mlp: torch.Tensor,
gate_mlp: torch.Tensor,
segments: list[tuple[int, int, int]],
) -> torch.Tensor:
h = modulate_segments(rms_norm(x, self.norm1_weight, self.norm_eps), shift_msa, scale_msa, segments)
x = gate_segments(x, self.attention(h, rope_rotation), gate_msa, segments)
h = modulate_segments(rms_norm(x, self.norm2_weight, self.norm_eps), shift_mlp, scale_mlp, segments)
return gate_segments(x, self.mlp(h), gate_mlp, segments)