Defining rules
A node computes three kinds of thing, and you define each with its own macro:
- a message towards one of its interfaces, with
@define_message_update_rule; - the joint marginal of a cluster of its interfaces, with
@define_marginal_update_rule; - its average energy, its term of the Bethe free energy, with
@define_average_energy.
Each definition names its node and its target. An average energy has no target. The definition also lists the inputs the rule takes, with their types, and gives its body, an ordinary lambda.
The macro turns the definition into a method of find_message_rule, find_marginal_rule or find_average_energy. Julia's dispatch then finds the rule by the node, the target, the algorithm and the types of the inputs, in any loaded package. A rule that omits algorithm belongs to its node's default algorithm, so you declare the node before the rule is loaded.
Your first node writes the rules of one node step by step. The Keyword reference lists every keyword of the three macros.
MessagePassingRulesBase.@define_message_update_rule — Macro
@define_message_update_rule(
node = ..., target = ..., args = (...), body = (...) -> ...,
algorithm = ..., logscale = ..., reads_logscale = ..., ctx = (...),
inplace = ..., preallocate = ..., scratch = ..., pure = ..., args_check = ...,
)Define the rule for the message a node sends towards one of its interfaces. The macro takes keyword arguments only; node, target, args and body are required, the rest optional. An unknown or repeated keyword is an error at definition time, naming the valid ones.
The rule becomes a method of find_message_rule: it is found for its node, target, algorithm and the types of the inputs args names, from any module, and its module's registry lists it for introspection (list_rules, check_rules).
Required keywords
node: the node, as declared with@define_factor_node: a type,NormalMeanVariance, or a function,+.target: the interface the message goes to::out, a single interface;(:m, k), any member of the groupm. The namekis bound to the member's index, anInt, inbodyand in thepreallocate,scratch,logscaleandargs_checkfunctions, without being listed among their parameters;argscan select by it.
args: the inputs the rule consumes, a tuple of entriescontainer[key]::T, the containermfor a message orqfor a marginal. A type left out isAny; the types are what the rule dispatches on. For a single interface:m[:μ]::T,q[:μ]::T: the message or the marginal onμ;q[:y, :x]::T, orq[(:y, :x)]::T: the joint marginal of a structural cluster, its members in interface order. A group in a cluster means all its members jointly:q[(:in,)]is the joint over the groupin.(:in, 1), with a literal index, is one member:q[:out, (:in, 1)]is the joint ofoutandin's first member.
For a group, whose value is a tuple in member order:
m[:in...]::T: every member, each of typeT;m[:in][k]::T: the target's own member,kbeing the name an indexed target binds;m[:in][!k]::T: every member but the target's own.
With a selection, the tuple keeps every position and holds
nothingwhere the selection leaves a member out, soargs.m[:in][k]is memberkwhatever was selected.defaultamong the entries stands for whatever inputs the default scheme delivers under the graph's factorisation. The rule then takes them all, requires the typed entries besidedefault, and walks them withrule_inputs: one rule for every factorisation.body: the rule itself, an ordinary lambda returning the message, whose parameters are some of the slots(output, scratch, algo, ctx, args, ann).Name only the slots the rule uses, in the order given; a slot that is misspelled, repeated or out of order is an error.
output: the buffer to write into, for an in-place rule only, and then first;scratch: the rule's working memory, when it declaresscratch;algo: the algorithm value the rule runs under, and so its parameters;ctx: theRuleContext, whose services the rule reads asctx.name;args: the inputs, read asargs.m[:μ],args.q[:y, :x],args.m[:in][k];ann: the annotations: those that arrived with the inputs,ann.m[:μ], and the rule's own, written withannotate!(ann, key, value).
Optional keywords
algorithm: the algorithm the rule runs under, a type, or a value whose type is used. Default: the node's own,default_algorithm(node), usuallyDefaultAlgorithm; the node must then be declared before the rule is loaded. Almost every rule leaves it out; naming one is for a rule switcher, aDefaultAlgorithmExtension, or a node's own algorithm. With a parametric algorithm typeT,algorithm = Tmatches everyT{…}, while leaving it out binds the rule to the type of the node's default instance only.logscale: the message's log scale, the scalar withmessage = exp(logscale) · resultfor the normalisedresultthe rule returns: a rule's result may stand for an unnormalised function, as a belief-propagation message does, and this is the log of its normaliser. One of:- a number:
logscale = 0, orlogscale = loghalf(StatsFuns'); anIrrationalkeeps the message's float type, while-logtwowould not, being aFloat64; - a function of the inputs, over the slots
(algo, ctx, args), in that order:logscale = (args) -> -log(abs(mean(args.m[:A]))); from_body: the body returnswith_logscale(result, logscale), for a log scale computed alongside the result;improper: the message has no normalising constant, so no log scale exists, as for an exact message whose integral is infinite.
Default: none declared. The message's log scale is then an
UndefinedLogScalenaming the rule, which propagates through products; onlyrequire_logscaleturns it into an error. Declaringimpropergives an undefined log scale too, whose reason says that none exists, where an omitted declaration says that it is not known.- a number:
reads_logscale:trueif the rule reads the log scales of its inbound messages, asargs.logscale.m[:x]. Its caller must then provide them: an engine does when it tracks log scales, and a call by hand takes them aslogscale = (...); without them the call is an error (check_reads_logscale). Default:false.ctx: the context services the rule reads, a tuple of symbols,ctx = (:rng,)orctx = (:node, :matrix_correction). Any name is allowed, so a rule may need a service of its own. An engine checks that its context supplies each one when it resolves the rule (check_services); a call by hand does not. Default:(), none.inplace:truefor a rule that writes its result into a buffer it is given rather than allocating one. It needspreallocate, and itsbodytakesoutputfirst. Default:false.preallocate: for an in-place rule, a function building the buffer from the inputs, over the slots(algo, ctx, args):preallocate = (args) -> similar(mean(args.m[:μ])). Allowed only withinplace = true.scratch: working memory, a function building it from the inputs over the slots(algo, ctx, args), given to the body as itsscratchslot:scratch = (args) -> (work = similar(mean(args.m[:μ])),). An engine keeps one per outbound stream and reuses it, so it is write-before-read: it carries nothing between calls, and the engine may keep, drop or rebuild it whenever it likes. It never leaves the rule, is never shared with another rule, keeps the rule pure, and combines withinplace. The body takesscratchexactly when this is given. An engine keeping it between calls infers its type from the inputs' types (rule_scratch_type), so a function whose result type infers, built from the inputs withsimilarorzeros(eltype(...), ...), runs the rule on a concretely typed scratch; one that does not infer runs it on an untyped one, a dynamic call per call. Default: none.pure:falsefor a rule with side effects,truefor a pure rule under an impure algorithm. Default: the algorithm's declaration,ispure,truefor almost every algorithm. A pure rule mutates neither its inputs nor state shared beyond one call, and draws randomness only fromctx.rng.args_check: a check of the inputs, run as the body starts, for what their types cannot say: a value in range, a vector's length, a variate form anAnyinput must have. A function over the slots(algo, ctx, args), in that order, returningtruewhen the inputs pass;falsewhen they do not: the rule raises aRuleInputErrorquoting the check's source,args_check = (args) -> length(mean(args.q[:out])) == 2;- a string when they do not, and the error says that instead: for a check that combines several conditions, whose source would read poorly, or that names the offending value,
args_check = (args) -> 0 <= mean(args.m[:out]) <= 1 || lazy"a probability in [0, 1]; got $(mean(args.m[:out]))".
A string is always a failure. A failed check is an error, not a reason to select another rule: resolution has already chosen this one.
The check costs nothing where it depends on the inputs' types only, since it folds away when the rule is compiled for them, and one comparison where it reads a value; its error path is outlined, so it allocates nothing until it fails. It runs after a
preallocateorscratchhelper, which see the same inputs. Performance trap: build the message only where the check fails. After||it is:cond || "...$x"builds nothing on a passing call. A message bound before the condition, or formatted by a helper before it decides, is built and allocated on every call. Alazy"..."string is the safe habit, since it defers the formatting until the error is shown. Default: none.
Examples
@define_message_update_rule(
node = NormalMeanVariance,
target = :out,
args = (m[:μ]::PointMass, m[:v]::PointMass),
logscale = 0,
body = (args) -> NormalMeanVariance(mean(args.m[:μ]), mean(args.m[:v])),
)
# Towards any member of the group `m`, reading the aligned member of the group `p`;
# `k` is bound in the body without being listed:
@define_message_update_rule(
node = NormalMixture,
target = (:m, k),
args = (q[:out]::Any, q[:switch]::Categorical, q[:p][k]::Any),
body = (args) -> … probvec(args.q[:switch])[k] …,
)MessagePassingRulesBase.@define_marginal_update_rule — Macro
@define_marginal_update_rule(
node = ..., target = (:y, :x), args = (...), body = (...) -> ...,
algorithm = ..., ctx = (...), inplace = ..., preallocate = ..., scratch = ..., pure = ...,
args_check = ...,
)Define the rule for the joint marginal of a structural cluster: the marginal of, say, (:y, :x) a node computes from the messages on the cluster's members and the marginals of its other interfaces. The macro takes keyword arguments only; node, target, args and body are required, the rest optional. An unknown or repeated keyword is an error at definition time, naming the valid ones.
The rule becomes a method of find_marginal_rule, found for its node, cluster, algorithm and input types, from any module. A marginal carries no log scale, so logscale and reads_logscale are not accepted.
Required keywords
node: the node, as declared with@define_factor_node: a type,NormalMeanVariance, or a function,+.target: the cluster, its members in interface order:(:y, :x), a cluster of interfaces;(:out, (:T, 1)), with a member of a group written with a literal index;- a bare name,
target = members: any cluster of the node, the name bound to the cluster's key inbodyand in thepreallocateandscratchfunctions, without being listed among their parameters. Used withdefaultinargs, for one rule over every factorisation.
args: the inputs the rule consumes, a tuple of entriescontainer[key]::T, the containermfor a message orqfor a marginal. A type left out isAny; the types are what the rule dispatches on. For a single interface:m[:μ]::T,q[:μ]::T: the message or the marginal onμ;q[:y, :x]::T, orq[(:y, :x)]::T: the joint marginal of a structural cluster, its members in interface order. A group in a cluster means all its members jointly:q[(:in,)]is the joint over the groupin.(:in, 1), with a literal index, is one member:q[:out, (:in, 1)]is the joint ofoutandin's first member.
For a group, whose value is a tuple in member order:
m[:in...]::T: every member, each of typeT;m[:in][k]::T: the target's own member,kbeing the name an indexed target binds;m[:in][!k]::T: every member but the target's own.
With a selection, the tuple keeps every position and holds
nothingwhere the selection leaves a member out, soargs.m[:in][k]is memberkwhatever was selected.defaultamong the entries stands for whatever inputs the default scheme delivers under the graph's factorisation. The rule then takes them all, requires the typed entries besidedefault, and walks them withrule_inputs: one rule for every factorisation.A marginal rule typically reads the messages on the cluster's members,
m[:y]andm[:x], and the marginals of the node's other interfaces.body: the rule itself, an ordinary lambda returning the joint marginal, whose parameters are some of the slots(output, scratch, algo, ctx, args, ann).Name only the slots the rule uses, in the order given; a slot that is misspelled, repeated or out of order is an error.
output: the buffer to write into, for an in-place rule only, and then first;scratch: the rule's working memory, when it declaresscratch;algo: the algorithm value the rule runs under, and so its parameters;ctx: theRuleContext, whose services the rule reads asctx.name;args: the inputs, read asargs.m[:μ],args.q[:y, :x],args.m[:in][k];ann: the annotations: those that arrived with the inputs,ann.m[:μ], and the rule's own, written withannotate!(ann, key, value).
Optional keywords
algorithm: the algorithm the rule runs under, a type, or a value whose type is used. Default: the node's own,default_algorithm(node), usuallyDefaultAlgorithm; the node must then be declared before the rule is loaded. Almost every rule leaves it out; naming one is for a rule switcher, aDefaultAlgorithmExtension, or a node's own algorithm. With a parametric algorithm typeT,algorithm = Tmatches everyT{…}, while leaving it out binds the rule to the type of the node's default instance only.ctx: the context services the rule reads, a tuple of symbols,ctx = (:rng,)orctx = (:node, :matrix_correction). Any name is allowed, so a rule may need a service of its own. An engine checks that its context supplies each one when it resolves the rule (check_services); a call by hand does not. Default:(), none.inplace:truefor a rule that writes its result into a buffer it is given rather than allocating one. It needspreallocate, and itsbodytakesoutputfirst. Default:false.preallocate: for an in-place rule, a function building the buffer from the inputs, over the slots(algo, ctx, args):preallocate = (args) -> similar(mean(args.m[:μ])). Allowed only withinplace = true.scratch: working memory, a function building it from the inputs over the slots(algo, ctx, args), given to the body as itsscratchslot:scratch = (args) -> (work = similar(mean(args.m[:μ])),). An engine keeps one per outbound stream and reuses it, so it is write-before-read: it carries nothing between calls, and the engine may keep, drop or rebuild it whenever it likes. It never leaves the rule, is never shared with another rule, keeps the rule pure, and combines withinplace. The body takesscratchexactly when this is given. An engine keeping it between calls infers its type from the inputs' types (rule_scratch_type), so a function whose result type infers, built from the inputs withsimilarorzeros(eltype(...), ...), runs the rule on a concretely typed scratch; one that does not infer runs it on an untyped one, a dynamic call per call. Default: none.pure:falsefor a rule with side effects,truefor a pure rule under an impure algorithm. Default: the algorithm's declaration,ispure,truefor almost every algorithm. A pure rule mutates neither its inputs nor state shared beyond one call, and draws randomness only fromctx.rng.args_check: a check of the inputs, run as the body starts, for what their types cannot say: a value in range, a vector's length, a variate form anAnyinput must have. A function over the slots(algo, ctx, args), in that order, returningtruewhen the inputs pass;falsewhen they do not: the rule raises aRuleInputErrorquoting the check's source,args_check = (args) -> length(mean(args.q[:out])) == 2;- a string when they do not, and the error says that instead: for a check that combines several conditions, whose source would read poorly, or that names the offending value,
args_check = (args) -> 0 <= mean(args.m[:out]) <= 1 || lazy"a probability in [0, 1]; got $(mean(args.m[:out]))".
A string is always a failure. A failed check is an error, not a reason to select another rule: resolution has already chosen this one.
The check costs nothing where it depends on the inputs' types only, since it folds away when the rule is compiled for them, and one comparison where it reads a value; its error path is outlined, so it allocates nothing until it fails. It runs after a
preallocateorscratchhelper, which see the same inputs. Performance trap: build the message only where the check fails. After||it is:cond || "...$x"builds nothing on a passing call. A message bound before the condition, or formatted by a helper before it decides, is built and allocated on every call. Alazy"..."string is the safe habit, since it defers the formatting until the error is shown. Default: none.
Example
@define_marginal_update_rule(
node = NormalMeanVariance,
target = (:out, :μ),
args = (m[:out]::NormalMeanVariance, m[:μ]::NormalMeanVariance, q[:v]::PointMass),
body = (args) -> …,
)MessagePassingRulesBase.@define_average_energy — Macro
@define_average_energy(node = ..., args = (...), body = (...) -> ..., algorithm = ..., ctx = (...), pure = ..., args_check = ...)Define a node's average energy, E_q[-log f] under the marginals of its clusters: the node's term of the Bethe free energy before the clusters' entropies are subtracted. The macro takes keyword arguments only; node, args and body are required, the rest optional. An unknown or repeated keyword is an error at definition time, naming the valid ones.
The energy becomes a method of find_average_energy, found for its node, algorithm and input types, from any module. It has no target, returns a number and has no log scale, so target, inplace, preallocate, scratch, logscale and reads_logscale are not accepted.
Required keywords
node: the node, as declared with@define_factor_node: a type,NormalMeanVariance, or a function,+.args: the marginals the energy reads, one per cluster of the node's factorisation, a tuple of entriesq[key]::T. A type left out isAny; the types are what the energy dispatches on.q[:μ]::T: the marginal of a single interface, a cluster of its own;q[:y, :x]::T, orq[(:y, :x)]::T: the joint marginal of a structural cluster, its members in interface order;q[(:in,)]is the joint over the groupin, and(:in, 1)one member;q[:in...]::T: every member of a group, each a cluster of its own, as a tuple in member order.
defaultamong the entries stands for whatever clusters the graph's factorisation delivers; the energy then takes them all, requires the typed entries besidedefault, and walks them withrule_inputs: one energy for every factorisation.body: the energy, an ordinary lambda returning a real number, whose parameters are some of the slots(algo, ctx, args, ann), named in that order:algo: the algorithm value, and so its parameters;ctx: theRuleContext, whose services the energy reads asctx.name;args: the marginals, read asargs.q[:μ],args.q[:y, :x];ann: the annotations that arrived with the marginals,ann.q[:μ].
Optional keywords
algorithm: the algorithm the rule runs under, a type, or a value whose type is used. Default: the node's own,default_algorithm(node), usuallyDefaultAlgorithm; the node must then be declared before the rule is loaded. Almost every rule leaves it out; naming one is for a rule switcher, aDefaultAlgorithmExtension, or a node's own algorithm. With a parametric algorithm typeT,algorithm = Tmatches everyT{…}, while leaving it out binds the rule to the type of the node's default instance only.ctx: the context services the rule reads, a tuple of symbols,ctx = (:rng,)orctx = (:node, :matrix_correction). Any name is allowed, so a rule may need a service of its own. An engine checks that its context supplies each one when it resolves the rule (check_services); a call by hand does not. Default:(), none.pure:falsefor a rule with side effects,truefor a pure rule under an impure algorithm. Default: the algorithm's declaration,ispure,truefor almost every algorithm. A pure rule mutates neither its inputs nor state shared beyond one call, and draws randomness only fromctx.rng.args_check: a check of the inputs, run as the body starts, for what their types cannot say: a value in range, a vector's length, a variate form anAnyinput must have. A function over the slots(algo, ctx, args), in that order, returningtruewhen the inputs pass;falsewhen they do not: the rule raises aRuleInputErrorquoting the check's source,args_check = (args) -> length(mean(args.q[:out])) == 2;- a string when they do not, and the error says that instead: for a check that combines several conditions, whose source would read poorly, or that names the offending value,
args_check = (args) -> 0 <= mean(args.m[:out]) <= 1 || lazy"a probability in [0, 1]; got $(mean(args.m[:out]))".
A string is always a failure. A failed check is an error, not a reason to select another rule: resolution has already chosen this one.
The check costs nothing where it depends on the inputs' types only, since it folds away when the rule is compiled for them, and one comparison where it reads a value; its error path is outlined, so it allocates nothing until it fails. It runs after a
preallocateorscratchhelper, which see the same inputs. Performance trap: build the message only where the check fails. After||it is:cond || "...$x"builds nothing on a passing call. A message bound before the condition, or formatted by a helper before it decides, is built and allocated on every call. Alazy"..."string is the safe habit, since it defers the formatting until the error is shown. Default: none.
Example
@define_average_energy(
node = NormalMeanVariance,
args = (q[:out]::Any, q[:μ]::Any, q[:v]::Any),
body = (args) -> begin
m_out, v_out = mean_var(args.q[:out])
m_μ, v_μ = mean_var(args.q[:μ])
(log(2π) + mean(log, args.q[:v]) + mean(inv, args.q[:v]) * (v_out + v_μ + abs2(m_out - m_μ))) / 2
end,
)Targets
A message rule's target is a single interface, target = :out, or any member of a group, target = (:in, k). The second form binds k to the member's index in the body, and the inputs can select members by it.
A marginal rule's target is a cluster, target = (:out, :μ), with its members in interface order. A member of a group is written (:T, 1) inside a cluster.
Internally a target is a type, so rules dispatch on it:
MessagePassingRulesBase.Target — Type
Target(edge::Symbol)
Target{E}()The target of a message rule towards the single interface E, written target = :out in a rule's declaration. The edge is a type parameter, so a rule dispatches on it. An engine and the message_passing_* calls pass one to find_message_rule; a rule's body never sees it.
julia> using MessagePassingRulesBase: Target, target_edge
julia> target_edge(Target(:out))
:outSee also IndexedTarget, ClusterTarget.
MessagePassingRulesBase.IndexedTarget — Type
IndexedTarget(edge::Symbol, index::Integer)
IndexedTarget{E}(index)The target of a message rule towards member index of the interface group E, written target = (:m, k) in a rule's declaration. The index is a field rather than a type parameter, so every member of a group shares one rule, which reads its member's index as the name k it binds.
julia> using MessagePassingRulesBase: IndexedTarget, target_edge, target_index
julia> t = IndexedTarget(:m, 2); (target_edge(t), target_index(t))
(:m, 2)See also Target, ClusterTarget.
MessagePassingRulesBase.ClusterTarget — Type
ClusterTarget(members::Tuple)
ClusterTarget{K}()The target of a marginal rule: the structural cluster K, written target = (:y, :x) in a rule's declaration, its members in interface-declaration order. A member of a group is written (:T, 1), as in (:out, (:T, 1)), the joint of out and the group T's first member. A rule declared with a bare name, target = members, is defined for every ClusterTarget of its node.
julia> using MessagePassingRulesBase: ClusterTarget, cluster_members
julia> cluster_members(ClusterTarget((:out, (:T, 1))))
(:out, (:T, 1))See also Target, IndexedTarget.
MessagePassingRulesBase.target_edge — Function
target_edge(target::Union{Target, IndexedTarget}) -> SymbolThe interface a message rule's target points at: the interface's name for a Target, the group's name for an IndexedTarget.
MessagePassingRulesBase.target_index — Function
target_index(target::IndexedTarget) -> IntThe index of the group member an IndexedTarget points at, counted from 1.
MessagePassingRulesBase.cluster_members — Function
cluster_members(target::ClusterTarget) -> TupleThe members of a cluster, in interface-declaration order: an interface's name, or (group, k) for member k of a group.
Inputs
args lists what a rule consumes:
m[:x], the message onx;q[:x], the marginal ofx;q[:y, :x], the joint marginal of the cluster(y, x);m[:in...], every member of the groupin;m[:in][k], for a group target, the member aligned with the target;m[:in][!k], for a group target, every other member.
A group arrives as a tuple in member order. The tuple holds nothing where the selection leaves a member out, so args.m[:in][k] is member k whatever was selected:
julia> using MessagePassingRulesBase
julia> struct Sum end # out = in₁ + in₂ + …
julia> @define_factor_node(node = Sum, type = Deterministic, interfaces = [:out, :in...])
julia> @define_message_update_rule(
node = Sum, target = :out, args = (m[:in...]::Real,),
body = (args) -> sum(args.m[:in]),
)
julia> @define_message_update_rule(
node = Sum, target = (:in, k), args = (m[:out]::Real, m[:in][!k]::Real),
body = (args) -> args.m[:out] - sum(x for x in args.m[:in] if x !== nothing),
)
julia> getresult(@call_message_update_rule(node = Sum, target = :out, m = (in = (1.0, 2.0, 3.0),)))
6.0
julia> getresult(@call_message_update_rule(node = Sum, target = (:in, 2), m = (out = 6.0, in = (1.0, nothing, 3.0))))
2.0The body receives the inputs as a RuleArgs: args.m is a Messages, and args.q is a Marginals.
The rule does not choose its inputs. The engine chooses them, and under the default algorithm they follow from the factorisation, as Algorithms and dependencies explains. You define a rule for the inputs it will be given. A deterministic node with a group writes group rules step by step.
MessagePassingRulesBase.RuleArgs — Type
RuleArgs(; m = NamedTuple(), q = NamedTuple(), logscale = nothing)
RuleArgs(m::Messages, q::Marginals[, logscale])The inputs a rule's body receives as args, and the value rules dispatch on.
args.m: the inbound messages, aMessages;args.q: the marginals, aMarginals;args.logscale: the log scales that arrived with the messages, aRuleLogScalesread asargs.logscale.m[:out], when the caller tracks them, andnothingotherwise.
Rules dispatch on the keys and types of m and q alone; a rule reads args.logscale only when declared with reads_logscale = true.
Keywords
m: the messages, aNamedTupleor aMessages. Default: none.q: the marginals of single interfaces, aNamedTupleor aMarginals. Joints need theMarginalsconstructor. Default: none.logscale: the messages' log scales, aNamedTuplekeyed likem, ornothingwhen they are not tracked. Default:nothing.
Examples
julia> using MessagePassingRulesBase: RuleArgs
julia> args = RuleArgs(m = (μ = 1.0,), q = (v = 2.0,));
julia> args.m[:μ], args.q[:v], args.logscale
(1.0, 2.0, nothing)MessagePassingRulesBase.Messages — Type
Messages(values::NamedTuple)The inbound messages a rule receives as args.m, keyed by interface name: args.m[:μ]. A group is a tuple of its members in order under the group's name, args.m[:inputs][k], with nothing for a member the rule does not take. The keys are held sorted, so the order a caller gives them in does not matter. keys(m) gives the names; indexing by a name m does not hold throws a KeyError.
julia> using MessagePassingRulesBase: Messages
julia> m = Messages((μ = 1.0, inputs = (2.0, 3.0)));
julia> m[:μ], m[:inputs][2], keys(m)
(1.0, 3.0, (:inputs, :μ))MessagePassingRulesBase.Marginals — Type
Marginals(singles::NamedTuple = NamedTuple())
Marginals(singles::NamedTuple, Val(cluster_keys), joints::Tuple)The marginals a rule receives as args.q. Single interfaces and groups are reached like messages, args.q[:name] and args.q[:p][k]. A structural cluster is reached by the tuple of its members, in interface-declaration order: args.q[(:y, :x)], or args.q[:y, :x] for short. Inside a cluster a group's name stands for all its members jointly, so args.q[(:in,)] is the joint over the group in, while args.q[:in] is the tuple of its members' marginals. One member of a group is (:in, 1): args.q[:out, (:in, 1)] is the joint of out and that member.
The second form takes the joints' keys, as a Val of a tuple of member tuples, and the joints in the same order. A cluster's key is carried in the type and never turned into a symbol, so no name is derived and none can collide: q[:y_x] below is an interface named y_x, not the cluster.
Throws
ArgumentErrorwhen the numbers of cluster keys and joints differ;KeyErrorwhen indexed by a name or a cluster it does not hold.
Examples
julia> using MessagePassingRulesBase: Marginals
julia> q = Marginals((τ = 2.0, y_x = 1.0), Val(((:y, :x),)), (0.5,));
julia> q[:τ], q[:y, :x], q[(:y, :x)], q[:y_x]
(2.0, 0.5, 0.5, 1.0)MessagePassingRulesBase.canonical_cluster_keys — Function
canonical_cluster_keys(keys) -> TupleSort joint-cluster keys into the order a Marginals holds them: by the names of their members, a group member's index after its group's name. A rule's dispatch signature and the containers an engine builds both use this order, so neither needs the other's spelling.
julia> using MessagePassingRulesBase: canonical_cluster_keys
julia> canonical_cluster_keys(((:y, :x), (:out, (:T, 2)), (:out, (:T, 1))))
((:out, (:T, 1)), (:out, (:T, 2)), (:y, :x))Rules over whatever the factorisation delivers
A tensor node, among others, computes the same thing under any factorisation. Such a node's rule declares default among its args. The rule then takes whatever inputs the default scheme delivers, and it requires the typed inputs named beside default.
The node below averages the means of its inputs x, whether they arrive as messages or as marginals. Its rule requires the marginal of v to be a point mass:
using MessagePassingRulesBase, BayesBase, ExponentialFamily
using MessagePassingRulesBase: rule_inputs
struct Average end # out ~ N(mean of x₁, x₂, …, v)
@define_factor_node(node = Average, type = Stochastic, interfaces = [:out, :v, :x...])
@define_message_update_rule(
node = Average, target = :out, args = (default, q[:v]::PointMass),
body = (args) -> begin
inputs = (rule_inputs(Average, args.m)..., rule_inputs(Average, args.q)...)
means = [mean(value) for (key, value) in inputs if key isa Tuple]
NormalMeanVariance(sum(means) / length(means), mean(args.q[:v]))
end,
)
@call_message_update_rule(
node = Average, target = :out,
m = (x = (NormalMeanVariance(1.0, 1.0), NormalMeanVariance(3.0, 1.0)),),
q = (v = PointMass(2.0),),
)Result
| value | ExponentialFamily.NormalMeanVariance{Float64}(μ=2.0, v=2.0) |
|---|---|
| type | ExponentialFamily.NormalMeanVariance{Float64} |
| log scale | undefined: the message rule for Average towards :out under DefaultAlgorithm declares no `logscale` |
Inputs
| edge | value | |
|---|---|---|
| v | q | PointMass{Float64}(2.0) |
Rule
| declared inputs | default, q[:v]::BayesBase.PointMass |
|---|---|
| algorithm | DefaultAlgorithm() |
| log scale | none |
| defined | rules.md:125 |
| body | args->begin
inputs = (rule_inputs(Average, args.m)..., rule_inputs(Average, args.q)...)
means = [mean(value) for (key, value) = inputs if key isa Tuple]
NormalMeanVariance(sum(means) / length(means), mean(args.q[:v]))
end |
The same rule runs when the x arrive as marginals:
@call_message_update_rule(
node = Average, target = :out,
q = (x = (NormalMeanVariance(1.0, 1.0), NormalMeanVariance(3.0, 1.0)), v = PointMass(2.0)),
)Result
| value | ExponentialFamily.NormalMeanVariance{Float64}(μ=2.0, v=2.0) |
|---|---|
| type | ExponentialFamily.NormalMeanVariance{Float64} |
| log scale | undefined: the message rule for Average towards :out under DefaultAlgorithm declares no `logscale` |
Inputs
| edge | value | |
|---|---|---|
| v | q | PointMass{Float64}(2.0) |
| x[1] | q | ExponentialFamily.NormalMeanVariance{Float64}(μ=1.0, v=1.0) |
| x[2] | q | ExponentialFamily.NormalMeanVariance{Float64}(μ=3.0, v=1.0) |
Rule
| declared inputs | default, q[:v]::BayesBase.PointMass |
|---|---|
| algorithm | DefaultAlgorithm() |
| log scale | none |
| defined | rules.md:125 |
| body | args->begin
inputs = (rule_inputs(Average, args.m)..., rule_inputs(Average, args.q)...)
means = [mean(value) for (key, value) = inputs if key isa Tuple]
NormalMeanVariance(sum(means) / length(means), mean(args.q[:v]))
end |
The body walks its inputs with rule_inputs, which returns key => value pairs. The key is an interface's name, (:x, k) for a member of a group, or a joint's key.
A marginal rule over any cluster names its target with a bare name, target = members. The body receives the cluster's key under that name.
A rule with explicit inputs for the same node and target is more specific, and it wins where it applies. You can therefore place a fast path on top of a default rule. When a typed input is missing or has another type, the lookup returns a RuleNotFound. A node has at most one default rule per target and algorithm.
MessagePassingRulesBase.rule_inputs — Function
rule_inputs(node, container::Union{Messages, Marginals}) -> Tuple{Vararg{Pair}}The inputs in a rule's args.m or args.q as a tuple of key => value pairs: an interface by its name, a member of a group as (group, k), and a joint by its key, as (:out, (:T, 1)). Members a group does not deliver, nothing in its tuple, are left out. node tells which names are groups. It is for a rule declared with default, whose body walks the inputs the factorisation delivered; the pairs have known types and constant keys, so the walk is type-stable.
julia> using MessagePassingRulesBase: Messages, rule_inputs
julia> struct Tensor end
julia> @define_factor_node(node = Tensor, type = Stochastic, interfaces = [:out, :T...])
julia> rule_inputs(Tensor, Messages((out = 1.0, T = (2.0, nothing, 3.0))))
((:T, 1) => 2.0, (:T, 3) => 3.0, :out => 1.0)The body
The body is a lambda over some of the slots (output, scratch, algo, ctx, args, ann). You name only the slots the body uses, in that order:
output: the buffer an in-place rule writes into;scratch: the rule's working memory;algo: the algorithm value the rule runs under, which carries its parameters (see Algorithms and dependencies);ctx: theRuleContext, which holds the services the rule declares withctx = (...)(see The rule context);args: the inputs;ann: the annotations.
A rule may compute its result with another rule, of its own node or of another, a packaged one included. It calls that rule with the inputs it builds and forwards its ctx, so the other rule sees the same services. The node below shifts the message of the Average node above:
struct ShiftedAverage end # out ~ N(1 + mean of x₁, x₂, …, v)
@define_factor_node(node = ShiftedAverage, type = Stochastic, interfaces = [:out, :v, :x...])
@define_message_update_rule(
node = ShiftedAverage, target = :out, args = (m[:x...]::NormalMeanVariance, q[:v]::PointMass),
body = (ctx, args) -> begin
average = getresult(call_message_update_rule(Average, :out; m = (x = args.m[:x],), q = (v = args.q[:v],), ctx))
NormalMeanVariance(mean(average) + 1.0, var(average))
end,
)
@call_message_update_rule(
node = ShiftedAverage, target = :out,
m = (x = (NormalMeanVariance(1.0, 1.0), NormalMeanVariance(3.0, 1.0)),),
q = (v = PointMass(2.0),),
)Result
| value | ExponentialFamily.NormalMeanVariance{Float64}(μ=3.0, v=2.0) |
|---|---|
| type | ExponentialFamily.NormalMeanVariance{Float64} |
| log scale | undefined: the message rule for ShiftedAverage towards :out under DefaultAlgorithm declares no `logscale` |
Inputs
| edge | value | |
|---|---|---|
| v | q | PointMass{Float64}(2.0) |
Rule
| declared inputs | m[:x...]::ExponentialFamily.NormalMeanVariance, q[:v]::BayesBase.PointMass |
|---|---|
| algorithm | DefaultAlgorithm() |
| log scale | none |
| defined | rules.md:188 |
| body | (ctx, args)->begin
average = getresult(call_message_update_rule(Average, :out; m = (x = args.m[:x],), q = (v = args.q[:v],), ctx))
NormalMeanVariance(mean(average) + 1.0, var(average))
end |
The keyword call allocates its arguments. Where that matters, call the positional message_passing_rule(Average, Target(:out), algorithm, RuleArgs(m = …, q = …), ctx) instead, which allocates nothing, or message_passing_marginalrule and message_passing_average_energy for the other kinds (see Resolving without the interactive layer).
In-place rules
A rule that writes its result into a buffer declares inplace = true. It also declares preallocate, a function over the slots (algo, ctx, args) that builds the buffer. Its body takes output first and returns it. An engine may keep the buffer between calls. buffer_like builds storage of the right kind from an input.
julia> struct Double end
julia> @define_factor_node(node = Double, type = Deterministic, interfaces = [:out, :in])
julia> @define_message_update_rule(
node = Double, target = :out, args = (m[:in]::Vector{Float64},),
inplace = true,
preallocate = (args) -> MessagePassingRulesBase.buffer_like(args.m[:in]),
body = (output, args) -> (output .= 2 .* args.m[:in]),
)
julia> getresult(@call_message_update_rule(node = Double, target = :out, m = (in = [1.0, 2.0],)))
2-element Vector{Float64}:
2.0
4.0MessagePassingRulesBase.buffer_like — Function
buffer_like(x)
buffer_like(x::AbstractArray, T::Type)Uninitialised storage shaped like x, with element type T if given, for an in-place rule's preallocate. It dispatches on the type of x, so the buffer is of a kind that matches: similar for arrays, which keeps a static array static-sized and mutable and a device array on its device; elementwise for tuples and named tuples. Array types and devices that need something else add a method.
Throws
ArgumentError for a number, which has no storage to write into.
Examples
julia> using MessagePassingRulesBase: buffer_like
julia> b = buffer_like([1.0, 2.0]); (typeof(b), size(b))
(Vector{Float64}, (2,))Scratch
A rule that needs working memory, its scratch, declares how to build it from its inputs. The body then takes it as its scratch slot:
using MessagePassingRulesBase
struct Summing end # out = 2 · sum(in), for a vector in
@define_factor_node(node = Summing, type = Deterministic, interfaces = [:out, :in])
@define_message_update_rule(
node = Summing, target = :out, args = (m[:in]::Vector{Float64},),
scratch = (args) -> (work = similar(args.m[:in]),),
body = (scratch, args) -> (scratch.work .= 2 .* args.m[:in]; sum(scratch.work)),
)
@call_message_update_rule(node = Summing, target = :out, m = (in = [1.0, 2.0],))Result
| value | 6.0 |
|---|---|
| type | Float64 |
| log scale | undefined: the message rule for Summing towards :out under DefaultAlgorithm declares no `logscale` |
Inputs
| edge | value |
|---|
Rule
| declared inputs | m[:in]::Vector{Float64} |
|---|---|
| algorithm | DefaultAlgorithm() |
| scratch | (work = [2.0, 4.0],) |
| log scale | none |
| defined | rules.md:250 |
| body | (scratch, args)->begin
scratch.work .= 2 .* args.m[:in]
sum(scratch.work)
end |
An engine keeps one scratch per outbound stream. It builds the scratch at the first call and passes the same one to every later call, so the memory is allocated once.
A scratch is write-before-read:
- A rule never relies on what an earlier call left in the scratch. The engine may drop or rebuild it at any time. The rule's result therefore depends on its inputs alone, and the rule stays pure.
- The scratch never leaves the rule. A rule does not return it, or a view into it.
- The scratch is never shared with another rule, not even one of the same node.
Scratch is independent of inplace. A rule may declare both, and its body then takes output and then scratch.
The scratch's type depends on the types of the inputs, their element types included. An engine that keeps the scratch between calls infers its type from them, with rule_scratch_type, and asserts the kept scratch to that type. Inputs of other types get a scratch of their own. You declare nothing for this.
A builder whose result type infers runs the rule on a concretely typed scratch. Builders made of similar or zeros(eltype(...), ...) over the inputs infer. A builder that does not infer, say because it reads a global, runs the rule on an untyped scratch, which costs a dynamic call per call.
Annotations
A rule may record facts about its result beside it, keyed by symbol, with annotate!(ann, key, value). It reads the annotations its inputs arrived with as ann.m[:x] and ann.q[:x]. Annotations never take part in dispatch.
The caller decides where the rule's annotations go. An engine passes its own store. A call by hand passes an AnnotationStore to collect them, or nothing to drop them. A message's log scale is not an annotation: a rule declares it with logscale (see Log scales).
julia> struct Solver end
julia> @define_factor_node(node = Solver, type = Deterministic, interfaces = [:out, :in])
julia> @define_message_update_rule(
node = Solver, target = :out, args = (m[:in]::Real,),
body = (args, ann) -> (MessagePassingRulesBase.annotate!(ann, :iterations, 3); sqrt(args.m[:in])),
)
julia> store = MessagePassingRulesBase.AnnotationStore();
julia> result = @call_message_update_rule(node = Solver, target = :out, m = (in = 4.0,), ann = store);
julia> MessagePassingRulesBase.getannotation(getannotations(result), :iterations)
3MessagePassingRulesBase.RuleAnnotations — Type
RuleAnnotations(; m = NamedTuple(), q = NamedTuple(), out = NoAnnotations())The ann a rule body receives, carrying annotations in both directions. The ones that arrived with the inputs are keyed exactly like them, ann.m[:out] and ann.q[:μ]; the rule writes its own with annotate!(ann, key, value), which lands in out, and reads them back with hasannotation and getannotation. Annotations never take part in dispatch.
Keywords
m: the annotations of the inbound messages, aNamedTuplekeyed likeargs.m. Default: none.q: the annotations of the marginals, keyed likeargs.q. Default: none.out: where the rule's own annotations go, anAnnotationStoreor aNoAnnotations. Default:NoAnnotations(), dropping them.
MessagePassingRulesBase.AnnotationStore — Type
AnnotationStore()A mutable store of annotations keyed by symbol, which collects what a rule writes with annotate!. Its dictionary is created on the first write, so an empty store allocates nothing more. Pass one as a call's ann to read afterwards what the rule annotated, through getannotations(result).
julia> using MessagePassingRulesBase: AnnotationStore, annotate!, hasannotation, getannotation
julia> store = AnnotationStore(); annotate!(store, :iterations, 3);
julia> hasannotation(store, :iterations), getannotation(store, :iterations), getannotation(store, :other, 0)
(true, 3, 0)See also NoAnnotations, RuleAnnotations.
MessagePassingRulesBase.NoAnnotations — Type
NoAnnotations()An annotation store that records nothing: annotate! on it does nothing, hasannotation is always false, and getannotation finds nothing. It has no fields, so a call that collects no annotations pays nothing for the rule's writes. It is the default for the message_passing_* calls, and the out of a RuleAnnotations unless one is given.
MessagePassingRulesBase.annotate! — Function
annotate!(annotations, key::Symbol, value) -> nothingRecord value under key, replacing any earlier value. annotations is an AnnotationStore, a NoAnnotations, which drops it, or a rule's ann, a RuleAnnotations, which writes to its out. A rule body calls it as annotate!(ann, key, value).
MessagePassingRulesBase.hasannotation — Function
hasannotation(annotations, key::Symbol) -> BoolWhether annotations holds a value under key; always false for a NoAnnotations. On a rule's ann it looks at what the rule wrote, not at the inputs' annotations.
MessagePassingRulesBase.getannotation — Function
getannotation(annotations, key::Symbol)
getannotation(annotations, key::Symbol, default)The annotation recorded under key, or default when there is none. On a rule's ann it reads what the rule wrote, not the inputs' annotations.
Throws
KeyError when key holds nothing and no default is given.
Checking inputs
A rule's args say which types its inputs have. Some rules need more: a probability between 0 and 1, a vector of a given length, a univariate marginal where the rule takes Any. The args_check keyword checks that as the body starts. It is a function of the inputs, over the same slots as logscale, (algo, ctx, args), and returns:
truewhen the inputs pass;falsewhen they do not, and the rule raises aRuleInputErrorthat quotes the check's source;- a string when they do not, and the error says that instead. Return one when the check combines several conditions, whose source would read poorly, or when the message should name the offending value.
using MessagePassingRulesBase, BayesBase, ExponentialFamily
struct Coin end # f(out, p) = Bernoulli(out | p)
@define_factor_node(node = Coin, type = Stochastic, interfaces = [:out, :p])
@define_message_update_rule(
node = Coin, target = :out, args = (m[:p]::PointMass,),
args_check = (args) -> 0 <= mean(args.m[:p]) <= 1 || lazy"`p` is a probability, between 0 and 1; got $(mean(args.m[:p]))",
body = (args) -> Bernoulli(mean(args.m[:p])),
)
try
@call_message_update_rule(node = Coin, target = :out, m = (p = PointMass(1.5),))
catch err
print(first(split(sprint(showerror, err), "\n rule at")))
endRuleInputError: the message rule for Coin towards :out under DefaultAlgorithm refuses its inputs: `p` is a probability, between 0 and 1; got 1.5
inputs: m[:p]::BayesBase.PointMass{Float64}The check runs wherever the rule runs: in an engine, in a call by hand and in a test. A failed check is an error, not a reason to select another rule, since resolution has already chosen this one. A combination of inputs a node does not support at all is a rule of its own whose body raises the error, found by dispatch like any other.
Cost. A check that reads only the inputs' types, such as args_check = (args) -> variate_form(typeof(args.q[:y])) === Univariate, costs nothing: the rule is compiled for those types, and the check folds away. A check that reads a value costs one comparison. The error path is outlined, so the rule allocates nothing until a check fails. A rule with a preallocate or scratch helper runs it first: the check guards the body, not the helpers.
Performance trap. Build the message only where the check fails. Placed after ||, as above, it is: cond || "...$(x)" builds nothing on a passing call, and costs what the plain condition does. A message bound before the condition, msg = "...$(x)"; cond || msg, or formatted by a helper before it decides, is built on every call and allocates there. Writing it as a lazy"..." string is the safe habit: the formatting waits until the error is shown, wherever the string is made.
Working types
A rule may compute in a type chosen for the arithmetic rather than for the reader. It then declares how that type converts to the type users expect, with public_equivalent. An engine applies the conversion to every marginal it forms.
MessagePassingRulesBase.public_equivalent — Function
public_equivalent(d)d as the type users expect: the same distribution, converted from an efficient working type to its public counterpart. The identity by default.
Rules may compute in types chosen for the arithmetic rather than for the reader. ExponentialFamily's WishartFast, for instance, stores the inverse of its scale matrix, which is what the Wishart rules and products need, while users and most code expect Distributions' Wishart. A package whose rules return such a working type adds a method for it:
MessagePassingRulesBase.public_equivalent(d::WishartFast) = convert(Wishart, d)An engine applies it to every marginal it forms from messages, so a posterior reaches its reader, and the rules that consume that marginal, in the public type. A method must return the same distribution, up to rounding; it is a change of representation, never of meaning. A working type that no reader should see is its use case; a type that is public already needs no method.