core.fusions.fused_mhc_kernels#

Fused kernels for mHC (Manifold-Constrained Hyper-Connections).

Uses Triton and cuda.tile (cuTile) kernels according to an explicit backend policy. With the default auto policy, unavailable accelerated operations fall back to PyTorch reference implementations. Reference implementations live in megatron.core.transformer.hyper_connection and are also used when the use_fused_mhc config flag is False.

Four fused operations:

  • sinkhorn: Sinkhorn-Knopp projection to doubly stochastic matrix

  • h_aggregate: weighted n-stream -> 1-stream aggregation

  • h_post_bda: fused H_res.T @ residual + H_post * (x + bias)

  • proj_rms_compute_h: fused projection + RMS normalization + compute_h

Determinism: The native, Triton, and cuTile backends are currently tagged unknown. Numerical and gradient parity are covered, but bit-exact repeatability has not been certified. The accelerated backends use timing-based autotuning.

Module Contents#

Classes#

FusedHAggregate

H_aggregate dispatched according to the configured backend policy.

FusedHPostBDA

H_post_bda dispatched according to the configured backend policy.

Functions#

is_cutile_available

Return True if cuTile fused kernels are enabled.

_get_tileiras_path

Return the tileiras compiler path if it can be found.

_cutile_supports_current_device

Return whether cuTile can compile for the current CUDA device.

is_triton_available

Return True if Triton is enabled for supported mHC kernels.

_validate_mhc_backend

Validate an explicit mHC backend policy before dispatch.

_backend_uses_triton

_backend_uses_cutile

_select_triton_cutile_native

_mhc_backend_status

Return backend description and whether every backend is native.

_mhc_backend_selection

Return a concise description of the selected mHC fused backends.

log_fused_mhc_backend_once

Log each configured fused mHC backend policy once per process.

fused_add_3

Add three tensors using the native torch.compile-backed implementation.

_get_triton_sinkhorn

_get_triton_h_aggregate_fwd

_get_triton_h_post_bda_fwd

_get_triton_h_post_bda_bwd

_torch_h_aggregate_bwd

_torch_h_post_bda_bwd

_torch_proj_rms_compute_h

fused_sinkhorn

Project logits according to the configured backend policy.

fused_h_aggregate

Aggregate n streams into one according to the configured backend policy.

fused_h_post_bda

Compute H_res.T @ residual + H_post * (x + bias) using the backend policy.

fused_proj_rms_compute_h

Compute projection, RMS norm, and H outputs using the backend policy.

Data#

API#

core.fusions.fused_mhc_kernels.logger#

‘getLogger(…)’

core.fusions.fused_mhc_kernels.LOG2E#

‘log2(…)’

core.fusions.fused_mhc_kernels.MHC_BACKEND_DETERMINISM: dict[str, str]#

None

core.fusions.fused_mhc_kernels.MHCBackend#

None

core.fusions.fused_mhc_kernels._VALID_MHC_BACKENDS#

(‘auto’, ‘native’, ‘triton’, ‘cutile’)

core.fusions.fused_mhc_kernels._CUTILE_AVAILABLE#

False

core.fusions.fused_mhc_kernels._CUTILE_EXPERIMENTAL_AVAILABLE#

False

core.fusions.fused_mhc_kernels._CUTILE_DEVICE_SUPPORT_CACHE: Optional[bool]#

None

core.fusions.fused_mhc_kernels._CUTILE_DEVICE_SUPPORT_ERROR: Optional[str]#

None

core.fusions.fused_mhc_kernels._TRITON_AVAILABLE#

False

core.fusions.fused_mhc_kernels.is_cutile_available() bool#

Return True if cuTile fused kernels are enabled.

core.fusions.fused_mhc_kernels._get_tileiras_path() Optional[str]#

Return the tileiras compiler path if it can be found.

core.fusions.fused_mhc_kernels._cutile_supports_current_device() bool#

Return whether cuTile can compile for the current CUDA device.

core.fusions.fused_mhc_kernels.is_triton_available() bool#

Return True if Triton is enabled for supported mHC kernels.

core.fusions.fused_mhc_kernels._validate_mhc_backend(
backend: core.fusions.fused_mhc_kernels.MHCBackend,
) None#

Validate an explicit mHC backend policy before dispatch.

core.fusions.fused_mhc_kernels._backend_uses_triton(
backend: core.fusions.fused_mhc_kernels.MHCBackend,
) bool#
core.fusions.fused_mhc_kernels._backend_uses_cutile(
backend: core.fusions.fused_mhc_kernels.MHCBackend,
) bool#
core.fusions.fused_mhc_kernels._TRITON_IMPLS#

None

core.fusions.fused_mhc_kernels._BACKEND_INFO_LOGGED: set[str]#

‘set(…)’

core.fusions.fused_mhc_kernels._select_triton_cutile_native(
triton_impl,
backend: core.fusions.fused_mhc_kernels.MHCBackend,
) str#
core.fusions.fused_mhc_kernels._mhc_backend_status(
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
) Tuple[str, bool]#

Return backend description and whether every backend is native.

core.fusions.fused_mhc_kernels._mhc_backend_selection(
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
) str#

Return a concise description of the selected mHC fused backends.

core.fusions.fused_mhc_kernels.log_fused_mhc_backend_once(
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
) None#

Log each configured fused mHC backend policy once per process.

core.fusions.fused_mhc_kernels.fused_add_3(
a: torch.Tensor,
b: torch.Tensor,
c: torch.Tensor,
) torch.Tensor#

Add three tensors using the native torch.compile-backed implementation.

core.fusions.fused_mhc_kernels._get_triton_sinkhorn(
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
)#
core.fusions.fused_mhc_kernels._get_triton_h_aggregate_fwd(
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
)#
core.fusions.fused_mhc_kernels._get_triton_h_post_bda_fwd(
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
)#
core.fusions.fused_mhc_kernels._get_triton_h_post_bda_bwd(
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
)#
core.fusions.fused_mhc_kernels._torch_h_aggregate_bwd(
grad_output: torch.Tensor,
x: torch.Tensor,
h_pre: torch.Tensor,
) Tuple[torch.Tensor, torch.Tensor]#
core.fusions.fused_mhc_kernels._torch_h_post_bda_bwd(
grad_output: torch.Tensor,
h_res: torch.Tensor,
original_residual: torch.Tensor,
h_post: torch.Tensor,
x: torch.Tensor,
bias: Optional[torch.Tensor],
) Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor, Optional[torch.Tensor]]#
core.fusions.fused_mhc_kernels._torch_proj_rms_compute_h(
x: torch.Tensor,
weight: torch.Tensor,
alpha_pre: torch.Tensor,
alpha_post: torch.Tensor,
alpha_res: torch.Tensor,
bias: torch.Tensor,
n: int,
eps: float,
compute_h_eps: float = 1e-06,
) Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]#
class core.fusions.fused_mhc_kernels.FusedHAggregate#

Bases: torch.autograd.Function

H_aggregate dispatched according to the configured backend policy.

static forward(
ctx,
x: torch.Tensor,
h_pre: torch.Tensor,
backend: core.fusions.fused_mhc_kernels.MHCBackend,
)#

Run h_aggregate forward using the best available backend.

static backward(ctx, grad_output)#

Run h_aggregate backward using the best available backend.

class core.fusions.fused_mhc_kernels.FusedHPostBDA#

Bases: torch.autograd.Function

H_post_bda dispatched according to the configured backend policy.

static forward(
ctx,
h_res: torch.Tensor,
original_residual: torch.Tensor,
h_post: torch.Tensor,
x: torch.Tensor,
bias: Optional[torch.Tensor],
backend: core.fusions.fused_mhc_kernels.MHCBackend,
)#

Run h_post_bda forward using the best available backend.

static backward(ctx, grad_output)#

Run h_post_bda backward using the best available backend.

core.fusions.fused_mhc_kernels.fused_sinkhorn(
input_logits: torch.Tensor,
num_iterations: int,
eps: float = 1e-06,
*,
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
) torch.Tensor#

Project logits according to the configured backend policy.

core.fusions.fused_mhc_kernels.fused_h_aggregate(
x: torch.Tensor,
h_pre: torch.Tensor,
*,
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
) torch.Tensor#

Aggregate n streams into one according to the configured backend policy.

core.fusions.fused_mhc_kernels.fused_h_post_bda(
h_res: torch.Tensor,
original_residual: torch.Tensor,
h_post: torch.Tensor,
x: torch.Tensor,
bias: Optional[torch.Tensor],
*,
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
) torch.Tensor#

Compute H_res.T @ residual + H_post * (x + bias) using the backend policy.

core.fusions.fused_mhc_kernels.fused_proj_rms_compute_h(
x: torch.Tensor,
weight: torch.Tensor,
alpha_pre: torch.Tensor,
alpha_post: torch.Tensor,
alpha_res: torch.Tensor,
bias: torch.Tensor,
n: int,
eps: float = 1e-06,
compute_h_eps: float = 1e-06,
*,
backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
) Tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]#

Compute projection, RMS norm, and H outputs using the backend policy.