core.transformer.residual_connection#

Runtime protocol for architecture-specific residual connections.

Module Contents#

Classes#

ResidualConnection

Read one branch input from, then write its update to, a residual stream.

Data#

API#

core.transformer.residual_connection.ResidualBranchOutput: TypeAlias#

None

core.transformer.residual_connection.ResidualConnectionState: TypeAlias#

None

core.transformer.residual_connection.ResidualConnectionWriteState: TypeAlias#

None

core.transformer.residual_connection.ResidualConnectionOperation: TypeAlias#

None

class core.transformer.residual_connection.ResidualConnection(
residual_stream_hidden_size: int,
branch_hidden_size: int,
)#

Bases: torch.nn.Module, abc.ABC

Read one branch input from, then write its update to, a residual stream.

forward validates the common contract and retains the incoming residual stream as the first tensor in ResidualConnectionState. Concrete connections own bias, dropout, mapping, and residual-update semantics.

Initialization

property residual_stream_hidden_size: int#

Hidden width accepted and returned by this connection.

property branch_hidden_size: int#

Hidden width produced by the read operation for the wrapped branch.

forward(
value: core.transformer.residual_connection.ResidualBranchOutput,
*,
operation: core.transformer.residual_connection.ResidualConnectionOperation,
state: core.transformer.residual_connection.ResidualConnectionState | None = None,
fp32_residual_connection: bool = False,
dropout_probability: float | None = None,
training: bool | None = None,
) → torch.Tensor | tuple[torch.Tensor, core.transformer.residual_connection.ResidualConnectionState]#

Execute one residual operation through the standard module call path.

_read_with_validation(
hidden_states: torch.Tensor,
*,
fp32_residual_connection: bool,
) → tuple[torch.Tensor, core.transformer.residual_connection.ResidualConnectionState]#
_write_with_validation(
branch_output: core.transformer.residual_connection.ResidualBranchOutput,
state: core.transformer.residual_connection.ResidualConnectionState,
*,
dropout_probability: float,
training: bool,
) → torch.Tensor#
static residual_stream(
state: core.transformer.residual_connection.ResidualConnectionState,
) → torch.Tensor#

Return the carried residual stream from a validated connection state.

static _validate_tensor_state(
state: tuple[torch.Tensor, ...],
*,
state_name: str,
allow_empty: bool,
) → None#
static _validate_branch_output(
branch_output: core.transformer.residual_connection.ResidualBranchOutput,
) → None#
abstractmethod _read(
hidden_states: torch.Tensor,
) → tuple[torch.Tensor, core.transformer.residual_connection.ResidualConnectionWriteState]#

Return the branch input and state needed only by the later write.

abstractmethod _write(
branch_output: core.transformer.residual_connection.ResidualBranchOutput,
state: core.transformer.residual_connection.ResidualConnectionState,
*,
dropout_probability: float,
training: bool,
) → torch.Tensor#

Implement the architecture-specific residual-stream update.