core.transformer.residual_connection#
Runtime protocol for architecture-specific residual connections.
Module Contents#
Classes#
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.ABCRead one branch input from, then write its update to, a residual stream.
forwardvalidates the common contract and retains the incoming residual stream as the first tensor inResidualConnectionState. Concrete connections own bias, dropout, mapping, and residual-update semantics.Initialization
Hidden width accepted and returned by this connection.
Hidden width produced by the
readoperation 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,
Execute one residual operation through the standard module call path.
- _read_with_validation(
- hidden_states: torch.Tensor,
- *,
- fp32_residual_connection: bool,
- _write_with_validation(
- branch_output: core.transformer.residual_connection.ResidualBranchOutput,
- state: core.transformer.residual_connection.ResidualConnectionState,
- *,
- dropout_probability: float,
- training: bool,
- static residual_stream( ) 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,
- static _validate_branch_output(
- branch_output: core.transformer.residual_connection.ResidualBranchOutput,
- abstractmethod _read(
- hidden_states: torch.Tensor,
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,
Implement the architecture-specific residual-stream update.