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,
Perform a forward pass through the GDN module.
- Returns:
(tuple[torch.Tensor, torch.Tensor]) GDN output and bias.