core.ssm.context_parallel.gdp#
FLA backend for Gated Delta Product chunkwise context parallelism.
Module Contents#
Classes#
GDP operands retained for the FLA backward kernels. |
|
FLA GDP workspace passed from forward prepare to forward apply. |
|
Non-tensor GDP state saved for backward. |
|
GDP state passed from backward prepare to backward apply. |
|
FLA-decorated form of the shared GDP autograd adapter. |
|
FLA implementation of the linear-attention chunkwise-CP interface. |
Functions#
Build the GDP WY representation used by both summary and output kernels. |
|
Compute packed |
|
Compute packed |
|
Compose an already-selected summary slice from zero into |
|
Apply FLA’s GDP backward kernels using Megatron-composed boundary state gradients. |
API#
- core.ssm.context_parallel.gdp._expand_cu_seqlens(
- cu_seqlens: torch.Tensor,
- num_householder: int,
- class core.ssm.context_parallel.gdp.GDPSavedInputs#
GDP operands retained for the FLA backward kernels.
- q: torch.Tensor#
None
- k: torch.Tensor#
None
- v: torch.Tensor#
None
- beta: torch.Tensor#
None
- cu_seqlens_dp: torch.Tensor | None#
None
- num_householder: int#
None
- scale: float#
None
- g_dtype: torch.dtype#
None
- class core.ssm.context_parallel.gdp.GDPLocalContext#
FLA GDP workspace passed from forward prepare to forward apply.
- inputs: megatron.core.ssm.context_parallel.gdp_common.GDPInputs#
None
- g_cumsum: torch.Tensor#
None
- g_interleaved: torch.Tensor#
None
- A: torch.Tensor#
None
- w: torch.Tensor#
None
- u: torch.Tensor#
None
- cu_seqlens_dp: torch.Tensor | None#
None
- chunk_indices: torch.Tensor | None#
None
- chunk_indices_dp: torch.Tensor | None#
None
- class core.ssm.context_parallel.gdp.GDPSavedMetadata#
Non-tensor GDP state saved for backward.
- num_householder: int#
None
- scale: float#
None
- has_cu_seqlens: bool#
None
- g_dtype: torch.dtype#
None
- class core.ssm.context_parallel.gdp.GDPBackwardContext#
GDP state passed from backward prepare to backward apply.
- inputs: core.ssm.context_parallel.gdp.GDPSavedInputs#
None
- g_interleaved: torch.Tensor#
None
- A: torch.Tensor#
None
- initial_state: torch.Tensor#
None
- chunk_indices_dp: torch.Tensor | None#
None
- q_interleaved: torch.Tensor#
None
- output_grad_interleaved: torch.Tensor#
None
- w: torch.Tensor#
None
- h: torch.Tensor#
None
- v_new: torch.Tensor#
None
- dv: torch.Tensor#
None
- class core.ssm.context_parallel.gdp.FLAGDPChunkwiseContextParallel#
Bases:
megatron.core.ssm.context_parallel.gdp_common.GDPChunkwiseContextParallelFLA-decorated form of the shared GDP autograd adapter.
- forward#
‘staticmethod(…)’
- backward#
‘staticmethod(…)’
- class core.ssm.context_parallel.gdp.FLAGatedDeltaProductCPBackend#
FLA implementation of the linear-attention chunkwise-CP interface.
- autograd_function#
None
- cp_forward_prepare(
- inputs: megatron.core.ssm.context_parallel.gdp_common.GDPInputs,
Compute the local state summary and reusable FLA forward workspace.
- cp_forward_apply(
- local_context: core.ssm.context_parallel.gdp.GDPLocalContext,
- preceding_summaries: megatron.core.ssm.context_parallel.chunkwise.CPForwardSummary,
Compose the incoming state and compute the local GDP output.
- cp_backward_prepare(
- output_grad: torch.Tensor,
- saved_context: megatron.core.ssm.context_parallel.chunkwise.CPSavedContext,
Compute the local output contribution to the incoming-state gradient.
- cp_backward_apply(
- backward_context: core.ssm.context_parallel.gdp.GDPBackwardContext,
- following_summaries: megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary,
Compose the outgoing-state gradient and compute all local GDP gradients.
- core.ssm.context_parallel.gdp._prepare_fla_forward(
- inputs: megatron.core.ssm.context_parallel.gdp_common.GDPInputs,
Build the GDP WY representation used by both summary and output kernels.
- core.ssm.context_parallel.gdp._compute_fragment_summary(
- k: torch.Tensor,
- w: torch.Tensor,
- u: torch.Tensor,
- g: torch.Tensor,
- cu_seqlens: torch.Tensor | None,
Compute packed
[delta_S | M]for the rank-local trailing boundary fragment.
- core.ssm.context_parallel.gdp._compute_backward_fragment_summary(
- q: torch.Tensor,
- k: torch.Tensor,
- w: torch.Tensor,
- g: torch.Tensor,
- do: torch.Tensor,
- dv: torch.Tensor,
- scale: float,
- cu_seqlens: torch.Tensor | None,
Compute packed
[gamma | M^T]for the rank-local leading boundary fragment.
- core.ssm.context_parallel.gdp._merge_affine_summaries(
- packed_summaries: torch.Tensor,
- output: torch.Tensor,
- forward: bool,
Compose an already-selected summary slice from zero into
output.
- core.ssm.context_parallel.gdp._merge_forward_summaries(
- summaries: megatron.core.ssm.context_parallel.chunkwise.CPForwardSummary,
- output: torch.Tensor,
- core.ssm.context_parallel.gdp._merge_backward_summaries(
- summaries: megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary,
- output_grad: torch.Tensor,
- core.ssm.context_parallel.gdp._state_shape(
- q: torch.Tensor,
- v: torch.Tensor,
- cu_seqlens: torch.Tensor | None,
- core.ssm.context_parallel.gdp._build_initial_state(
- inputs: megatron.core.ssm.context_parallel.gdp_common.GDPInputs,
- summaries: megatron.core.ssm.context_parallel.chunkwise.CPForwardSummary,
- core.ssm.context_parallel.gdp._build_final_state_grad(
- inputs: core.ssm.context_parallel.gdp.GDPSavedInputs,
- summaries: megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary,
- core.ssm.context_parallel.gdp._restore_saved_context(
- saved_context: megatron.core.ssm.context_parallel.chunkwise.CPSavedContext,
- core.ssm.context_parallel.gdp._interleave_last_update(
- tensor: torch.Tensor,
- num_householder: int,
- core.ssm.context_parallel.gdp._apply_fla_backward(
- context: core.ssm.context_parallel.gdp.GDPBackwardContext,
- final_state_grad: torch.Tensor,
Apply FLA’s GDP backward kernels using Megatron-composed boundary state gradients.