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

1"""Variational autoencoder components for NeuroCausalFactorAnalysis.""" 

2 

3import math 

4 

5import torch 

6from torch import nn 

7from torch.nn import functional as F 

8 

9 

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__() 

23 

24 if encoder_hidden_dim is None: 

25 encoder_hidden_dim = max(num_meas, 64) 

26 

27 self.encoder = Encoder( 

28 num_latent=num_latent * latent_width, 

29 num_meas=num_meas, 

30 hidden_dim=encoder_hidden_dim, 

31 ) 

32 

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 ) 

42 

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 

48 

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 

55 

56 

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) 

67 

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 

74 

75 

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__() 

88 

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 

94 

95 self.latent_dim = num_latent * latent_width 

96 self.hidden_dim = num_meas * meas_width 

97 

98 if biadj is None: 

99 biadj = torch.ones(num_meas, num_latent) 

100 else: 

101 biadj = torch.as_tensor(biadj, dtype=torch.float32) 

102 

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) 

106 

107 self.linear_in = SparseLinear( 

108 in_features=self.latent_dim, 

109 out_features=self.hidden_dim, 

110 mask=first_mask, 

111 ) 

112 

113 self.bn_in = nn.BatchNorm1d(self.hidden_dim) 

114 

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 ) 

125 

126 self.hidden_bns = nn.ModuleList( 

127 [nn.BatchNorm1d(self.hidden_dim) for _ in range(num_hidden_layers)] 

128 ) 

129 

130 self.linear_out = SparseLinear( 

131 in_features=self.hidden_dim, 

132 out_features=self.num_meas * num_classes, 

133 mask=output_mask, 

134 ) 

135 

136 self.activation = nn.GELU() 

137 

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 ) 

143 

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) 

149 

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) 

155 

156 def forward(self, z): 

157 h = self.activation(self.bn_in(self.linear_in(z))) 

158 

159 for layer, bn in zip(self.hidden_layers, self.hidden_bns): 

160 h = self.activation(bn(layer(h))) 

161 

162 x_recon = self.linear_out(h) 

163 return x_recon 

164 

165 

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} 

178 

179 self.in_features = in_features 

180 self.out_features = out_features 

181 

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 ) 

190 

191 self.register_buffer("mask", mask.float()) 

192 

193 self.weight = nn.Parameter( 

194 torch.empty((out_features, in_features), **factory_kwargs) 

195 ) 

196 

197 if bias: 

198 self.bias = nn.Parameter(torch.empty(out_features, **factory_kwargs)) 

199 else: 

200 self.register_parameter("bias", None) 

201 

202 self.reset_parameters() 

203 

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) 

210 

211 def forward(self, x): 

212 return F.linear(x, self.weight * self.mask, self.bias)