Rule fallbacks

Where no rule fits, an engine may consult a rule fallback instead of reporting the error. A rule fallback is a callable, fallback(node, target, args), that takes the node, the target and the rule's arguments. It returns a message, or nothing when it has none either.

An engine consults the fallback only when resolution returns a RuleNotFound. A fallback therefore never replaces a rule that exists, and an error inside a rule is never turned into a fallback. ReactiveMP takes a fallback as its activation option rulefallback. A message that a fallback computes has an undefined log scale.

The node function fallback

NodeFunctionRuleFallback is the fallback this package provides. It applies to a stochastic node without groups. Its message towards a target is the node's log-density as a function of the target, with every other input collapsed to a point, by default its mean. The message is a NodeFunctionLogPdf, an unnormalised log-density. A form constraint, or a product with a proper distribution, turns it into a distribution.

The node below is a normal distribution given by a function of its mean and variance. It has no rules at all:

using MessagePassingRulesBase, BayesBase, ExponentialFamily

gaussian(μ, v) = NormalMeanVariance(μ, v)   # the distribution of out, given μ and v

@define_factor_node(node = gaussian, type = Stochastic, interfaces = [:out, :μ, :v])

fallback = NodeFunctionRuleFallback()
args = MessagePassingRulesBase.RuleArgs(m = (out = PointMass(2.0),), q = (v = PointMass(0.5),))
message = fallback(gaussian, MessagePassingRulesBase.Target(:μ), args)

logpdf(message, 1.0), logpdf(NormalMeanVariance(1.0, 0.5), 2.0)
(-1.5723649429247002, -1.5723649429247002)

The fallback's message towards μ evaluates the node's log-density, $\log \mathcal{N}(2 \mid \mu, 0.5)$, at $\mu = 1$. The observation out = 2 and the variance v = 0.5 were collapsed to their means. An input that is not a point mass is collapsed the same way:

args = MessagePassingRulesBase.RuleArgs(m = (out = NormalMeanVariance(2.0, 3.0),), q = (v = PointMass(0.5),))
logpdf(fallback(gaussian, MessagePassingRulesBase.Target(:μ), args), 1.0)
-1.5723649429247002

The message on out has the mean 2.0, so the result is the same. Its variance is lost.

Where the fallback has nothing

The fallback returns nothing for a deterministic node, a node with groups, a member of a group, or an input that is a joint marginal. A deterministic node has no log-density:

julia> using MessagePassingRulesBase

julia> struct Plain end

julia> @define_factor_node(node = Plain, type = Deterministic, interfaces = [:out, :in])

julia> NodeFunctionRuleFallback()(Plain, MessagePassingRulesBase.Target(:out), MessagePassingRulesBase.RuleArgs(m = (in = 1.0,))) === nothing
true
MessagePassingRulesBase.NodeFunctionRuleFallback — Type
NodeFunctionRuleFallback(extract = mean)
(fallback::NodeFunctionRuleFallback)(node, target, args::RuleArgs)

A rule fallback: the message a stochastic node sends when no rule matches, computed from its nodefunction, the log-density its declaration defines. The message towards out or another interface is the node's log-density in that interface, every other input, a message or a marginal, collapsed to a point by extract. It is a NodeFunctionLogPdf, unnormalised, which a form constraint or a product with a proper distribution turns into one. Its log scale is undefined.

Arguments

  • extract: the function collapsing each input to a point. Default: mean.

Returns

Called as fallback(node, target, args): a NodeFunctionLogPdf, or nothing where the node function does not apply: a deterministic node, a node with groups, a member of a group, an input that is a joint, or an interface other than the target that has no input.

An engine consults a fallback only when resolution finds no rule, so an error inside a rule is never turned into a fallback; ReactiveMP takes it as the activation option rulefallback. A fallback of one's own is any callable of (node, target, args) returning a message or nothing.

Examples

fallback = NodeFunctionRuleFallback()
message = fallback(NormalMeanVariance, MessagePassingRulesBase.Target(:out),
                   MessagePassingRulesBase.RuleArgs(m = (μ = PointMass(0.0), v = PointMass(1.0))))
logpdf(message, 0.5)   # logpdf(NormalMeanVariance(0.0, 1.0), 0.5)
source