Skip to content

[Design proposal]: shared residual-connection integration for mHC and Wide Residual #7888

Description

@kwangjunahn

Proposal

Share the layer-integration machinery for mHC and Wide Residual while preserving their separate algorithms, kernels, and supported behavior. The goal is to reduce duplicated code and make future integration changes easier to review and maintain. This shared structure could also make it easier to implement other wide-residual architectures from the literature without duplicating the layer-integration code.

Motivation

Both implementations read an ordinary-width input from a wide residual stream, execute a branch, and write its update back to the stream. WR uses a ResidualConnection contract through shared Transformer hooks. mHC has separate layer overrides that repeat related normalization, checkpoint, and offload bookkeeping.

HybridStack also coordinates separate mHC and WR replay paths, although both use CheckpointWithoutOutputManager. Sharing these mechanics could reduce duplicated maintenance without changing the mathematical operations.

Proposed Structure

Use a common connection lifecycle, conceptually:

branch_input, state = connection.read(wide_stream)
branch_update = branch(norm(branch_input))
next_stream = connection.write(branch_update, state)

The connection owns its maps, carried tensor state, bias/dropout semantics, and residual update. The layer owns branch execution. Replay coordination owns block boundaries and manager lifetime.

WR will retain its token-independent streamwise maps and optional discounting/retention. mHC will retain its input-dependent maps, Sinkhorn normalization, and cross-stream mixing. Sharing the layer-integration structure should also make it easier to add other wide-residual architectures emerging in the literature while keeping their algorithm-specific behavior explicit.

Start from the existing connection contract, extending it only where necessary to preserve mHC's state and execution requirements.

Plans for Incremental PRs

Stage Scope Expected Result
1. Connection contract Expose mHC's existing operations through a compatible interface mHC caller uses the contract while retaining its execution schedule.
2. Transformer execution Reuse shared read, normalization, and writeback hooks, retaining necessary architecture-specific policies. Less duplicated layer code, with existing mHC and WR modes preserved.
3. Replay coordination Share block planning, manager creation, and finalization; preserve each algorithm's operation-level checkpoint policy. Common coordination with unchanged replay semantics.
4. HybridStack orchestration Simplify dispatch and expansion/readout integration. A more generic stack without changing model boundaries.

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

enhancementNew feature or request

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions