core.fusions.fused_gated_norm#

GDN output RMSNorm/SiLU fusion preserving intermediate activation rounding.

Module Contents#

Classes#

_FusedGatedNorm

First-order autograd for fused output gating.

Functions#

_row_offset

Address logical [batch, sequence, head] rows without copying tensor views.

gated_norm_fwd

Apply output RMSNorm and SiLU gating.

gated_norm_bwd

Differentiate output RMSNorm and SiLU gating.

forward_impl

Normalize and gate [batch, sequence, head, dimension] tensor views.

backward_impl

Compute activation, gate, and RMSNorm-weight gradients.

fused_gated_norm

Run fused RMSNorm/SiLU gating with first-order autograd support.

validate_gated_norm

Reject unsupported GDN output fusion configurations and layouts.

API#

core.fusions.fused_gated_norm._row_offset(
row,
SEQUENCE: triton.language.constexpr,
HEADS: triton.language.constexpr,
STRIDE_B: triton.language.constexpr,
STRIDE_S: triton.language.constexpr,
STRIDE_H: triton.language.constexpr,
FLAT_TOKENS: triton.language.constexpr,
)#

Address logical [batch, sequence, head] rows without copying tensor views.

core.fusions.fused_gated_norm.gated_norm_fwd(
x,
gate,
weight,
out,
rstd,
ROWS: triton.language.constexpr,
HEADS: triton.language.constexpr,
D: triton.language.constexpr,
SEQUENCE: triton.language.constexpr,
X_STRIDES: triton.language.constexpr,
GATE_STRIDES: triton.language.constexpr,
X_FLAT_TOKENS: triton.language.constexpr,
GATE_FLAT_TOKENS: triton.language.constexpr,
EPS: triton.language.constexpr,
ZERO_CENTERED: triton.language.constexpr,
BT: triton.language.constexpr,
)#

Apply output RMSNorm and SiLU gating.

core.fusions.fused_gated_norm.gated_norm_bwd(
x,
gate,
weight,
rstd,
dy,
dx,
dgate,
dw_partial,
ROWS: triton.language.constexpr,
HEADS: triton.language.constexpr,
D: triton.language.constexpr,
SEQUENCE: triton.language.constexpr,
X_STRIDES: triton.language.constexpr,
GATE_STRIDES: triton.language.constexpr,
X_FLAT_TOKENS: triton.language.constexpr,
GATE_FLAT_TOKENS: triton.language.constexpr,
ZERO_CENTERED: triton.language.constexpr,
BT: triton.language.constexpr,
)#

Differentiate output RMSNorm and SiLU gating.

core.fusions.fused_gated_norm.forward_impl(
x: torch.Tensor,
gate: torch.Tensor,
weight: torch.Tensor,
eps: float,
zero_centered: bool = False,
bt: int = 4,
warps: int = 2,
) → tuple[torch.Tensor, torch.Tensor]#

Normalize and gate [batch, sequence, head, dimension] tensor views.

core.fusions.fused_gated_norm.backward_impl(
x: torch.Tensor,
gate: torch.Tensor,
weight: torch.Tensor,
rstd: torch.Tensor,
dy: torch.Tensor,
zero_centered: bool = False,
bt: int = 16,
warps: int = 2,
) → tuple[torch.Tensor, torch.Tensor, torch.Tensor]#

Compute activation, gate, and RMSNorm-weight gradients.

class core.fusions.fused_gated_norm._FusedGatedNorm#

Bases: torch.autograd.Function

First-order autograd for fused output gating.

static forward(ctx, x, gate, weight, eps, zero_centered)#

Save activations and inverse norms for backward.

static backward(ctx, dy)#

Return activation, gate, and norm-weight gradients.

core.fusions.fused_gated_norm.fused_gated_norm(
x: torch.Tensor,
gate: torch.Tensor,
weight: torch.Tensor,
eps: float,
zero_centered: bool = False,
) → torch.Tensor#

Run fused RMSNorm/SiLU gating with first-order autograd support.

core.fusions.fused_gated_norm.validate_gated_norm(
module: torch.nn.Module,
x: torch.Tensor,
gate: torch.Tensor,
) → None#

Reject unsupported GDN output fusion configurations and layouts.

Called on every fused module forward, including output-norm recomputation. These checks inspect host-side metadata only and do not synchronize CUDA.