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#
Non-tensor fields needed to reconstruct the CuTeDSL backend’s saved CP context. |
|
Adapt |
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,
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,
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,
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,
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,
- core.ssm.context_parallel.gdp_cutedsl._backward_packed_tensor(
- summary: megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary,
- core.ssm.context_parallel.gdp_cutedsl._restore_saved_context(
- saved_context: megatron.core.ssm.context_parallel.chunkwise.CPSavedContext,