core.ssm.gated_delta_net.gdn#

Module Contents#

Classes#

API#

class core.ssm.gated_delta_net.gdn.GatedDeltaNet(
config: megatron.core.transformer.TransformerConfig,
submodules: megatron.core.ssm.gated_delta_net.common.GatedDeltaNetSubmodules,
layer_number: int = None,
bias: bool = False,
conv_bias: bool = False,
conv_init: float | None = None,
use_qk_l2norm: bool = True,
A_init_range: tuple[float, float] = (1, 16),
pg_collection: megatron.core.process_groups_config.ProcessGroupCollection = None,
*,
name: str | None = None,
cp_comm_type: str | None = None,
)#

Bases: megatron.core.ssm.gated_delta_net.common._GDNBase

_setup_variant_attrs()#

Set the GDN in_proj sizing, split tables, gate parameter dims, and kernel.

forward(
hidden_states: torch.Tensor,
attention_mask: torch.Tensor,
inference_context: Optional[megatron.core.inference.contexts.BaseInferenceContext] = None,
packed_seq_params: Optional[megatron.core.packed_seq_params.PackedSeqParams] = None,
sequence_len_offset: Optional[int] = None,
*,
inference_params: Optional[megatron.core.inference.contexts.BaseInferenceContext] = None,
**kwargs,
) tuple[torch.Tensor, torch.Tensor | None]#

Perform a forward pass through the GDN module.

Returns:

(tuple[torch.Tensor, torch.Tensor]) GDN output and bias.