Skip to content

Commit d32fe02

Browse files
Merge pull request #76 from francescoTheSantis/mid_cleaning
Mid cleaning
2 parents 1db89c3 + d337c61 commit d32fe02

56 files changed

Lines changed: 2625 additions & 1921 deletions

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎doc/guides/using_mid_level.rst‎

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -183,6 +183,12 @@ Concept Bottleneck Model ``input → latent → concepts → task`` as a probabi
183183
point predictions.
184184
- :class:`~torch_concepts.nn.AncestralSamplingInference` — draws a (reparameterised) sample per
185185
variable in topological order.
186+
- :class:`~torch_concepts.nn.BeliefPropagation` — differentiable sum-product marginals over any
187+
factor graph (directed, undirected or mixed). Exact on trees, the standard approximation with
188+
loops. This is the engine to *train* an undirected model with.
189+
- :class:`~torch_concepts.nn.PgmpyVariableElimination` — **exact** marginals on any factor
190+
graph, via pgmpy. Evaluation only: it exports a static table per observation, so it carries
191+
no gradients.
186192
- :class:`~torch_concepts.nn.ForwardInference`, :class:`~torch_concepts.nn.IndependentInference`,
187193
:class:`~torch_concepts.nn.RejectionSampling`, :class:`~torch_concepts.nn.ImportanceSampling`,
188194
and the Pyro-backed :class:`~torch_concepts.nn.VariationalInference` provide further

‎doc/modules/mid_level_api.rst‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,6 @@ Probabilistic Models
5858
ProbabilisticModel
5959
BayesianNetwork
6060
MarkovNetwork
61-
ChainGraph
6261

6362
Inference
6463
---------
@@ -73,10 +72,10 @@ Inference
7372
AncestralSamplingInference
7473
MAPForwardInference
7574
BeliefPropagation
75+
PgmpyVariableElimination
7676
RejectionSampling
7777
ImportanceSampling
7878
VariationalInference
79-
PyroImportanceSampling
8079
BaseProposal
8180
MutilatedNetworkProposal
8281

@@ -90,3 +89,4 @@ Base Classes
9089
BaseInference
9190
TorchBaseInference
9291
PyroBaseInference
92+
PgmpyBaseInference

‎examples/experiments/plate_refactor_acceptance.py‎

Lines changed: 1 addition & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -177,18 +177,13 @@ def verify():
177177
# 9. (post-refactor) Pyro engines, skipped if pyro is absent
178178
try:
179179
import pyro # noqa: F401
180-
from torch_concepts.nn import VariationalInference, PyroImportanceSampling
180+
from torch_concepts.nn import VariationalInference
181181
vi = VariationalInference(pgm)
182182
seed_everything(0)
183183
y_hi = vi.query(query=["y"], evidence={"x": xt, "c1": ones}).probs["y"]
184184
seed_everything(0)
185185
y_lo = vi.query(query=["y"], evidence={"x": xt, "c1": zeros}).probs["y"]
186186
assert not torch.allclose(y_hi, y_lo), "member evidence ignored by model_fn"
187-
with warnings.catch_warnings():
188-
warnings.simplefilter("ignore")
189-
pis = PyroImportanceSampling(pgm, n_samples=200)
190-
p = pis.query({"y": ones4}, evidence={"c1": ones4}).probabilities
191-
assert p.shape == (B4,) and torch.isfinite(p).all()
192187
print("9. Pyro engines honor member evidence OK")
193188
except ImportError:
194189
print("9. pyro-ppl not installed SKIPPED")

‎examples/utilization/1_pgm/2_test_time_inference.py‎

Lines changed: 17 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -26,23 +26,23 @@
2626
from torch_concepts.data import BnLearnDataset
2727
from torch_concepts.nn import ParametricCPD, BayesianNetwork, \
2828
AncestralSamplingInference, LearnablePrior, RejectionSampling, \
29-
ImportanceSampling, MutilatedNetworkProposal, PyroImportanceSampling, \
30-
BeliefPropagation
29+
ImportanceSampling, MutilatedNetworkProposal, BeliefPropagation, \
30+
PgmpyVariableElimination
3131

3232

3333
def _compare_query(
3434
query_name, evidence_names, targets,
35-
rejection_engine, torch_is_engine, pyro_is_engine, bp_engine,
35+
rejection_engine, is_engine, bp_engine, ve_engine,
3636
):
3737
"""Estimate P(query_name=1 | evidence) with every engine and the data.
3838
3939
Enumerates every {0,1} assignment of ``evidence_names`` (each a single
40-
binary node), queries the four inference engines plus the empirical
40+
binary node), queries every inference engine plus the empirical
4141
frequency in ``targets``, and prints one row per estimator.
4242
"""
4343
combos = list(itertools.product([0., 1.], repeat=len(evidence_names)))
44-
rows = {"empirical": [], "reject S": [], "torch IS": [], "pyro IS": [],
45-
"belief prop": []}
44+
rows = {"empirical": [], "reject S": [], "torch IS": [],
45+
"belief prop": [], "exact VE": []}
4646

4747
for combo in combos:
4848
evidence = {n: torch.tensor([[v]]) for n, v in zip(evidence_names, combo)}
@@ -51,12 +51,13 @@ def _compare_query(
5151
rows["reject S"].append(
5252
rejection_engine.query(query, evidence).probabilities.item())
5353
rows["torch IS"].append(
54-
torch_is_engine.query(query, evidence).probabilities.item())
55-
rows["pyro IS"].append(
56-
pyro_is_engine.query(query, evidence).probabilities.item())
54+
is_engine.query(query, evidence).probabilities.item())
5755
rows["belief prop"].append(
5856
bp_engine.query(query=[query_name], evidence=evidence)
5957
.probs[query_name].item())
58+
rows["exact VE"].append(
59+
ve_engine.query(query=[query_name], evidence=evidence)
60+
.probs[query_name].item())
6061

6162
mask = torch.ones(targets[query_name].shape[0], dtype=torch.bool)
6263
for n, v in zip(evidence_names, combo):
@@ -71,8 +72,8 @@ def _compare_query(
7172
fmt = lambda vals: " ".join(f"{v:6.3f}" for v in vals)
7273
print(f"\n=== P({query_name}=1 | {', '.join(evidence_names)}) ===")
7374
print(f"{'':>14} | {header}")
74-
for label in ("empirical", "reject S", "torch IS", "pyro IS", "belief prop"):
75-
print(f"{label:>14} | {fmt(rows[label])}")
75+
for label, vals in rows.items():
76+
print(f"{label:>14} | {fmt(vals)}")
7677

7778

7879
def main():
@@ -229,12 +230,14 @@ def main():
229230
# a purely forward computation.
230231
n_test_samples = 20_000
231232
rejection_engine = RejectionSampling(concept_model, n_samples=n_test_samples)
232-
torch_is_engine = ImportanceSampling(
233+
is_engine = ImportanceSampling(
233234
concept_model, MutilatedNetworkProposal(concept_model),
234235
n_samples=n_test_samples, initial_temperature=0.1,
235236
)
236-
pyro_is_engine = PyroImportanceSampling(concept_model, n_samples=n_test_samples)
237237
bp_engine = BeliefPropagation(concept_model, iters=20, damping=0.2)
238+
# Exact marginals: the reference the approximate engines are judged
239+
# against. Evaluation only -- it carries no gradients.
240+
ve_engine = PgmpyVariableElimination(concept_model)
238241

239242
for query_name, evidence_names in [
240243
("tub", ("either", "lung")),
@@ -243,7 +246,7 @@ def main():
243246
]:
244247
_compare_query(
245248
query_name, evidence_names, targets,
246-
rejection_engine, torch_is_engine, pyro_is_engine, bp_engine,
249+
rejection_engine, is_engine, bp_engine, ve_engine,
247250
)
248251

249252
return

‎tests/distributions/test_delta.py‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -202,9 +202,9 @@ def test_dtype_preservation(self):
202202
self.assertEqual(dist_float64.mean.dtype, torch.float64)
203203

204204
def test_batch_shape(self):
205-
"""Test batch_shape attribute."""
205+
"""batch_shape follows the value, so Delta composes with Independent."""
206206
dist = Delta([1.0, 2.0])
207-
self.assertEqual(dist.batch_shape, torch.Size([]))
207+
self.assertEqual(dist.batch_shape, torch.Size([2]))
208208

209209
def test_multiple_samples_consistency(self):
210210
"""Test that multiple samples are consistent."""

‎tests/nn/modules/high/models/test_cbvae.py‎

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -90,9 +90,9 @@ def test_categorical_concepts_normalise_per_concept(self, categorical_annotation
9090
def test_a_categorical_plates_cpd_emits_one_block_per_member(self):
9191
"""The plate's concept CPD emits one score block per member.
9292
93-
Asserted on the CPD's own output: the emitted width is the plate's
94-
flattened width, laid out member-major so the distribution can normalise
95-
each member's block independently.
93+
Asserted on the CPD's own output, which is in member layout
94+
``(batch, n_members, states)``, so the distribution can normalise each
95+
member's block independently.
9696
"""
9797
# Same cardinality on both concepts, so they can share one plate.
9898
annotations = Annotations(
@@ -110,8 +110,8 @@ def test_a_categorical_plates_cpd_emits_one_block_per_member(self):
110110
model(query=list(model.pgm.variables), input=torch.rand(6, INPUT_SIZE))
111111

112112
logits = emitted["logits"]
113-
assert logits.shape == (6, 6) # 2 members x 3 states, member-major
114-
probs = logits.reshape(6, 2, 3).softmax(-1)
113+
assert logits.shape == (6, 2, 3) # 2 members x 3 states
114+
probs = logits.softmax(-1)
115115
assert torch.allclose(probs.sum(-1), torch.ones(6, 2), atol=1e-5)
116116

117117
def test_the_latent_prior_and_guide_produce_a_positive_scale(self, binary_annotations):

‎tests/nn/modules/high/models/test_cvae.py‎

Lines changed: 8 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -92,8 +92,9 @@ def test_binary_concepts_report_logits(self, binary_annotations):
9292
def test_categorical_marginals_normalise_per_concept(self, categorical_annotations):
9393
model = build_model(categorical_annotations, plate=False)
9494
for name, cardinality in (("digit", 4), ("color", 3)):
95+
# A CPD reports the member layout: one member of `cardinality` states.
9596
logits = model.pgm.factors[name](parent_values={})["logits"]
96-
assert logits.shape == (cardinality,)
97+
assert logits.shape == (1, cardinality)
9798
assert torch.allclose(logits.softmax(-1).sum(), torch.ones(()))
9899

99100
def test_a_categorical_plates_marginal_normalises_each_member(self):
@@ -111,7 +112,8 @@ def test_a_categorical_plates_marginal_normalises_each_member(self):
111112
model = build_model(annotations, plate=True)
112113
assert "concepts" in model.pgm.variables # one plate, both members
113114

114-
assert model.pgm.factors["concepts"](parent_values={})["logits"].shape == (8,)
115+
# Member layout on the CPD (2 members x 4 states); flat on the engine output.
116+
assert model.pgm.factors["concepts"](parent_values={})["logits"].shape == (2, 4)
115117
drawn = AncestralSamplingInference(model.pgm).query(
116118
query=["concepts"], evidence={}, n_samples=5
117119
).samples["concepts"]
@@ -534,7 +536,8 @@ def test_a_multivariate_normal_concept_sizes_its_cholesky_factor(self):
534536
plate=False,
535537
variable_distributions={"continuous": MultivariateNormal},
536538
)
539+
# Member layout: one member, whose `scale_tril` carries the extra rank.
537540
params = model.pgm.factors["v"](parent_values={})
538-
assert params["loc"].shape == (3,)
539-
assert params["scale_tril"].shape == (3, 3)
540-
assert bool((params["scale_tril"].diagonal() > 0).all())
541+
assert params["loc"].shape == (1, 3)
542+
assert params["scale_tril"].shape == (1, 3, 3)
543+
assert bool((params["scale_tril"].diagonal(dim1=-2, dim2=-1) > 0).all())

‎tests/nn/modules/mid/inference/test_assembly_caching.py‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -168,7 +168,9 @@ def test_samples_assembly_is_cached_consistently():
168168
"""_assemble_samples shares the cached chunks with _assemble_params but
169169
keys its annotation separately, so both must stay correct."""
170170
eng = _engine()
171-
per_variable = {"concepts": torch.rand(3, 3), "y": torch.randn(3, 2)}
171+
# Realisations arrive in member layout (B, n_members, member_size), as the
172+
# engines cache them.
173+
per_variable = {"concepts": torch.rand(3, 3, 1), "y": torch.randn(3, 1, 2)}
172174
for _ in range(2):
173175
samples = eng._assemble_samples(per_variable, ["concepts", "y"])
174176
assert list(samples.annotation.labels) == ["c1", "c2", "c3", "y"]

‎tests/nn/modules/mid/inference/test_belief_propagation.py‎

Lines changed: 19 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,10 @@
1111
from torch_concepts.nn.modules.mid.variable import ConceptVariable
1212
from torch_concepts.nn.modules.mid.factors.cpd import ParametricCPD
1313
from 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+
)
1518
from torch_concepts.nn.modules.mid.graph.probabilistic_model import ProbabilisticModel
1619
from torch_concepts.nn.modules.mid.graph.markov_network import MarkovNetwork
1720
from 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

Comments
 (0)