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

Class GaussianZ

tokenizer.py:424–466  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

422
423@si_module
424class GaussianZ(nn.Module):
425 class Config:
426 dim: int
427 latent_dim: int
428 bias: bool = False
429 use_weight_norm: bool = False
430
431 def __init__(self, c: Config):
432 super().__init__()
433
434 self.proj_in = nn.Linear(c.dim, c.latent_dim * 2, bias=c.bias)
435 self.proj_out = nn.Linear(c.latent_dim, c.dim, bias=c.bias)
436
437 if c.use_weight_norm:
438 self.proj_in = weight_norm(self.proj_in)
439 self.proj_out = weight_norm(self.proj_out)
440
441 def reparam(self, mu, logvar):
442 std = T.exp(logvar / 2)
443 eps = T.randn_like(std)
444 return mu + eps * std
445
446 def kl_divergence(self, mu, logvar):
447 return T.mean(-0.5 * T.sum(
448 1 + logvar - mu.pow(2) - logvar.exp(),
449 dim=(1, 2))
450 )
451
452 def repr_from_latent(self, latent: Union[dict, T.Tensor]):
453 if isinstance(latent, T.Tensor):
454 z = latent
455 else:
456 z = self.reparam(latent['mu'], latent['logvar'])
457 l = self.proj_out(z)
458 return l
459
460 def forward(self, x: T.Tensor) -> Tuple[T.Tensor, dict]:
461 mu, logvar = self.proj_in(x).chunk(2, dim=-1)
462 kl_div = self.kl_divergence(mu, logvar)
463 z = self.reparam(mu, logvar)
464 xhat = self.proj_out(z)
465 latent = {'mu': mu, 'logvar': logvar, 'z': z, 'kl_divergence': kl_div}
466 return xhat, latent
467
468
469

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected