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#
H_aggregate dispatched according to the configured backend policy. |
|
H_post_bda dispatched according to the configured backend policy. |
Functions#
Return True if cuTile fused kernels are enabled. |
|
Return the tileiras compiler path if it can be found. |
|
Return whether cuTile can compile for the current CUDA device. |
|
Return True if Triton is enabled for supported mHC kernels. |
|
Validate an explicit mHC backend policy before dispatch. |
|
Return backend description and whether every backend is native. |
|
Return a concise description of the selected mHC fused backends. |
|
Log each configured fused mHC backend policy once per process. |
|
Add three tensors using the native torch.compile-backed implementation. |
|
Project logits according to the configured backend policy. |
|
Aggregate n streams into one according to the configured backend policy. |
|
Compute H_res.T @ residual + H_post * (x + bias) using the backend policy. |
|
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( ) None#
Validate an explicit mHC backend policy before dispatch.
- core.fusions.fused_mhc_kernels._backend_uses_triton( ) bool#
- core.fusions.fused_mhc_kernels._backend_uses_cutile( ) 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,
- core.fusions.fused_mhc_kernels._mhc_backend_status(
- backend: core.fusions.fused_mhc_kernels.MHCBackend = 'auto',
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',
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',
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,
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,
- 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],
- 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,
- class core.fusions.fused_mhc_kernels.FusedHAggregate#
Bases:
torch.autograd.FunctionH_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.FunctionH_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',
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',
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',
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',
Compute projection, RMS norm, and H outputs using the backend policy.