MCPcopy Create free account
hub / github.com/Standard-Intelligence/hertz-dev / ResNetBlock

Class ResNetBlock

tokenizer.py:230–297  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

228
229
230class ResNetBlock(nn.Module):
231 def __init__(self,
232 in_channels,
233 out_channels,
234 stride,
235 kernel_size=7,
236 dilations=(1, 3, 9),
237 bias=True,
238 mode='encoder',
239 ):
240 super().__init__()
241 assert mode in ('encoder', 'decoder'), f"Mode ({mode}) is not supported!"
242
243 self.mode = mode
244 self.stride = stride
245
246 ConvUnit = CausalConv1d if mode == 'encoder' else CausalConvTranspose1d
247
248 res_channels = in_channels if mode == 'encoder' else out_channels
249
250 res_units = [CausalResUnit(
251 res_channels,
252 res_channels,
253 kernel_size=kernel_size,
254 dilation=dilation,
255 ) for dilation in dilations]
256
257 if in_channels == out_channels:
258 if mode == 'encoder':
259 self.pool = nn.AvgPool1d(kernel_size=stride, stride=stride)
260 if mode == 'decoder':
261 self.upsample = nn.Upsample(scale_factor=stride, mode='nearest')
262 conv_unit = nn.Conv1d(
263 in_channels=in_channels,
264 out_channels=out_channels,
265 kernel_size=1,
266 bias=bias,
267 ) if in_channels != out_channels else nn.Identity()
268 else:
269 conv_unit = ConvUnit(
270 in_channels=in_channels,
271 out_channels=out_channels,
272 kernel_size=(2 * stride),
273 stride=stride,
274 bias=bias,
275 )
276
277 if mode == 'encoder':
278 if in_channels == out_channels:
279 self.res_block = nn.Sequential(*res_units, self.pool, conv_unit)
280 else:
281 self.res_block = nn.Sequential(*res_units, conv_unit)
282 elif mode == 'decoder':
283 if in_channels == out_channels:
284 self.res_block = nn.Sequential(self.upsample, conv_unit, *res_units)
285 else:
286 self.res_block = nn.Sequential(conv_unit, *res_units)
287

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected