1111from torch_concepts .nn .modules .mid .variable import ConceptVariable
1212from torch_concepts .nn .modules .mid .factors .cpd import ParametricCPD
1313from torch_concepts .nn .modules .mid .factors .potential import ParametricPotential
14- from torch_concepts .nn .modules .mid .inference .utils import enumerable_cardinality
14+ from torch_concepts .nn .modules .mid .inference .utils import (
15+ enumerable_cardinality ,
16+ factor_table ,
17+ )
1518from torch_concepts .nn .modules .mid .graph .probabilistic_model import ProbabilisticModel
1619from torch_concepts .nn .modules .mid .graph .markov_network import MarkovNetwork
1720from torch_concepts .nn .modules .mid .inference .torch .belief_propagation import BeliefPropagation
@@ -357,8 +360,7 @@ def test_cells_get_independent_masks(self):
357360 pot .parametrization ["energy" ][0 ].weight .zero_ ()
358361 fg = MarkovNetwork (variables = [a , b ], factors = [pot ])
359362 fg .train ()
360- eng = BeliefPropagation (fg , iters = 3 )
361- table = eng ._factor_table (
363+ table = factor_table (
362364 pot , ["a" , "b" ], {"a" : a , "b" : b }, {}, torch .Size ([1 ]),
363365 torch .float32 , torch .device ("cpu" ),
364366 )
@@ -545,7 +547,7 @@ def _member_probs(self, fg, out, owner_name, member):
545547 """The engine's marginal for one member, widened to a state distribution."""
546548 owner = fg .variables [owner_name ]
547549 return _state_marginal (
548- owner .member (member ), out .probs [owner_name ][..., owner . column_of ( member ) ]
550+ owner .member (member ), out .probs [member ]
549551 )
550552
551553 # -- test-time queries --------------------------------------------------
@@ -571,9 +573,9 @@ def test_partial_plate_evidence_matches_exact(self):
571573 assert torch .allclose (bp , exact [member ], atol = 1e-5 ), (member , bp , exact [member ])
572574 # An observed member reports its evidence: conditioning makes it a point mass.
573575 g = fg .variables ["g" ]
574- assert torch .allclose (out .probs ["g" ][..., g . column_of ( " g1") ], torch .ones (1 , 1 ))
576+ assert torch .allclose (out .probs ["g1" ], torch .ones (1 , 1 ))
575577 h = fg .variables ["h" ]
576- assert torch .allclose (out .probs ["h" ][..., h . column_of ( " h3") ], torch .zeros (1 , 1 ))
578+ assert torch .allclose (out .probs ["h3" ], torch .zeros (1 , 1 ))
577579
578580 def test_evidence_on_a_member_moves_the_other_plate (self ):
579581 """Both plates hang off ``root``, so conditioning one must move the other."""
@@ -651,7 +653,7 @@ def losses():
651653 out = eng .query (query = ["g" , "h" ], evidence = evidence )
652654 return sum (
653655 _binary_ce (
654- out .probs [owner ][..., fg . variables [ owner ]. column_of ( m ) ], target
656+ out .probs [m ], target
655657 )
656658 for owner , m , target in [
657659 ("g" , "g2" , targets ["g2" ]),
@@ -669,9 +671,9 @@ def losses():
669671 assert loss .item () < first * 0.5
670672
671673 out = eng .query (query = ["g" , "h" ], evidence = evidence )
672- assert out .probs ["g " ][0 , fg . variables [ "g" ]. column_of ( "g2" ) ].item () > 0.9
673- assert out .probs ["h " ][0 , fg . variables [ "h" ]. column_of ( "h2" ) ].item () > 0.9
674- assert out .probs ["h " ][0 , fg . variables [ "h" ]. column_of ( "h3" ) ].item () < 0.1
674+ assert out .probs ["g2 " ][0 ].item () > 0.9
675+ assert out .probs ["h2 " ][0 ].item () > 0.9
676+ assert out .probs ["h3 " ][0 ].item () < 0.1
675677
676678 # -- categorical plates (needs build_distribution's per-member split) ----
677679 def _mixed_plate_bn (self , seed = 32 ):
@@ -706,13 +708,13 @@ def test_categorical_plate_members_are_independent(self):
706708 out = BeliefPropagation (fg , iters = 30 ).query (query = ["h" ], evidence = {})
707709 h = fg .variables ["h" ]
708710 for member in h .members :
709- block = out .probs ["h" ][..., h . column_of ( member ) ]
711+ block = out .probs [member ]
710712 assert block .shape == (1 , 3 )
711713 assert torch .allclose (block .sum (- 1 ), torch .ones (1 ), atol = 1e-5 ), member
712714 exact = self ._exact_member_marginals (fg , {})
713715 for member in h .members :
714716 assert torch .allclose (
715- out .probs ["h" ][..., h . column_of ( member ) ], exact [member ], atol = 1e-5
717+ out .probs [member ], exact [member ], atol = 1e-5
716718 ), member
717719
718720 def test_categorical_plate_partial_evidence (self ):
@@ -723,14 +725,14 @@ def test_categorical_plate_partial_evidence(self):
723725 )
724726 exact = self ._exact_member_marginals (fg , {"h1" : 2 , "g1" : 1 })
725727 h , g = fg .variables ["h" ], fg .variables ["g" ]
726- assert torch .allclose (out .probs ["h" ][..., h . column_of ( " h2") ], exact ["h2" ], atol = 1e-5 )
728+ assert torch .allclose (out .probs ["h2" ], exact ["h2" ], atol = 1e-5 )
727729 assert torch .allclose (
728- _state_marginal (g .member ("g2" ), out .probs ["g" ][..., g . column_of ( " g2") ]),
730+ _state_marginal (g .member ("g2" ), out .probs ["g2" ]),
729731 exact ["g2" ], atol = 1e-5 ,
730732 )
731733 # observed members echo their evidence
732734 assert torch .allclose (
733- out .probs ["h" ][..., h . column_of ( " h1") ], torch .tensor ([[0.0 , 0.0 , 1.0 ]])
735+ out .probs ["h1" ], torch .tensor ([[0.0 , 0.0 , 1.0 ]])
734736 )
735737
736738 def test_categorical_plate_trains (self ):
@@ -745,7 +747,7 @@ def test_categorical_plate_trains(self):
745747 def loss_fn ():
746748 out = eng .query (query = ["h" ], evidence = evidence )
747749 return torch .nn .functional .nll_loss (
748- out .probs ["h" ][..., h . column_of ( " h2") ].log (), target
750+ out .probs ["h2" ].log (), target
749751 )
750752
751753 first = loss_fn ().item ()
@@ -756,4 +758,4 @@ def loss_fn():
756758 opt .step ()
757759 assert loss .item () < first * 0.5
758760 out = eng .query (query = ["h" ], evidence = evidence )
759- assert out .probs ["h " ][0 , h . column_of ( "h2" ) ][2 ].item () > 0.9
761+ assert out .probs ["h2 " ][0 ][2 ].item () > 0.9
0 commit comments