| 422 | |
| 423 | @si_module |
| 424 | class 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 |
nothing calls this directly
no outgoing calls
no test coverage detected