core.ssm.context_parallel.gdp#

FLA backend for Gated Delta Product chunkwise context parallelism.

Module Contents#

Classes#

GDPSavedInputs

GDP operands retained for the FLA backward kernels.

GDPLocalContext

FLA GDP workspace passed from forward prepare to forward apply.

GDPSavedMetadata

Non-tensor GDP state saved for backward.

GDPBackwardContext

GDP state passed from backward prepare to backward apply.

FLAGDPChunkwiseContextParallel

FLA-decorated form of the shared GDP autograd adapter.

FLAGatedDeltaProductCPBackend

FLA implementation of the linear-attention chunkwise-CP interface.

Functions#

_expand_cu_seqlens

_prepare_fla_forward

Build the GDP WY representation used by both summary and output kernels.

_compute_fragment_summary

Compute packed [delta_S | M] for the rank-local trailing boundary fragment.

_compute_backward_fragment_summary

Compute packed [gamma | M^T] for the rank-local leading boundary fragment.

_merge_affine_summaries

Compose an already-selected summary slice from zero into output.

_merge_forward_summaries

_merge_backward_summaries

_state_shape

_build_initial_state

_build_final_state_grad

_restore_saved_context

_interleave_last_update

_apply_fla_backward

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,
) torch.Tensor#
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.GDPChunkwiseContextParallel

FLA-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,
) tuple[megatron.core.ssm.context_parallel.chunkwise.CPForwardSummary, core.ssm.context_parallel.gdp.GDPLocalContext]#

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,
) tuple[torch.Tensor, megatron.core.ssm.context_parallel.chunkwise.CPSavedContext]#

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,
) tuple[megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary, core.ssm.context_parallel.gdp.GDPBackwardContext]#

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,
) megatron.core.ssm.context_parallel.gdp_common.GDPInputGradients#

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,
) core.ssm.context_parallel.gdp.GDPLocalContext#

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,
) torch.Tensor#

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,
) torch.Tensor#

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,
) torch.Tensor#

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,
) torch.Tensor#
core.ssm.context_parallel.gdp._merge_backward_summaries(
summaries: megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary,
output_grad: torch.Tensor,
) torch.Tensor#
core.ssm.context_parallel.gdp._state_shape(
q: torch.Tensor,
v: torch.Tensor,
cu_seqlens: torch.Tensor | None,
) tuple[int, int, int, int]#
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,
) torch.Tensor#
core.ssm.context_parallel.gdp._build_final_state_grad(
inputs: core.ssm.context_parallel.gdp.GDPSavedInputs,
summaries: megatron.core.ssm.context_parallel.chunkwise.CPBackwardSummary,
) torch.Tensor#
core.ssm.context_parallel.gdp._restore_saved_context(
saved_context: megatron.core.ssm.context_parallel.chunkwise.CPSavedContext,
) tuple[core.ssm.context_parallel.gdp.GDPSavedInputs, torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor | None]#
core.ssm.context_parallel.gdp._interleave_last_update(
tensor: torch.Tensor,
num_householder: int,
) torch.Tensor#
core.ssm.context_parallel.gdp._apply_fla_backward(
context: core.ssm.context_parallel.gdp.GDPBackwardContext,
final_state_grad: torch.Tensor,
) megatron.core.ssm.context_parallel.gdp_common.GDPInputGradients#

Apply FLA’s GDP backward kernels using Megatron-composed boundary state gradients.