core.ssm.context_parallel.gdp_cutedsl#

CuTeDSL backend for Gated Delta Product chunkwise context parallelism.

This backend is deterministic, but not batch invariant.

Module Contents#

Classes#

CuTeDSLGDPSavedMetadata

Non-tensor fields needed to reconstruct the CuTeDSL backend’s saved CP context.

CuTeDSLGatedDeltaProductCPBackend

Adapt gdp_attn’s four-function CP protocol to Megatron collectives.

Functions#

API#

class core.ssm.context_parallel.gdp_cutedsl.CuTeDSLGDPSavedMetadata#

Non-tensor fields needed to reconstruct the CuTeDSL backend’s saved CP context.

scale: float#

None

num_householder: int#

None

use_qk_l2norm_in_kernel: bool#

None

recompute_chunk_num: int#

None

has_initial_states: bool#

None

class core.ssm.context_parallel.gdp_cutedsl.CuTeDSLGatedDeltaProductCPBackend(recompute_chunk_num: int = 0)#

Adapt gdp_attn’s four-function CP protocol to Megatron collectives.

Initialization

autograd_function#

None

cp_forward_prepare(
inputs: megatron.core.ssm.context_parallel.gdp_common.GDPInputs,
) tuple[megatron.core.ssm.context_parallel.chunkwise.CPForwardSummary, gdp_attn.chunk_gated_delta_product.GdpCpForwardLocalContext]#

Compute the trailing-boundary [E | M] summary for this rank.

cp_forward_apply(
local_context: gdp_attn.chunk_gated_delta_product.GdpCpForwardLocalContext,
preceding_summaries: megatron.core.ssm.context_parallel.chunkwise.CPForwardSummary,
) tuple[torch.Tensor, megatron.core.ssm.context_parallel.chunkwise.CPSavedContext]#

Apply the causal prefix and retain the CuTeDSL backward checkpoint state.

cp_backward_prepare(
output_grad: torch.Tensor,
saved_context: megatron.core.ssm.context_parallel.chunkwise.CPSavedContext,
) tuple[megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary, gdp_attn.chunk_gated_delta_product.GdpCpBackwardContext]#

Compute the leading-boundary [gamma | M^T] summary for this rank.

cp_backward_apply(
backward_context: gdp_attn.chunk_gated_delta_product.GdpCpBackwardContext,
following_summaries: megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary,
) megatron.core.ssm.context_parallel.gdp_common.GDPInputGradients#

Apply the reverse-causal suffix and compute all local input gradients.

core.ssm.context_parallel.gdp_cutedsl._forward_packed_tensor(
summary: megatron.core.ssm.context_parallel.chunkwise.CPForwardSummary,
) torch.Tensor#
core.ssm.context_parallel.gdp_cutedsl._backward_packed_tensor(
summary: megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary,
) torch.Tensor#
core.ssm.context_parallel.gdp_cutedsl._restore_saved_context(
saved_context: megatron.core.ssm.context_parallel.chunkwise.CPSavedContext,
) gdp_attn.chunk_gated_delta_product.GdpCpSavedContext#