Coverage for /opt/conda/lib/python3.14/site-packages/medil/_vae.py: 95%
103 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-23 02:29 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-07-23 02:29 +0000
1"""Variational autoencoder components for NeuroCausalFactorAnalysis."""
3import math
5import torch
6from torch import nn
7from torch.nn import functional as F
10class VariationalAutoencoder(nn.Module):
11 def __init__(
12 self,
13 num_latent,
14 num_meas,
15 num_hidden_layers,
16 latent_width,
17 meas_width,
18 biadj=None,
19 encoder_hidden_dim=None,
20 num_classes=1,
21 ):
22 super().__init__()
24 if encoder_hidden_dim is None:
25 encoder_hidden_dim = max(num_meas, 64)
27 self.encoder = Encoder(
28 num_latent=num_latent * latent_width,
29 num_meas=num_meas,
30 hidden_dim=encoder_hidden_dim,
31 )
33 self.decoder = Decoder(
34 num_latent=num_latent,
35 num_meas=num_meas,
36 num_hidden_layers=num_hidden_layers,
37 latent_width=latent_width,
38 meas_width=meas_width,
39 biadj=biadj,
40 num_classes=num_classes,
41 )
43 def forward(self, x):
44 mu, logvar = self.encoder(x)
45 z = self.latent_sample(mu, logvar)
46 x_recon = self.decoder(z)
47 return x_recon, mu, logvar
49 def latent_sample(self, mu, logvar):
50 if self.training:
51 std = torch.exp(0.5 * logvar)
52 eps = torch.randn_like(std)
53 return mu + eps * std
54 return mu
57class Encoder(nn.Module):
58 def __init__(self, num_latent, num_meas, hidden_dim=64):
59 super().__init__()
60 self.enc1 = nn.Linear(num_meas, hidden_dim)
61 self.bn1 = nn.BatchNorm1d(hidden_dim)
62 self.enc2 = nn.Linear(hidden_dim, hidden_dim)
63 self.bn2 = nn.BatchNorm1d(hidden_dim)
64 self.activation = nn.GELU()
65 self.fc_mu = nn.Linear(hidden_dim, num_latent)
66 self.fc_logvar = nn.Linear(hidden_dim, num_latent)
68 def forward(self, x):
69 h = self.activation(self.bn1(self.enc1(x)))
70 h = self.activation(self.bn2(self.enc2(h)))
71 mu = self.fc_mu(h)
72 logvar = self.fc_logvar(h)
73 return mu, logvar
76class Decoder(nn.Module):
77 def __init__(
78 self,
79 num_latent,
80 num_meas,
81 num_hidden_layers,
82 latent_width,
83 meas_width,
84 biadj=None,
85 num_classes=1,
86 ):
87 super().__init__()
89 self.num_latent = num_latent
90 self.num_meas = num_meas
91 self.num_classes = num_classes
92 self.latent_width = latent_width
93 self.meas_width = meas_width
95 self.latent_dim = num_latent * latent_width
96 self.hidden_dim = num_meas * meas_width
98 if biadj is None:
99 biadj = torch.ones(num_meas, num_latent)
100 else:
101 biadj = torch.as_tensor(biadj, dtype=torch.float32)
103 first_mask = self._expand_biadj(biadj, meas_width, latent_width)
104 hidden_mask = self._make_hidden_block_mask(num_meas, meas_width)
105 output_mask = self._make_output_mask(num_meas, meas_width, num_classes)
107 self.linear_in = SparseLinear(
108 in_features=self.latent_dim,
109 out_features=self.hidden_dim,
110 mask=first_mask,
111 )
113 self.bn_in = nn.BatchNorm1d(self.hidden_dim)
115 self.hidden_layers = nn.ModuleList(
116 [
117 SparseLinear(
118 in_features=self.hidden_dim,
119 out_features=self.hidden_dim,
120 mask=hidden_mask,
121 )
122 for _ in range(num_hidden_layers)
123 ]
124 )
126 self.hidden_bns = nn.ModuleList(
127 [nn.BatchNorm1d(self.hidden_dim) for _ in range(num_hidden_layers)]
128 )
130 self.linear_out = SparseLinear(
131 in_features=self.hidden_dim,
132 out_features=self.num_meas * num_classes,
133 mask=output_mask,
134 )
136 self.activation = nn.GELU()
138 @staticmethod
139 def _expand_biadj(biadj, meas_width, latent_width):
140 return biadj.repeat_interleave(meas_width, dim=0).repeat_interleave(
141 latent_width, dim=1
142 )
144 @staticmethod
145 def _make_hidden_block_mask(num_meas, width_per_meas):
146 block = torch.ones(width_per_meas, width_per_meas)
147 blocks = [block for _ in range(num_meas)]
148 return torch.block_diag(*blocks)
150 @staticmethod
151 def _make_output_mask(num_meas, width_per_meas, num_classes=1):
152 block = torch.ones(num_classes, width_per_meas)
153 blocks = [block for _ in range(num_meas)]
154 return torch.block_diag(*blocks)
156 def forward(self, z):
157 h = self.activation(self.bn_in(self.linear_in(z)))
159 for layer, bn in zip(self.hidden_layers, self.hidden_bns):
160 h = self.activation(bn(layer(h)))
162 x_recon = self.linear_out(h)
163 return x_recon
166class SparseLinear(nn.Module):
167 def __init__(
168 self,
169 in_features,
170 out_features,
171 mask=None,
172 bias=True,
173 device=None,
174 dtype=None,
175 ):
176 super().__init__()
177 factory_kwargs = {"device": device, "dtype": dtype}
179 self.in_features = in_features
180 self.out_features = out_features
182 if mask is None:
183 mask = torch.ones(out_features, in_features)
184 else:
185 if mask.shape != (out_features, in_features):
186 raise ValueError(
187 f"mask must have shape {(out_features, in_features)}, "
188 f"got {tuple(mask.shape)}"
189 )
191 self.register_buffer("mask", mask.float())
193 self.weight = nn.Parameter(
194 torch.empty((out_features, in_features), **factory_kwargs)
195 )
197 if bias:
198 self.bias = nn.Parameter(torch.empty(out_features, **factory_kwargs))
199 else:
200 self.register_parameter("bias", None)
202 self.reset_parameters()
204 def reset_parameters(self):
205 nn.init.orthogonal_(self.weight)
206 if self.bias is not None:
207 fan_in, _ = nn.init._calculate_fan_in_and_fan_out(self.weight)
208 bound = 1 / math.sqrt(fan_in) if fan_in > 0 else 0
209 nn.init.uniform_(self.bias, -bound, bound)
211 def forward(self, x):
212 return F.linear(x, self.weight * self.mask, self.bias)