Your first node
This tutorial builds a factor node step by step, starting from an empty module: a normal distribution with a known variance,
\[f(y, x, v) = \mathcal{N}(y \mid x, v),\]
where y is the output, x the mean and v the variance. You define the distribution, declare it as a node, write its belief propagation and variational rules, give it an average energy, and check what it can compute. Every step runs the rules by hand, as a test would. An engine such as ReactiveMP runs the same rules in a graph.
Observations and constants reach a rule as a PointMass from BayesBase, and the distribution interface (mean, var, logpdf) comes from Distributions.
using MessagePassingRulesBase, BayesBase, DistributionsWhat a node is
A node is a Julia value that names a factor of a model: a type, or a function such as +. The node is what everything else refers to. A rule names it, node = Gaussian, and finding a rule is Julia's dispatch on its type, so any loaded package can add rules for a node another package declared (Defining rules explains how). This package builds no models: an engine such as ReactiveMP places the node in a graph, and a model written with RxInfer names it in a statement such as y ~ Gaussian(x, v).
Most stochastic nodes are probability distributions, and the rule packages use the distribution's own type as the node. StandardMessagePassingRules, for example, declares ExponentialFamily's NormalMeanVariance as a node. One type then does three jobs:
- it names the factor in a model,
y ~ NormalMeanVariance(μ, v)in RxInfer's syntax; - it is the factor's density: the declaration derives the node's log-density from it,
logpdf(NormalMeanVariance(μ, v), y), which rule fallbacks and rule tests use; - it is often a message: the message towards
outfrom point-mass inputs is the node's own density,NormalMeanVariance(μ, v), and many other rules return the same family.
Not every message has the node's type: the message towards a variance is another family. Nor is every node a distribution. A deterministic function is a node of its own, as A deterministic node with a group shows with +, and a factor that has no distribution type is named by an empty type, struct MyFactor end, used only as a name.
This tutorial follows the pattern with a small normal distribution of its own. A real package takes the type from ExponentialFamily, which also defines products and conversions of messages.
struct Gaussian{T <: Real} <: ContinuousUnivariateDistribution
μ::T # the mean
v::T # the variance
end
Gaussian(μ::Real, v::Real) = Gaussian(promote(μ, v)...)
Distributions.mean(d::Gaussian) = d.μ
Distributions.var(d::Gaussian) = d.v
Distributions.logpdf(d::Gaussian, x::Real) = -(log(2π * d.v) + abs2(x - d.μ) / d.v) / 2Declare the node
@define_factor_node declares the type as a node. The declaration names the node's interfaces, its edges in a graph, in order: the output first, then the distribution's parameters in the order its constructor takes them. The mean also answers to mean, an alias.
@define_factor_node(
node = Gaussian,
type = Stochastic,
interfaces = [:out, (:μ, aliases = [:mean]), :v],
)Each keyword has a section in the Keyword reference: node, type and interfaces, and the optional ones this node does not need. Stochastic says the node is a density over its interfaces, as opposed to a Deterministic function. The declaration is data, which nodespec returns and which draws itself:
MessagePassingRulesBase.nodespec(Gaussian)| interface | aliases | |
|---|---|---|
out | ||
μ | mean | |
v |
| default algorithm | DefaultAlgorithm() |
|---|---|
| static inputs | none |
| defined | first-node.md:80 |
Because the node is a distribution, callable with its parameters, the declaration also gives its log-density as a function of the interfaces, nodefunction:
f = MessagePassingRulesBase.nodefunction(Gaussian)
f(out = 1.0, μ = 0.0, v = 2.0) ≈ logpdf(Normal(0.0, sqrt(2.0)), 1.0)trueThe node has no rules yet, so it can compute no message. Defining nodes covers every part of a declaration.
A message towards the output
The message towards out says what the node knows about y from the messages on its other edges. For a normal message on the mean, $\mathcal{N}(x \mid m, s)$, and a known variance $v$, belief propagation integrates the mean out:
\[\mu_{f \to y}(y) = \int \mathcal{N}(y \mid x, v)\, \mathcal{N}(x \mid m, s)\, \mathrm{d}x = \mathcal{N}(y \mid m, s + v).\]
@define_message_update_rule defines the rule. Its target is out; its args are the message on μ, a Gaussian or a point mass, and the message on v, a point mass; its body computes the result from them. The result is the node's own type:
@define_message_update_rule(
node = Gaussian,
target = :out,
args = (m[:μ]::Union{Gaussian, PointMass}, m[:v]::PointMass),
logscale = 0,
body = (args) -> Gaussian(mean(args.m[:μ]), var(args.m[:μ]) + mean(args.m[:v])),
)m[:μ] reads as "the message on μ". The body receives the inputs as args, and args.m[:μ] is that message; a point mass has variance zero, so one body serves both. The logscale keyword states the logarithm of the message's normalising constant, which The log scale of a message explains.
@call_message_update_rule calls the rule with inputs of your choice, as an engine would:
@call_message_update_rule(
node = Gaussian, target = :out,
m = (μ = Gaussian(1.0, 2.0), v = PointMass(0.5)),
)Result
| value | Gaussian{Float64}(μ=1.0, v=2.5) |
|---|---|
| type | Gaussian{Float64} |
| log scale | 0 (declared) |
Inputs
| edge | value |
|---|
Rule
| declared inputs | m[:μ]::Union{BayesBase.PointMass, Gaussian}, m[:v]::BayesBase.PointMass |
|---|---|
| algorithm | DefaultAlgorithm() |
| log scale | 0 |
| defined | first-node.md:127 |
| body | args->Gaussian(mean(args.m[:μ]), var(args.m[:μ]) + mean(args.m[:v])) |
The result draws the node with the edges the rule read, messages as solid arrows in, and the target as the arrow out. Its value is Gaussian(1.0, 2.5): the variances add. With the mean observed too, the message is the node's density itself:
@call_message_update_rule(node = Gaussian, target = :out, m = (μ = PointMass(1.0), v = PointMass(0.5)))Result
| value | Gaussian{Float64}(μ=1.0, v=0.5) |
|---|---|
| type | Gaussian{Float64} |
| log scale | 0 (declared) |
Inputs
| edge | value |
|---|
Rule
| declared inputs | m[:μ]::Union{BayesBase.PointMass, Gaussian}, m[:v]::BayesBase.PointMass |
|---|---|
| algorithm | DefaultAlgorithm() |
| log scale | 0 |
| defined | first-node.md:127 |
| body | args->Gaussian(mean(args.m[:μ]), var(args.m[:μ]) + mean(args.m[:v])) |
A message towards the mean
The message towards μ is the same integral taken the other way. When y is observed, its message is a point mass at the observation, and the message towards μ is the likelihood of the observation as a function of the mean, $\mathcal{N}(y \mid x, v)$, a normal in $x$. The density is symmetric in y and x, so the rule mirrors the one towards out:
@define_message_update_rule(
node = Gaussian,
target = :μ,
args = (m[:out]::Union{Gaussian, PointMass}, m[:v]::PointMass),
logscale = 0,
body = (args) -> Gaussian(mean(args.m[:out]), var(args.m[:out]) + mean(args.m[:v])),
)
@call_message_update_rule(node = Gaussian, target = :μ, m = (out = PointMass(3.0), v = PointMass(0.5)))Result
| value | Gaussian{Float64}(μ=3.0, v=0.5) |
|---|---|
| type | Gaussian{Float64} |
| log scale | 0 (declared) |
Inputs
| edge | value |
|---|
Rule
| declared inputs | m[:out]::Union{BayesBase.PointMass, Gaussian}, m[:v]::BayesBase.PointMass |
|---|---|
| algorithm | DefaultAlgorithm() |
| log scale | 0 |
| defined | first-node.md:169 |
| body | args->Gaussian(mean(args.m[:out]), var(args.m[:out]) + mean(args.m[:v])) |
Which inputs a rule takes
The rules so far take messages, because every interface of the node is in one cluster. Which inputs a rule receives is not the rule's choice. It follows from the factorisation of the approximate posterior: a rule takes the messages on the other interfaces of its target's cluster, and the marginals of the other clusters.
| factorisation | clusters | the rule towards out takes | which is |
|---|---|---|---|
q(y, x, v) | (out, μ, v) | m[:μ], m[:v] | belief propagation |
q(y) q(x) q(v) | (out), (μ), (v) | q[:μ], q[:v] | mean-field variational message passing |
q(y, x) q(v) | (out, μ), (v) | m[:μ], q[:v] | structured variational message passing |
So a node supports a factorisation when it has rules for the inputs that factorisation delivers. Algorithms and dependencies describes this scheme in full.
A variational rule
Under the mean-field factorisation, the rule towards out takes the marginals q(x) and q(v), written q[:μ] and q[:v] among its args. Variational message passing sends the exponentiated expected log-density:
\[\mu_{f \to y}(y) \propto \exp \mathbb{E}_{q(x)}\big[\log \mathcal{N}(y \mid x, v)\big] \propto \mathcal{N}\big(y \mid \mathbb{E}[x], v\big).\]
Only the mean of q(x) matters, so the rule accepts any marginal with a mean:
@define_message_update_rule(
node = Gaussian,
target = :out,
args = (q[:μ]::Any, q[:v]::PointMass),
body = (args) -> Gaussian(mean(args.q[:μ]), mean(args.q[:v])),
)
@call_message_update_rule(
node = Gaussian, target = :out,
q = (μ = Gaussian(1.0, 2.0), v = PointMass(0.5)),
)Result
| value | Gaussian{Float64}(μ=1.0, v=0.5) |
|---|---|
| type | Gaussian{Float64} |
| log scale | undefined: the message rule for Gaussian towards :out under DefaultAlgorithm declares no `logscale` |
Inputs
| edge | value | |
|---|---|---|
| μ | q | Gaussian{Float64}(μ=1.0, v=2.0) |
| v | q | PointMass{Float64}(0.5) |
Rule
| declared inputs | q[:μ]::Any, q[:v]::BayesBase.PointMass |
|---|---|
| algorithm | DefaultAlgorithm() |
| log scale | none |
| defined | first-node.md:211 |
| body | args->Gaussian(mean(args.q[:μ]), mean(args.q[:v])) |
Other rules for this target (1)
first-node.md:127| ✓ | algorithm DefaultAlgorithm |
| ✗ | m[:μ]::Union{BayesBase.PointMass, Gaussian} not provided |
| ✗ | m[:v]::BayesBase.PointMass not provided |
| ✗ | q[:v]::BayesBase.PointMass{Float64} provided but not consumed |
| ✗ | q[:μ]::Gaussian{Float64} provided but not consumed |
The inputs are drawn dashed, as marginals, and the variance of q(x) no longer reaches the result. The rule declares no logscale, so the result's log scale is undefined: an expected log-density has no normalising constant with a meaning of its own.
The log scale of a message
A message is a distribution up to a constant, and the log scale is the logarithm of that constant. The belief propagation rules above declare logscale = 0 because their integrals are already normalised: a convolution of normals integrates to one. Summed along a graph, log scales give the model's evidence, so a rule that declares one must state it correctly. Log scales covers the other ways to declare one.
The average energy
The Bethe free energy, the quantity message passing minimises, needs each node's average energy, its expected negative log-density under the marginals of its clusters. Under the mean-field factorisation with a known variance:
\[U = \tfrac{1}{2}\log(2\pi v) + \frac{\operatorname{var}[y] + \operatorname{var}[x] + (\mathbb{E}[y] - \mathbb{E}[x])^2}{2v}.\]
@define_average_energy defines it, with the same args and body as a rule and no target, and @call_average_energy calls it:
@define_average_energy(
node = Gaussian,
args = (q[:out]::Any, q[:μ]::Any, q[:v]::PointMass),
body = (args) -> begin
y, x, v = args.q[:out], args.q[:μ], mean(args.q[:v])
(log(2v * π) + (var(y) + var(x) + (mean(y) - mean(x))^2) / v) / 2
end,
)
@call_average_energy(
node = Gaussian,
q = (out = Gaussian(0.0, 1.0), μ = Gaussian(1.0, 2.0), v = PointMass(0.5)),
)Result
| value | 4.57236 |
|---|---|
| type | Float64 |
Inputs
| edge | value | |
|---|---|---|
| out | q | Gaussian{Float64}(μ=0.0, v=1.0) |
| μ | q | Gaussian{Float64}(μ=1.0, v=2.0) |
| v | q | PointMass{Float64}(0.5) |
Rule
| declared inputs | q[:out]::Any, q[:μ]::Any, q[:v]::BayesBase.PointMass |
|---|---|
| algorithm | DefaultAlgorithm() |
| defined | first-node.md:251 |
| body | args->begin
(y, x, v) = (args.q[:out], args.q[:μ], mean(args.q[:v]))
(log((2v) * π) + (var(y) + var(x) + (mean(y) - mean(x)) ^ 2) / v) / 2
end |
What the node can compute
rule_coverage tabulates the node's rules: a row per target and one for the average energy, a column per algorithm.
MessagePassingRulesBase.rule_coverage(Gaussian)| DefaultAlgorithm | |
|---|---|
| → out | ✓×2 |
| → μ | ✓ |
| → v | |
| average energy | ✓ |
which_message_update_rule finds the rule a call would run, without running it, and shows what it consumes:
which_message_update_rule(Gaussian, :out; q = (μ = Gaussian(1.0, 2.0), v = PointMass(0.5)))| inputs | q[:μ]::Any, q[:v]::BayesBase.PointMass |
|---|---|
| in-place | no |
| scratch | no |
| pure | yes |
| services | none |
| log scale | none |
| defined | first-node.md:211 |
| body | args->Gaussian(mean(args.q[:μ]), mean(args.q[:v])) |
check_rules compares every rule with its node's declaration, and returns the problems it finds; a rule package's tests run it. Inspecting rules covers these queries.
MessagePassingRulesBase.check_rules(@__MODULE__)MessagePassingRulesBase.RuleIssue[]When no rule fits
The node has no rule towards v. Asking for one throws a RuleNotFoundError, which says what was asked and, for every rule of the node and target, why it does not fit:
try
@call_message_update_rule(node = Gaussian, target = :v, m = (out = PointMass(3.0), μ = PointMass(1.0)))
catch err
showerror(stdout, err)
endRuleNotFoundError: no message rule for Gaussian towards :v under DefaultAlgorithm() takes the inputs (m[:out]::BayesBase.PointMass{Float64}, m[:μ]::BayesBase.PointMass{Float64})
no rule exists for this node and target under any algorithm
what to try: no loaded package defines this rule: load the package that defines the node's rules, or define the rule with `@define_message_update_rule`The same error reaches a model's user when a graph needs a rule that no package defines. Calling rules describes every part of the report.
Next steps
- A deterministic node with a group writes rules for
out = in₁ + in₂ + …. - A node with its own algorithm gives a node a parametrised algorithm and declares what its rules take.
- Defining nodes, Defining rules and the Keyword reference cover every option.
- To use a node in a model, see ReactiveMP and RxInfer.