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

Class WaveCodec

tokenizer.py:471–538  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

469
470@si_module
471class WaveCodec(nn.Module):
472 class Config:
473 resnet_config: ResNetStack.Config = None
474 sample_rate: int = 16_000
475 use_weight_norm: bool = False
476
477 compressor_config: dataclass = None
478
479 norm_stddev: float = 1.0
480
481 def __init__(self, c: Config):
482 super().__init__()
483 self.norm_stddev = c.norm_stddev
484 self.encoder = c.resnet_config(mode='encoder')
485 self.sample_rate = c.sample_rate
486
487 self.total_stride = 1
488 for stride in c.resnet_config.strides:
489 self.total_stride *= stride
490 self.tokens_per_second = self.sample_rate / self.total_stride
491
492 self.compressor = c.compressor_config(dim=self.encoder.middle_channels)
493
494 self.decoder = c.resnet_config(mode='decoder')
495
496 if c.use_weight_norm:
497 self.encoder.apply_weight_norm()
498 self.decoder.apply_weight_norm()
499 self.encoder.reset_parameters()
500 self.decoder.reset_parameters()
501
502 def encode(self, data):
503 return self.encoder(data/self.norm_stddev)
504
505 def decode(self, latent):
506 return self.decoder(latent.transpose(1, 2))*self.norm_stddev
507
508 @T.no_grad()
509 def latent_from_data(self, data, get_parameters=False):
510 x = self.encode(data)
511 l_in = x.transpose(1, 2)
512 l, latent = self.compressor(l_in)
513 return latent['z'] if not get_parameters else {
514 'mu': latent['mu'],
515 'logvar': latent['logvar'],
516 'z': latent['z'],
517 }
518
519 @T.no_grad()
520 def data_from_latent(self, latent):
521 l = self.compressor.repr_from_latent(latent)
522 x = self.decode(l)
523 return x
524
525 def process(self, x):
526 return self.latent_from_data(x)
527
528 def unprocess(self, latent):

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected