Skip to content

Approximate inference in aGrUM (pyAgrum)

Creative Commons LicenseaGrUMinteractive online version

There are several approximate inference for BN in aGrUM (pyAgrum). They share the same API than exact inference.

  • Loopy Belief Propagation : LBP is an approximate inference that uses exact calculous methods (when the BN os a tree) even if the BN is not a tree. LBP is a special case of inference : the algorithm may not converge and even if it converges, it may converge to anything (but the exact posterior). LBP however is fast and usually gives not so bad results.
  • Sampling inference : Sampling inference use sampling to compute the posterior. The sampling may be (very) slow but those algorithms converge to the exac distribution. aGrUM implements :
    • Montecarlo sampling,
    • Weighted sampling,
    • Importance sampling,
    • Gibbs sampling.
  • Finally, aGrUM propose the so-called ‘loopy version’ of the sampling algorithms : the idea is to use LBP as a Dirichlet prior for the sampling algorithm. A loopy version of each sampling algorithm is proposed.
%matplotlib inline
from pylab import *
import matplotlib.pyplot as plt
def unsharpen(bn):
"""
Force the parameters of the BN not to be a bit more far from 0 or 1
"""
for nod in bn.nodes():
bn.cpt(nod).translate(bn.maxParam() / 10).normalizeAsCPT()
def compareInference(ie, ie2, ax=None):
"""
compare 2 inference by plotting all the points from (posterior(ie),posterior(ie2))
"""
exact = []
appro = []
errmax = 0
for node in bn.nodes():
# Tensors as list
exact += ie.posterior(node).tolist()
appro += ie2.posterior(node).tolist()
errmax = max(errmax, (ie.posterior(node) - ie2.posterior(node)).abs().max())
if errmax < 1e-10:
errmax = 0
if ax == None:
fig = plt.Figure(figsize=(4, 4))
ax = plt.gca() # default axis for plt
ax.plot(exact, appro, "ro")
ax.set_title(
"{} vs {}\n {}\nMax error {:2.4} in {:2.4} seconds".format(
str(type(ie)).split(".")[2].split("_")[0][0:-2], # name of first inference
str(type(ie2)).split(".")[2].split("_")[0][0:-2], # name of second inference
ie2.messageApproximationScheme(),
errmax,
ie2.currentTime(),
)
)
import pyagrum as gum
import pyagrum.lib.notebook as gnb
bn = gum.loadBN("res/alarm.bgum")
unsharpen(bn)
ie = gum.LazyPropagation(bn)
ie.makeInference()
gnb.showBN(bn, size="8")

svg

gnb.sideBySide(gnb.getJunctionTreeMap(bn), gnb.getInference(bn, size="8")) # using LazyPropagation by default
print(ie.posterior("KINKEDTUBE"))
0 0~16 0--0~16 1 1~32 1--1~32 2 2~33 2--2~33 3 3~4 3--3~4 4 4~22 4--4~22 5 5~22 5--5~22 6 6~23 6--6~23 7 7~26 7--7~26 8 8~17 8--8~17 10 10~14 10--10~14 11 11~16 11--11~16 12 12~13 12--12~13 13 13~30 13--13~30 14 14~26 14--14~26 16 16~17 16--16~17 17 17~24 17--17~24 19 19~27 19--19~27 20 20~33 20--20~33 22 22~33 22--22~33 23 23~27 23--23~27 23~31 23--23~31 24 24~26 24--24~26 26 26~27 26--26~27 27 30 30~31 30--30~31 31 31~32 31--31~32 32 32~33 32--32~33 33 19~27--27 12~13--13 2~33--33 23~27--27 22~33--33 11~16--16 24~26--26 31~32--32 10~14--14 26~27--27 13~30--30 5~22--22 7~26--26 20~33--33 16~17--17 32~33--33 23~31--31 8~17--17 1~32--32 3~4--4 4~22--22 17~24--24 14~26--26 30~31--31 6~23--23 0~16--16
structs Inference in   1.21ms KINKEDTUBE 2026-09-28T17:47:29.289543 image/svg+xml Matplotlib v3.11.2, VENTLUNG 2026-09-28T17:47:30.139971 image/svg+xml Matplotlib v3.11.2, KINKEDTUBE->VENTLUNG PRESS 2026-09-28T17:47:30.226475 image/svg+xml Matplotlib v3.11.2, KINKEDTUBE->PRESS HYPOVOLEMIA 2026-09-28T17:47:29.334564 image/svg+xml Matplotlib v3.11.2, STROKEVOLUME 2026-09-28T17:47:29.757254 image/svg+xml Matplotlib v3.11.2, HYPOVOLEMIA->STROKEVOLUME LVEDVOLUME 2026-09-28T17:47:29.858947 image/svg+xml Matplotlib v3.11.2, HYPOVOLEMIA->LVEDVOLUME INTUBATION 2026-09-28T17:47:29.365748 image/svg+xml Matplotlib v3.11.2, SHUNT 2026-09-28T17:47:29.968286 image/svg+xml Matplotlib v3.11.2, INTUBATION->SHUNT INTUBATION->VENTLUNG MINVOL 2026-09-28T17:47:30.196898 image/svg+xml Matplotlib v3.11.2, INTUBATION->MINVOL INTUBATION->PRESS VENTALV 2026-09-28T17:47:30.265380 image/svg+xml Matplotlib v3.11.2, INTUBATION->VENTALV MINVOLSET 2026-09-28T17:47:29.398224 image/svg+xml Matplotlib v3.11.2, VENTMACH 2026-09-28T17:47:29.899849 image/svg+xml Matplotlib v3.11.2, MINVOLSET->VENTMACH PULMEMBOLUS 2026-09-28T17:47:29.456177 image/svg+xml Matplotlib v3.11.2, PAP 2026-09-28T17:47:29.716692 image/svg+xml Matplotlib v3.11.2, PULMEMBOLUS->PAP PULMEMBOLUS->SHUNT INSUFFANESTH 2026-09-28T17:47:29.486570 image/svg+xml Matplotlib v3.11.2, CATECHOL 2026-09-28T17:47:30.513909 image/svg+xml Matplotlib v3.11.2, INSUFFANESTH->CATECHOL ERRLOWOUTPUT 2026-09-28T17:47:29.518175 image/svg+xml Matplotlib v3.11.2, HRBP 2026-09-28T17:47:30.580499 image/svg+xml Matplotlib v3.11.2, ERRLOWOUTPUT->HRBP ERRCAUTER 2026-09-28T17:47:29.546049 image/svg+xml Matplotlib v3.11.2, HRSAT 2026-09-28T17:47:30.618230 image/svg+xml Matplotlib v3.11.2, ERRCAUTER->HRSAT HREKG 2026-09-28T17:47:30.688481 image/svg+xml Matplotlib v3.11.2, ERRCAUTER->HREKG FIO2 2026-09-28T17:47:29.569750 image/svg+xml Matplotlib v3.11.2, PVSAT 2026-09-28T17:47:30.331174 image/svg+xml Matplotlib v3.11.2, FIO2->PVSAT LVFAILURE 2026-09-28T17:47:29.626588 image/svg+xml Matplotlib v3.11.2, LVFAILURE->STROKEVOLUME LVFAILURE->LVEDVOLUME HISTORY 2026-09-28T17:47:29.998404 image/svg+xml Matplotlib v3.11.2, LVFAILURE->HISTORY DISCONNECT 2026-09-28T17:47:29.652642 image/svg+xml Matplotlib v3.11.2, VENTTUBE 2026-09-28T17:47:30.046111 image/svg+xml Matplotlib v3.11.2, DISCONNECT->VENTTUBE ANAPHYLAXIS 2026-09-28T17:47:29.684541 image/svg+xml Matplotlib v3.11.2, TPR 2026-09-28T17:47:29.830261 image/svg+xml Matplotlib v3.11.2, ANAPHYLAXIS->TPR CO 2026-09-28T17:47:30.653696 image/svg+xml Matplotlib v3.11.2, STROKEVOLUME->CO TPR->CATECHOL BP 2026-09-28T17:47:30.712319 image/svg+xml Matplotlib v3.11.2, TPR->BP PCWP 2026-09-28T17:47:29.941611 image/svg+xml Matplotlib v3.11.2, LVEDVOLUME->PCWP CVP 2026-09-28T17:47:30.093268 image/svg+xml Matplotlib v3.11.2, LVEDVOLUME->CVP VENTMACH->VENTTUBE SAO2 2026-09-28T17:47:30.399630 image/svg+xml Matplotlib v3.11.2, SHUNT->SAO2 VENTTUBE->VENTLUNG VENTTUBE->PRESS VENTLUNG->MINVOL VENTLUNG->VENTALV EXPCO2 2026-09-28T17:47:30.478728 image/svg+xml Matplotlib v3.11.2, VENTLUNG->EXPCO2 ARTCO2 2026-09-28T17:47:30.294646 image/svg+xml Matplotlib v3.11.2, VENTALV->ARTCO2 VENTALV->PVSAT ARTCO2->EXPCO2 ARTCO2->CATECHOL PVSAT->SAO2 SAO2->CATECHOL HR 2026-09-28T17:47:30.548223 image/svg+xml Matplotlib v3.11.2, CATECHOL->HR HR->HRBP HR->HRSAT HR->CO HR->HREKG CO->BP
KINKEDTUBE │
TRUE │FALSE │
─────────│─────────│
0.1167 │ 0.8833 │

Gibbs inference iterations can be stopped :

  • by the value of error (epsilon)
  • by the rate of change of epsilon (MinEpsilonRate)
  • by the number of iteration (MaxIteration)
  • by the duration of the algorithm (MaxTime)
ie2 = gum.GibbsSampling(bn)
ie2.setEpsilon(1e-2)
gnb.showInference(bn, engine=ie2, size="8")
print(ie2.posterior("KINKEDTUBE"))
print(ie2.messageApproximationScheme())
compareInference(ie, ie2)

svg

KINKEDTUBE │
TRUE │FALSE │
─────────│─────────│
0.1068 │ 0.8932 │
stopped with rate=0.006737946999085467

svg

With default parameters, this inference has been stopped by a low value of rate.

ie2 = gum.GibbsSampling(bn)
ie2.setMaxIter(1000)
ie2.setEpsilon(5e-3)
ie2.makeInference()
print(ie2.posterior(2))
print(ie2.messageApproximationScheme())
INTUBATION │
NORMAL │ESOPHAGEA│ONESIDED │
─────────│─────────│─────────│
0.6870 │ 0.1910 │ 0.1220 │
stopped with max iteration=1000
compareInference(ie, ie2)

svg

ie2 = gum.GibbsSampling(bn)
ie2.setMaxTime(3)
ie2.makeInference()
print(ie2.posterior(2))
print(ie2.messageApproximationScheme())
compareInference(ie, ie2)
INTUBATION │
NORMAL │ESOPHAGEA│ONESIDED │
─────────│─────────│─────────│
0.8000 │ 0.1967 │ 0.0033 │
stopped with epsilon=0.20189651799465538

svg

ie2 = gum.GibbsSampling(bn)
ie2.setEpsilon(10**-1.8)
ie2.setBurnIn(300)
ie2.setPeriodSize(300)
ie2.setDrawnAtRandom(True)
gnb.animApproximationScheme(ie2)
ie2.makeInference()

svg

svg

compareInference(ie, ie2)

svg

ie4 = gum.ImportanceSampling(bn)
ie4.setEpsilon(10**-1.8)
ie4.setMaxTime(10) # 10 seconds for inference
ie4.setPeriodSize(300)
ie4.makeInference()
compareInference(ie, ie4)

svg

Every sampling inference has a ‘hybrid’ version which consists in using a first loopy belief inference as a prior for the probability estimations by sampling.

ie3 = gum.LoopyGibbsSampling(bn)
ie3.setEpsilon(10**-1.8)
ie3.setMaxTime(10) # 10 seconds for inference
ie3.setPeriodSize(300)
ie3.makeInference()
compareInference(ie, ie3)

svg

These computations may be a bit long

def compareAllInference(bn, evs={}, epsilon=10**-1.6, epsilonRate=1e-8, maxTime=20):
ies = [
gum.LazyPropagation(bn),
gum.LoopyBeliefPropagation(bn),
gum.GibbsSampling(bn),
gum.LoopyGibbsSampling(bn),
gum.WeightedSampling(bn),
gum.LoopyWeightedSampling(bn),
gum.ImportanceSampling(bn),
gum.LoopyImportanceSampling(bn),
]
# burn in for Gibbs samplings
for i in [2, 3]:
ies[i].setBurnIn(300)
ies[i].setDrawnAtRandom(True)
for i in range(2, len(ies)):
ies[i].setEpsilon(epsilon)
ies[i].setMinEpsilonRate(epsilonRate)
ies[i].setPeriodSize(300)
ies[i].setMaxTime(maxTime)
for i in range(len(ies)):
ies[i].setEvidence(evs)
ies[i].makeInference()
fig, axes = plt.subplots(1, len(ies) - 1, figsize=(35, 3), num="gpplot")
for i in range(len(ies) - 1):
compareInference(ies[0], ies[i + 1], axes[i])
compareAllInference(bn, epsilon=1e-1)

svg

compareAllInference(bn, epsilon=1e-2)

svg

compareAllInference(bn, maxTime=1, epsilon=1e-8)

svg

compareAllInference(bn, maxTime=2, epsilon=1e-8)

svg

funny = {"BP": 1, "PCWP": 2, "EXPCO2": 0, "HISTORY": 0}
compareAllInference(bn, maxTime=1, evs=funny, epsilon=1e-8)

svg

compareAllInference(bn, maxTime=4, evs=funny, epsilon=1e-8)

svg

compareAllInference(bn, maxTime=10, evs=funny, epsilon=1e-8)

svg