core.fusions.fused_gated_norm#
GDN output RMSNorm/SiLU fusion preserving intermediate activation rounding.
Module Contents#
Classes#
First-order autograd for fused output gating. |
Functions#
Address logical [batch, sequence, head] rows without copying tensor views. |
|
Apply output RMSNorm and SiLU gating. |
|
Differentiate output RMSNorm and SiLU gating. |
|
Normalize and gate [batch, sequence, head, dimension] tensor views. |
|
Compute activation, gate, and RMSNorm-weight gradients. |
|
Run fused RMSNorm/SiLU gating with first-order autograd support. |
|
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,
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,
Compute activation, gate, and RMSNorm-weight gradients.
- class core.fusions.fused_gated_norm._FusedGatedNorm#
Bases:
torch.autograd.FunctionFirst-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,
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,
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.