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

Method __init__

tokenizer.py:321–361  ·  view source on GitHub ↗
(self, c: Config)

Source from the content-addressed store, hash-verified

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:
361 self.multiplier_proj = Conv1d1x1(self.middle_channels, self.middle_channels * self.multiplier, bias=c.bias)
362
363 def forward(self, x, return_feats=False):
364 if self.c.latent_dim is not None and self.mode == 'decoder':

Callers

nothing calls this directly

Calls 4

ResNetBlockClass · 0.85
CausalConv1dClass · 0.85
Conv1d1x1Function · 0.85
__init__Method · 0.45

Tested by

no test coverage detected