55
66"""Radial and hierarchical consistency losses for p-adic VAE."""
77
8- from typing import Any , Dict , Tuple
8+ from typing import Any , Dict , Optional , Tuple
99
1010import torch
1111import torch .nn as nn
@@ -67,8 +67,10 @@ def forward(
6767 ) -> Tuple [torch .Tensor , MetricsDict ]:
6868 device = z_hyp .device
6969 batch_size = z_hyp .size (0 )
70+ from typing import cast
7071 cur_c = kwargs .get ("curvature" , self .curvature )
71- valuations = self ._valuation_fn (batch_indices ).double ()
72+ v_raw = self ._valuation_fn (batch_indices )
73+ valuations = cast (torch .Tensor , v_raw ).double ()
7274 actual_radius = hyperbolic_radius (z_hyp , c = cur_c )
7375
7476 target_radii_all = _euclidean_to_hyperbolic_radius (
@@ -182,7 +184,8 @@ def forward(
182184 if batch_size < 2 :
183185 return torch .tensor (0.0 , device = device , dtype = torch .float64 ), {"n_levels" : 0 }
184186
185- valuations = self ._valuation_fn (batch_indices )
187+ from typing import cast
188+ valuations = cast (torch .Tensor , self ._valuation_fn (batch_indices ))
186189 radii = hyperbolic_radius (z_hyp , c = cur_c )
187190
188191 target_radii_all = _euclidean_to_hyperbolic_radius (
@@ -191,7 +194,7 @@ def forward(
191194 )
192195
193196 dim_size = self .max_valuation + 1
194- vals_long = valuations .long ()
197+ vals_long = cast ( torch . LongTensor , valuations .long (). to ( device ) )
195198 present_mask = level_has_data (vals_long , dim_size = dim_size )
196199 levels_present = present_mask .nonzero (as_tuple = False ).squeeze (- 1 ).tolist ()
197200
@@ -226,20 +229,19 @@ def forward(
226229 target_loss = F .mse_loss (level_means , target_guidance )
227230 total_loss = loss + self .target_loss_weight * target_loss
228231
229- per_gap_tensors : Dict [ str , torch . Tensor ] = {}
232+ per_gap_tensors : MetricsDict = {}
230233 for i in range (n_levels - 1 ):
231- per_gap_tensors [f"gap_viol_tensor_v { levels_present [i ]} " ] = F . relu ( violations [i ])
234+ per_gap_tensors [f"gap_viol_v { levels_present [i ]} " ] = float ( max ( 0.0 , violations [i ]. item ()) )
232235
233- metrics = {
236+ metrics : MetricsDict = {
234237 "n_levels" : n_levels ,
235- "margin_violations" : ( violations > 0 ).sum ().item (),
236- "monotonic_loss" : loss .item (),
237- "target_loss" : target_loss .item (),
238+ "margin_violations" : int (( violations > 0 ).sum ().item () ),
239+ "monotonic_loss" : float ( loss .item () ),
240+ "target_loss" : float ( target_loss .item () ),
238241 }
239242 for i in range (n_levels ):
240- metrics [f"r_v{ levels_present [i ]} " ] = level_means [i ].item ()
241- for i in range (n_levels - 1 ):
242- metrics [f"gap_viol_v{ levels_present [i ]} " ] = float (max (0.0 , violations [i ].item ()))
243+ metrics [f"r_v{ levels_present [i ]} " ] = float (level_means [i ].item ())
244+
243245 metrics .update (per_gap_tensors )
244246
245247 return total_loss , metrics
@@ -273,20 +275,22 @@ def __init__(
273275 )
274276 self .register_buffer ("target_radii" , target_radii )
275277
276- def forward (
278+ def forward ( # type: ignore[override]
277279 self ,
278280 z_hyp : torch .Tensor ,
279281 batch_indices : torch .Tensor ,
280282 ** kwargs : Any ,
281283 ) -> Tuple [Dict [str , torch .Tensor ], MetricsDict ]:
282- logits , targets = kwargs .get ("logits" ), kwargs .get ("targets" )
284+ logits : Optional [torch .Tensor ] = kwargs .get ("logits" )
285+ targets : Optional [torch .Tensor ] = kwargs .get ("targets" )
283286 device , cur_c = z_hyp .device , kwargs .get ("curvature" , self .curvature )
284287 radii = hyperbolic_radius (z_hyp , c = cur_c )
285288 target_radii_adj = _euclidean_to_hyperbolic_radius (
286289 _exponential_target_radii (9 , self .inner_radius , self .outer_radius , scale = 3.0 ).to (device ),
287290 c = cur_c
288291 )
289- valuations = self ._valuation_fn (batch_indices ).long ().to (device )
292+ from typing import cast
293+ valuations = cast (torch .LongTensor , self ._valuation_fn (batch_indices ).long ().to (device ))
290294
291295 dim_size = 10
292296 present_mask = level_has_data (valuations , dim_size = dim_size )
@@ -327,13 +331,13 @@ def forward(
327331 margin = torch .maximum (min_m , target_radii_adj [v ] - target_radii_adj [v_next ])
328332 separation_loss = separation_loss + F .relu (means_all [v_next ] - means_all [v ] + margin )
329333
330- metrics = {"hierarchy" : hierarchy_loss .item (), "coverage" : coverage_loss .item (),
331- "separation" : separation_loss .item (), "variance" : variance_loss .item ()}
334+ metrics : MetricsDict = {"hierarchy" : float ( hierarchy_loss .item ()) , "coverage" : float ( coverage_loss .item () ),
335+ "separation" : float ( separation_loss .item ()) , "variance" : float ( variance_loss .item () )}
332336 with torch .no_grad ():
333337 for v in range (dim_size ):
334338 if present_mask [v ]:
335- metrics [f"r_mean_v{ v } " ] = means_all [v ].item ()
336- metrics [f"r_std_v{ v } " ] = stds_all [v ].item ()
339+ metrics [f"r_mean_v{ v } " ] = float ( means_all [v ].item () )
340+ metrics [f"r_std_v{ v } " ] = float ( stds_all [v ].item () )
337341
338342 return {"hierarchy" : total_hier_loss , "coverage" : coverage_loss , "separation" : separation_loss }, metrics
339343
0 commit comments