Defining rules

A node computes three kinds of thing, and you define each with its own macro:

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 group m. The name k is bound to the member's index, an Int, in body and in the preallocate, scratch, logscale and args_check functions, without being listed among their parameters; args can select by it.
  • args: the inputs the rule consumes, a tuple of entries container[key]::T, the container m for a message or q for a marginal. A type left out is Any; 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, or q[(: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 group in. (:in, 1), with a literal index, is one member: q[:out, (:in, 1)] is the joint of out and in's first member.

    For a group, whose value is a tuple in member order:

    • m[:in...]::T: every member, each of type T;
    • m[:in][k]::T: the target's own member, k being 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 nothing where the selection leaves a member out, so args.m[:in][k] is member k whatever was selected.

    default among 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 beside default, and walks them with rule_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 declares scratch;
    • algo: the algorithm value the rule runs under, and so its parameters;
    • ctx: the RuleContext, whose services the rule reads as ctx.name;
    • args: the inputs, read as args.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 with annotate!(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), usually DefaultAlgorithm; the node must then be declared before the rule is loaded. Almost every rule leaves it out; naming one is for a rule switcher, a DefaultAlgorithmExtension, or a node's own algorithm. With a parametric algorithm type T, algorithm = T matches every T{…}, 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 with message = exp(logscale) · result for the normalised result the 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, or logscale = loghalf (StatsFuns'); an Irrational keeps the message's float type, while -logtwo would not, being a Float64;
    • 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 returns with_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 UndefinedLogScale naming the rule, which propagates through products; only require_logscale turns it into an error. Declaring improper gives an undefined log scale too, whose reason says that none exists, where an omitted declaration says that it is not known.

  • reads_logscale: true if the rule reads the log scales of its inbound messages, as args.logscale.m[:x]. Its caller must then provide them: an engine does when it tracks log scales, and a call by hand takes them as logscale = (...); 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,) or ctx = (: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: true for a rule that writes its result into a buffer it is given rather than allocating one. It needs preallocate, and its body takes output first. 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 with inplace = true.

  • scratch: working memory, a function building it from the inputs over the slots (algo, ctx, args), given to the body as its scratch slot: 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 with inplace. The body takes scratch exactly 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 with similar or zeros(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: false for a rule with side effects, true for a pure rule under an impure algorithm. Default: the algorithm's declaration, ispure, true for almost every algorithm. A pure rule mutates neither its inputs nor state shared beyond one call, and draws randomness only from ctx.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 an Any input must have. A function over the slots (algo, ctx, args), in that order, returning

    • true when the inputs pass;
    • false when they do not: the rule raises a RuleInputError quoting 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 preallocate or scratch helper, 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. A lazy"..." 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] …,
)
source
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 in body and in the preallocate and scratch functions, without being listed among their parameters. Used with default in args, for one rule over every factorisation.
  • args: the inputs the rule consumes, a tuple of entries container[key]::T, the container m for a message or q for a marginal. A type left out is Any; 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, or q[(: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 group in. (:in, 1), with a literal index, is one member: q[:out, (:in, 1)] is the joint of out and in's first member.

    For a group, whose value is a tuple in member order:

    • m[:in...]::T: every member, each of type T;
    • m[:in][k]::T: the target's own member, k being 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 nothing where the selection leaves a member out, so args.m[:in][k] is member k whatever was selected.

    default among 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 beside default, and walks them with rule_inputs: one rule for every factorisation.

    A marginal rule typically reads the messages on the cluster's members, m[:y] and m[: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 declares scratch;
    • algo: the algorithm value the rule runs under, and so its parameters;
    • ctx: the RuleContext, whose services the rule reads as ctx.name;
    • args: the inputs, read as args.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 with annotate!(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), usually DefaultAlgorithm; the node must then be declared before the rule is loaded. Almost every rule leaves it out; naming one is for a rule switcher, a DefaultAlgorithmExtension, or a node's own algorithm. With a parametric algorithm type T, algorithm = T matches every T{…}, 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,) or ctx = (: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: true for a rule that writes its result into a buffer it is given rather than allocating one. It needs preallocate, and its body takes output first. 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 with inplace = true.

  • scratch: working memory, a function building it from the inputs over the slots (algo, ctx, args), given to the body as its scratch slot: 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 with inplace. The body takes scratch exactly 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 with similar or zeros(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: false for a rule with side effects, true for a pure rule under an impure algorithm. Default: the algorithm's declaration, ispure, true for almost every algorithm. A pure rule mutates neither its inputs nor state shared beyond one call, and draws randomness only from ctx.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 an Any input must have. A function over the slots (algo, ctx, args), in that order, returning

    • true when the inputs pass;
    • false when they do not: the rule raises a RuleInputError quoting 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 preallocate or scratch helper, 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. A lazy"..." 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) -> …,
)
source
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 entries q[key]::T. A type left out is Any; 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, or q[(:y, :x)]::T: the joint marginal of a structural cluster, its members in interface order; q[(:in,)] is the joint over the group in, 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.

    default among the entries stands for whatever clusters the graph's factorisation delivers; the energy then takes them all, requires the typed entries beside default, and walks them with rule_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: the RuleContext, whose services the energy reads as ctx.name;
    • args: the marginals, read as args.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), usually DefaultAlgorithm; the node must then be declared before the rule is loaded. Almost every rule leaves it out; naming one is for a rule switcher, a DefaultAlgorithmExtension, or a node's own algorithm. With a parametric algorithm type T, algorithm = T matches every T{…}, 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,) or ctx = (: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: false for a rule with side effects, true for a pure rule under an impure algorithm. Default: the algorithm's declaration, ispure, true for almost every algorithm. A pure rule mutates neither its inputs nor state shared beyond one call, and draws randomness only from ctx.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 an Any input must have. A function over the slots (algo, ctx, args), in that order, returning

    • true when the inputs pass;
    • false when they do not: the rule raises a RuleInputError quoting 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 preallocate or scratch helper, 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. A lazy"..." 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,
)
source

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))
:out

See also IndexedTarget, ClusterTarget.

source
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.

source
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.

source

Inputs

args lists what a rule consumes:

  • m[:x], the message on x;
  • q[:x], the marginal of x;
  • q[:y, :x], the joint marginal of the cluster (y, x);
  • m[:in...], every member of the group in;
  • 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.0

The 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, a Messages;
  • args.q: the marginals, a Marginals;
  • args.logscale: the log scales that arrived with the messages, a RuleLogScales read as args.logscale.m[:out], when the caller tracks them, and nothing otherwise.

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, a NamedTuple or a Messages. Default: none.
  • q: the marginals of single interfaces, a NamedTuple or a Marginals. Joints need the Marginals constructor. Default: none.
  • logscale: the messages' log scales, a NamedTuple keyed like m, or nothing when 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)
source
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, :μ))

See also Marginals, RuleArgs.

source
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

  • ArgumentError when the numbers of cluster keys and joints differ;
  • KeyError when 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)

See also Messages, RuleArgs.

source
MessagePassingRulesBase.canonical_cluster_keys — Function
canonical_cluster_keys(keys) -> Tuple

Sort 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))
source

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),),
)
RuleResult: message of Average towards :outmessages and marginals
vx[1]x[2]outAverage
Result
valueExponentialFamily.NormalMeanVariance{Float64}(μ=2.0, v=2.0)
typeExponentialFamily.NormalMeanVariance{Float64}
log scaleundefined: the message rule for Average towards :out under DefaultAlgorithm declares no `logscale`
Inputs
edgevalue
vqPointMass{Float64}(2.0)
x[1]mExponentialFamily.NormalMeanVariance{Float64}(μ=1.0, v=1.0)
x[2]mExponentialFamily.NormalMeanVariance{Float64}(μ=3.0, v=1.0)
Rule
declared inputsdefault, q[:v]::BayesBase.PointMass
algorithmDefaultAlgorithm()
log scalenone
definedrules.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)),
)
RuleResult: message of Average towards :outvariational
vx[1]x[2]outAverage
Result
valueExponentialFamily.NormalMeanVariance{Float64}(μ=2.0, v=2.0)
typeExponentialFamily.NormalMeanVariance{Float64}
log scaleundefined: the message rule for Average towards :out under DefaultAlgorithm declares no `logscale`
Inputs
edgevalue
vqPointMass{Float64}(2.0)
x[1]qExponentialFamily.NormalMeanVariance{Float64}(μ=1.0, v=1.0)
x[2]qExponentialFamily.NormalMeanVariance{Float64}(μ=3.0, v=1.0)
Rule
declared inputsdefault, q[:v]::BayesBase.PointMass
algorithmDefaultAlgorithm()
log scalenone
definedrules.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)
source

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:

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),),
)
RuleResult: message of ShiftedAverage towards :outmessages and marginals
vx[1]x[2]outShiftedAverage
Result
valueExponentialFamily.NormalMeanVariance{Float64}(μ=3.0, v=2.0)
typeExponentialFamily.NormalMeanVariance{Float64}
log scaleundefined: the message rule for ShiftedAverage towards :out under DefaultAlgorithm declares no `logscale`
Inputs
edgevalue
vqPointMass{Float64}(2.0)
x[1]mExponentialFamily.NormalMeanVariance{Float64}(μ=1.0, v=1.0)
x[2]mExponentialFamily.NormalMeanVariance{Float64}(μ=3.0, v=1.0)
Rule
declared inputsm[:x...]::ExponentialFamily.NormalMeanVariance, q[:v]::BayesBase.PointMass
algorithmDefaultAlgorithm()
log scalenone
definedrules.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.0
MessagePassingRulesBase.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,))
source

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],))
RuleResult: message of Summing towards :outbelief propagation
inoutSumming
Result
value6.0
typeFloat64
log scaleundefined: the message rule for Summing towards :out under DefaultAlgorithm declares no `logscale`
Inputs
edgevalue
inm[1.0, 2.0]
Rule
declared inputsm[:in]::Vector{Float64}
algorithmDefaultAlgorithm()
scratch(work = [2.0, 4.0],)
log scalenone
definedrules.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)
3
MessagePassingRulesBase.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, a NamedTuple keyed like args.m. Default: none.
  • q: the annotations of the marginals, keyed like args.q. Default: none.
  • out: where the rule's own annotations go, an AnnotationStore or a NoAnnotations. Default: NoAnnotations(), dropping them.
source
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.

source
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.

source

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:

  • true when the inputs pass;
  • false when they do not, and the rule raises a RuleInputError that 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")))
end
RuleInputError: 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.

source