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

Class ResNetStack

tokenizer.py:303–419  ·  view source on GitHub ↗

ResNet encoder or decoder stack. Channel ratios and strides take the default order of from data/io-layer, to the middle of the model.

Source from the content-addressed store, hash-verified

301
302@si_module
303class ResNetStack(nn.Module):
304 """
305 ResNet encoder or decoder stack. Channel ratios
306 and strides take the default order of from
307 data/io-layer, to the middle of the model.
308 """
309 class Config:
310 input_channels: int = 1
311 output_channels: int = 1
312 encode_channels: int = 32
313 decode_channel_multiplier: int = 1
314 latent_dim: int = None
315 kernel_size: int = 7
316 bias: bool = True
317 channel_ratios: Tuple[int, ...] = (2, 4, 8, 16)
318 strides: Tuple[int, ...] = (3, 4, 5, 5)
319 mode: Literal['encoder', 'decoder'] = 'encoder'
320
321 def __init__(self, c: Config):
322 super().__init__()
323 assert c.mode in ('encoder', 'decoder'), f"Mode ({c.mode}) is not supported!"
324
325 self.mode = c.mode
326
327 assert len(c.channel_ratios) == len(c.strides)
328 channel_ratios = (1,) + c.channel_ratios
329 strides = c.strides
330 self.middle_channels = c.encode_channels * channel_ratios[-1]
331 if c.mode == 'decoder':
332 channel_ratios = tuple(reversed(channel_ratios))
333 strides = tuple(reversed(strides))
334
335 self.multiplier = c.decode_channel_multiplier if c.mode == 'decoder' else 1
336 res_blocks = [ResNetBlock(
337 c.encode_channels * channel_ratios[s_idx] * self.multiplier,
338 c.encode_channels * channel_ratios[s_idx+1] * self.multiplier,
339 stride,
340 kernel_size=c.kernel_size,
341 bias=c.bias,
342 mode=c.mode,
343 ) for s_idx, stride in enumerate(strides)]
344
345 data_conv = CausalConv1d(
346 in_channels=c.input_channels if c.mode == 'encoder' else c.encode_channels * self.multiplier,
347 out_channels=c.encode_channels if c.mode == 'encoder' else c.output_channels,
348 kernel_size=c.kernel_size,
349 stride=1,
350 bias=False,
351 )
352
353 if c.mode == 'encoder':
354 self.res_stack = nn.Sequential(data_conv, *res_blocks)
355 elif c.mode == 'decoder':
356 self.res_stack = nn.Sequential(*res_blocks, data_conv)
357
358 if c.latent_dim is not None:
359 self.latent_proj = Conv1d1x1(self.middle_channels, c.latent_dim, bias=c.bias) if c.mode == 'encoder' else Conv1d1x1(c.latent_dim, self.middle_channels, bias=c.bias)
360 if self.multiplier != 1:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected