77"""
88from __future__ import annotations
99
10+ from typing import Sequence
11+
1012import torch
1113import torch .nn as nn
1214
@@ -44,27 +46,96 @@ def set_lambda(self, lambda_: float):
4446# ----------------------------
4547
4648
49+ def _make_activation (name : str | None ) -> nn .Module | None :
50+ if name is None :
51+ return None
52+ name = name .lower ()
53+ if name == "relu" :
54+ return nn .ReLU ()
55+ if name == "tanh" :
56+ return nn .Tanh ()
57+ if name == "sigmoid" :
58+ return nn .Sigmoid ()
59+ if name in {"silu" , "swish" }:
60+ return nn .SiLU ()
61+ if name == "gelu" :
62+ return nn .GELU ()
63+ raise ValueError (f"Unsupported activation: { name } " )
64+
65+
4766class ResidualBlock (nn .Module ):
48- """
49- Residual block with skip connection: x + FFN(x) + LayerNorm
50- """
67+ """Feed-forward residual block with optional projection skip."""
5168
52- def __init__ (self , hidden_size : int , dropout : float = 0.0 ):
69+ def __init__ (
70+ self ,
71+ in_features : int ,
72+ out_features : int ,
73+ * ,
74+ hidden_features : int | None = None ,
75+ dropout : float = 0.0 ,
76+ activation : str = "silu" ,
77+ use_layer_norm : bool = True ,
78+ ) -> None :
5379 super ().__init__ ()
54- self .ffn = nn .Sequential (
55- nn .Linear (hidden_size , hidden_size ),
56- nn .SiLU (),
57- nn .Dropout (dropout ),
58- nn .Linear (hidden_size , hidden_size ),
59- nn .Dropout (dropout ),
60- )
61- self .layer_norm = nn .LayerNorm (hidden_size )
80+ hidden_features = out_features if hidden_features is None else hidden_features
81+
82+ ff_layers : list [nn .Module ] = [nn .Linear (in_features , hidden_features )]
83+ act = _make_activation (activation )
84+ if act is not None :
85+ ff_layers .append (act )
86+ if dropout > 0 :
87+ ff_layers .append (nn .Dropout (dropout ))
88+ ff_layers .append (nn .Linear (hidden_features , out_features ))
89+ if dropout > 0 :
90+ ff_layers .append (nn .Dropout (dropout ))
91+ self .ffn = nn .Sequential (* ff_layers )
92+
93+ if in_features == out_features :
94+ self .shortcut : nn .Module = nn .Identity ()
95+ else :
96+ self .shortcut = nn .Linear (in_features , out_features , bias = False )
6297
63- def forward (self , x ):
64- return self .layer_norm (x + self .ffn (x ))
98+ self .layer_norm = nn .LayerNorm (out_features ) if use_layer_norm else nn .Identity ()
99+
100+ def forward (self , x : torch .Tensor ) -> torch .Tensor :
101+ return self .layer_norm (self .shortcut (x ) + self .ffn (x ))
65102
66103
67- def make_mlp (sizes , dropout = 0.0 , last_activation = None , use_residual = False ):
104+ class ResidualStack (nn .Module ):
105+ """Stack of residual blocks that can change feature dimensionality."""
106+
107+ def __init__ (
108+ self ,
109+ sizes : Sequence [int ],
110+ * ,
111+ dropout : float = 0.0 ,
112+ activation : str = "silu" ,
113+ ) -> None :
114+ super ().__init__ ()
115+ if len (sizes ) < 2 :
116+ raise ValueError ("ResidualStack requires at least two layer sizes" )
117+
118+ blocks = [
119+ ResidualBlock (
120+ in_features = in_f ,
121+ out_features = out_f ,
122+ dropout = dropout ,
123+ activation = activation ,
124+ )
125+ for in_f , out_f in zip (sizes [:- 1 ], sizes [1 :])
126+ ]
127+ self .blocks = nn .Sequential (* blocks )
128+
129+ def forward (self , x : torch .Tensor ) -> torch .Tensor :
130+ return self .blocks (x )
131+
132+
133+ def make_mlp (
134+ sizes : Sequence [int ],
135+ dropout : float = 0.0 ,
136+ last_activation : str | None = None ,
137+ use_residual : bool = False ,
138+ ):
68139 """
69140 Hidden layers: Linear -> LayerNorm -> SiLU -> Dropout
70141 Output layer: optional activation per arg
@@ -78,52 +149,34 @@ def make_mlp(sizes, dropout=0.0, last_activation=None, use_residual=False):
78149 "Residual networks need at least input, hidden, and output layers"
79150 )
80151
81- # Check that all hidden layers have the same size for residual connections
82- hidden_sizes = sizes [1 :- 1 ]
83- if len (set (hidden_sizes )) > 1 :
84- raise ValueError (
85- f"For residual connections, all hidden layer sizes must be the same. Got: { hidden_sizes } "
86- )
87-
88- hidden_size = hidden_sizes [0 ]
89- num_hidden_layers = len (hidden_sizes )
90-
91- layers = []
92-
93- # Input projection to hidden size
94- layers .append (nn .Linear (sizes [0 ], hidden_size ))
95- layers .extend ([nn .LayerNorm (hidden_size ), nn .SiLU (), nn .Dropout (dropout )])
96-
97- # Residual blocks
98- for _ in range (num_hidden_layers ):
99- layers .append (ResidualBlock (hidden_size , dropout ))
100-
101- # Output layer
102- layers .append (nn .Linear (hidden_size , sizes [- 1 ]))
103- if last_activation == "relu" :
104- layers .append (nn .ReLU ())
105- elif last_activation == "tanh" :
106- layers .append (nn .Tanh ())
107- elif last_activation == "sigmoid" :
108- layers .append (nn .Sigmoid ())
109-
152+ trunk_sizes = sizes [:- 1 ]
153+ layers : list [nn .Module ] = [
154+ ResidualStack (trunk_sizes , dropout = dropout , activation = "silu" )
155+ ]
156+ layers .append (nn .Linear (trunk_sizes [- 1 ], sizes [- 1 ]))
157+ final_act = _make_activation (last_activation )
158+ if final_act is not None :
159+ layers .append (final_act )
110160 return nn .Sequential (* layers )
111161
112162 else :
113163 # Original implementation
114- layers = []
164+ layers : list [ nn . Module ] = []
115165 for i in range (len (sizes ) - 1 ):
116166 in_f , out_f = sizes [i ], sizes [i + 1 ]
117167 layers .append (nn .Linear (in_f , out_f ))
118- if i < len (sizes ) - 2 :
119- layers += [nn .LayerNorm (out_f ), nn .SiLU (), nn .Dropout (dropout )]
168+ is_last = i == len (sizes ) - 2
169+ if not is_last :
170+ layers .append (nn .LayerNorm (out_f ))
171+ act = _make_activation ("silu" )
172+ if act is not None :
173+ layers .append (act )
174+ if dropout > 0 :
175+ layers .append (nn .Dropout (dropout ))
120176 else :
121- if last_activation == "relu" :
122- layers += [nn .ReLU ()]
123- elif last_activation == "tanh" :
124- layers += [nn .Tanh ()]
125- elif last_activation == "sigmoid" :
126- layers += [nn .Sigmoid ()]
177+ final_act = _make_activation (last_activation )
178+ if final_act is not None :
179+ layers .append (final_act )
127180 return nn .Sequential (* layers )
128181
129182
0 commit comments