core.transformer.wide_residual_layer#
Construction-time wide-residual connections and model-boundary helpers.
Module Contents#
Classes#
Bounded learned retention with one controller per full-width stream. |
|
Positive scalar controller for each contiguous full-width stream. |
|
Positive streamwise maps around one ordinary-width residual branch. |
|
Transformer layer carrying a wide stream around ordinary-width branches. |
|
Learned streamwise readout from D’ to D before final normalization. |
Functions#
Create a parameter vector padded for small distributed-optimizer shards. |
|
Return the active prefix of a potentially padded vector parameter. |
|
Attach optimizer and TP-gradient metadata to a replicated residual map. |
|
Mark a replicated controller for the correct TP gradient reduction. |
|
Attach optimizer and TP-gradient metadata to a replicated retention parameter. |
|
Build the streamwise model-boundary readout when wide residuals are enabled. |
|
Replicate an ordinary-width activation into contiguous residual streams. |
Data#
API#
- core.transformer.wide_residual_layer._MIN_CONTROLLER_NUMEL#
128
- core.transformer.wide_residual_layer._make_padded_vector(values: torch.Tensor) torch.nn.Parameter#
Create a parameter vector padded for small distributed-optimizer shards.
- core.transformer.wide_residual_layer._active_vector(param: torch.Tensor, length: int) torch.Tensor#
Return the active prefix of a potentially padded vector parameter.
- core.transformer.wide_residual_layer._mark_residual_map_parameter(
- param: torch.nn.Parameter,
- config: megatron.core.transformer.transformer_config.TransformerConfig,
Attach optimizer and TP-gradient metadata to a replicated residual map.
- core.transformer.wide_residual_layer._mark_replicated_tp_parameter(
- param: torch.nn.Parameter,
- config: megatron.core.transformer.transformer_config.TransformerConfig,
Mark a replicated controller for the correct TP gradient reduction.
- core.transformer.wide_residual_layer._mark_retention_parameter(
- param: torch.nn.Parameter,
- config: megatron.core.transformer.transformer_config.TransformerConfig,
Attach optimizer and TP-gradient metadata to a replicated retention parameter.
- class core.transformer.wide_residual_layer.LearnedWideResidualRetention(
- config: megatron.core.transformer.transformer_config.TransformerConfig,
- layer_number: int,
- branch_name: str,
- *,
- num_streams: int,
Bases:
torch.nn.ModuleBounded learned retention with one controller per full-width stream.
Initialization
- factors() torch.Tensor#
Return one FP32 retention factor for every full-width stream.
- forward(*, return_logits: bool = False) torch.Tensor#
Return padded logits or materialized retention factors.
- class core.transformer.wide_residual_layer.StreamwiseSigmoidMap(
- config: megatron.core.transformer.transformer_config.TransformerConfig,
- *,
- map_kind: Literal[read, write],
Bases:
torch.nn.ModulePositive scalar controller for each contiguous full-width stream.
Initialization
- factors() torch.Tensor#
Return one positive factor for every full-width residual stream.
- forward(*, return_logits: bool = False) torch.Tensor#
Return padded logits or one positive factor per full-width stream.
- class core.transformer.wide_residual_layer.StreamwiseSigmoidWideResidualConnection(
- config: megatron.core.transformer.transformer_config.TransformerConfig,
- layer_number: int,
- branch_name: str,
- pg_collection: megatron.core.process_groups_config.ProcessGroupCollection,
- name: str | None = None,
Bases:
megatron.core.transformer.residual_connection.ResidualConnectionPositive streamwise maps around one ordinary-width residual branch.
Initialization
- _read(
- hidden_states: torch.Tensor,
- _write(
- branch_output: megatron.core.transformer.residual_connection.ResidualBranchOutput,
- state: megatron.core.transformer.residual_connection.ResidualConnectionState,
- *,
- dropout_probability: float,
- training: bool,
- class core.transformer.wide_residual_layer.WideResidualTransformerLayer(
- config: megatron.core.transformer.transformer_config.TransformerConfig,
- submodules: megatron.core.transformer.transformer_layer.TransformerLayerSubmodules,
- layer_number: int = 1,
- hidden_dropout: Optional[float] = None,
- pg_collection: Optional[megatron.core.process_groups_config.ProcessGroupCollection] = None,
- vp_stage: Optional[int] = None,
- is_mtp_layer: bool = False,
- add_layer_offset: bool = True,
- pp_layer_offset: Optional[int] = None,
- name: str | None = None,
Bases:
megatron.core.transformer.transformer_layer.TransformerLayerTransformer layer carrying a wide stream around ordinary-width branches.
Initialization
- Parameters:
name (str | None) – module instance name passed top-down from its paranet module
- supports_wide_residual_connections: bool#
True
- _get_self_attention_residual_connection() megatron.core.transformer.residual_connection.ResidualConnection | None#
Return the connection surrounding the self-attention branch.
- _get_mlp_residual_connection() megatron.core.transformer.residual_connection.ResidualConnection | None#
Return the connection surrounding the MLP or MoE branch.
- class core.transformer.wide_residual_layer.StreamwiseSigmoidResidualReadout(
- config: megatron.core.transformer.transformer_config.TransformerConfig,
Bases:
torch.nn.ModuleLearned streamwise readout from D’ to D before final normalization.
Initialization
- forward(hidden_states: torch.Tensor) torch.Tensor#
Mix the full-width streams into one backbone-width activation.
- core.transformer.wide_residual_layer.build_wide_residual_readout(
- config: megatron.core.transformer.transformer_config.TransformerConfig,
Build the streamwise model-boundary readout when wide residuals are enabled.
- core.transformer.wide_residual_layer.expand_wide_residual_stream(
- hidden_states: torch.Tensor,
- num_streams: int,
Replicate an ordinary-width activation into contiguous residual streams.