core.transformer.wide_residual_layer#

Construction-time wide-residual connections and model-boundary helpers.

Module Contents#

Classes#

LearnedWideResidualRetention

Bounded learned retention with one controller per full-width stream.

StreamwiseSigmoidMap

Positive scalar controller for each contiguous full-width stream.

StreamwiseSigmoidWideResidualConnection

Positive streamwise maps around one ordinary-width residual branch.

WideResidualTransformerLayer

Transformer layer carrying a wide stream around ordinary-width branches.

StreamwiseSigmoidResidualReadout

Learned streamwise readout from D’ to D before final normalization.

Functions#

_make_padded_vector

Create a parameter vector padded for small distributed-optimizer shards.

_active_vector

Return the active prefix of a potentially padded vector parameter.

_mark_residual_map_parameter

Attach optimizer and TP-gradient metadata to a replicated residual map.

_mark_replicated_tp_parameter

Mark a replicated controller for the correct TP gradient reduction.

_mark_retention_parameter

Attach optimizer and TP-gradient metadata to a replicated retention parameter.

build_wide_residual_readout

Build the streamwise model-boundary readout when wide residuals are enabled.

expand_wide_residual_stream

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,
) → None#

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,
) → None#

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,
) → None#

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.Module

Bounded 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.Module

Positive 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.ResidualConnection

Positive streamwise maps around one ordinary-width residual branch.

Initialization

_read(
hidden_states: torch.Tensor,
) → tuple[torch.Tensor, megatron.core.transformer.residual_connection.ResidualConnectionWriteState]#
_write(
branch_output: megatron.core.transformer.residual_connection.ResidualBranchOutput,
state: megatron.core.transformer.residual_connection.ResidualConnectionState,
*,
dropout_probability: float,
training: bool,
) → torch.Tensor#
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.TransformerLayer

Transformer 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.Module

Learned 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,
) → core.transformer.wide_residual_layer.StreamwiseSigmoidResidualReadout | None#

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

Replicate an ordinary-width activation into contiguous residual streams.