From d472db156a8dc038ccd4cccf2ab5210f58026a23 Mon Sep 17 00:00:00 2001 From: Bagaev Dmitry Date: Tue, 6 Oct 2026 10:03:36 +0200 Subject: [PATCH 1/2] Address the review in #9: improper log scales, a new-scheme tutorial, docs - `logscale = improper` declares a message rule whose message has no normalising constant, as an exact message may (the likelihood of a variance integrates to infinity): its log scale is an UndefinedLogScale with cause `:improper`, and require_logscale says that none exists, where an omitted `logscale` says that it is not known. - A tutorial, "A new message passing scheme": natural-gradient message passing for a Poisson count with a log rate, as an algorithm, a dependency declaration and a rule, iterated to its fixed point; it cites Lukashchuk et al., Information Geometry of Message Passing (2026). - What the registry is for, with an example; "resolution" defined in the glossary and linked from each page; purity explained by what a pure rule may and may not do; improper messages on the Log scales page and in the glossary; a shorter first example on the overview; the expectation propagation entry pointing to the tutorial that builds such a rule; the default-scheme example told through its own node; the first tutorial's opening reworded. - A Makefile: test (as CI runs it, with test_args), test-fast, docs, docs-serve (LiveServer in the docs environment), format, check-format, clean. Co-Authored-By: Claude Opus 5.5 (1M context) --- CHANGELOG.md | 19 +++++ Makefile | 40 ++++++++++ README.md | 11 ++- docs/Project.toml | 1 + docs/make.jl | 1 + docs/src/algorithms.md | 30 ++++++-- docs/src/calling.md | 2 +- docs/src/fallbacks.md | 2 +- docs/src/glossary.md | 14 +++- docs/src/index.md | 45 ++++-------- docs/src/inspecting.md | 52 +++++++++++-- docs/src/internals.md | 1 + docs/src/keywords.md | 3 +- docs/src/logscales.md | 48 +++++++++++- docs/src/rules.md | 2 +- docs/src/tutorials/algorithm.md | 2 + docs/src/tutorials/first-node.md | 3 +- docs/src/tutorials/new-scheme.md | 122 +++++++++++++++++++++++++++++++ src/MessagePassingRulesBase.jl | 2 +- src/algorithms.jl | 12 +-- src/logscale.jl | 34 ++++++++- src/registry.jl | 3 +- src/result_show.jl | 1 + src/rule_macro.jl | 9 ++- src/rulespec.jl | 4 +- test/logscale_tests.jl | 32 +++++++- 26 files changed, 422 insertions(+), 73 deletions(-) create mode 100644 Makefile create mode 100644 docs/src/tutorials/new-scheme.md diff --git a/CHANGELOG.md b/CHANGELOG.md index c9567d1..0a4314d 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -6,6 +6,25 @@ All notable changes to MessagePassingRulesBase.jl are documented here. The forma ## [Unreleased] +### Added + +- `logscale = improper`, a declaration for a message rule whose message has no normalising + constant, as an exact message may: its log scale is an `UndefinedLogScale` whose cause is + `:improper`, and `require_logscale` says that none exists, where an omitted `logscale` says that + the log scale is not known (#9). +- A tutorial, *A new message passing scheme*: natural-gradient message passing for a Poisson count + with a log rate, as an algorithm, a dependency declaration and a rule (#9). +- A Makefile: `make test` (as CI runs it, with `test_args` to select items), `test-fast`, `docs`, + `docs-serve`, `format`, `check-format`, `clean` (#9). + +### Changed + +- Documentation, from the review in #9: what a pure rule may and may not do; what the registry is + for, with an example; *resolution* defined in the glossary and linked; improper messages on the + *Log scales* page; a shorter first example on the overview; the expectation propagation entry + pointing to the tutorial that builds such a rule; and an example in *Algorithms and + dependencies* told through its own node. + ## [1.0.0] The first release: the rule system of the ReactiveMP ecosystem, developed in the diff --git a/Makefile b/Makefile new file mode 100644 index 0000000..ab48d29 --- /dev/null +++ b/Makefile @@ -0,0 +1,40 @@ +SHELL := /bin/bash +JULIA ?= julia +.DEFAULT_GOAL := help + +# The variables CI sets: the items tagged `:slow` run, and FastCholesky throws on a matrix that is +# not symmetric instead of warning, so a local run fails where CI would. +CI_ENV := TEST_ALL=true JULIA_FASTCHOLESKY_THROW_ERROR_NON_SYMMETRIC=1 + +# Optional selection, passed to the suite: `make test test_args="tag:alloc name:routing"`. +test_args ?= +TEST_ARGS := $(if $(test_args),test_args = split("$(test_args)") .|> string,) + +.PHONY: help test test-fast format check-format docs docs-serve clean + +help: ## Show this help + @grep -E '^[a-zA-Z_-]+:.*?## .*$$' $(MAKEFILE_LIST) \ + | awk 'BEGIN {FS = ":.*?## "}; {printf " \033[36m%-14s\033[0m %s\n", $$1, $$2}' + +test: ## Run the suite as CI does; select with test_args="tag: name: " + $(CI_ENV) $(JULIA) --project=. -e 'import Pkg; Pkg.test($(TEST_ARGS))' + +test-fast: ## Run the suite without the items tagged :slow + $(JULIA) --project=. -e 'import Pkg; Pkg.test($(TEST_ARGS))' + +format: ## Format the tree with Runic + $(JULIA) -e 'import Pkg; Pkg.activate(temp = true); Pkg.add("Runic"); using Runic; exit(Runic.main(["--inplace", "."]))' + +check-format: ## Check the formatting with Runic, without changing files + $(JULIA) -e 'import Pkg; Pkg.activate(temp = true); Pkg.add("Runic"); using Runic; exit(Runic.main(["--check", "--diff", "."]))' + +docs: ## Build the documentation into docs/build, running every example and doctest + $(JULIA) --project=docs -e 'import Pkg; Pkg.instantiate()' + $(JULIA) --project=docs docs/make.jl + +docs-serve: ## Build and serve the documentation with live reload + $(JULIA) --project=docs -e 'import Pkg; Pkg.instantiate(); using LiveServer; servedocs()' + +clean: ## Remove the documentation build and coverage files + rm -rf docs/build lcov.info + find . \( -name '*.jl.cov' -o -name '*.jl.*.cov' -o -name '*.jl.mem' \) -delete diff --git a/README.md b/README.md index 39c91e5..10bea7c 100644 --- a/README.md +++ b/README.md @@ -35,12 +35,11 @@ result = @call_message_update_rule(node = Shift, target = :out, m = (in = 1.0,)) getresult(result), getlogscale(result) # (2.0, 0) ``` -- Documentation: . Build it - locally with `julia --project=docs -e 'import Pkg; Pkg.instantiate()'` and then - `julia --project=docs docs/make.jl`, into `docs/build`. -- Tests: `julia --project -e 'import Pkg; Pkg.test()'`; with `TEST_ALL=true` in the environment - they include the items tagged `:slow`, and `Pkg.test(test_args = ["tag:alloc"])` selects by tag, - `name:` by name and a path by file. +- Documentation: . `make docs` + builds it locally, into `docs/build`. +- Tests: `make test` runs the suite as CI does; `make test test_args="tag:alloc name:routing"` + selects items by tag, by name or by file. `make help` lists the other targets: `docs`, + `docs-serve`, `format`, `check-format`, `test-fast`. - Depends on BayesBase, MacroTools, Compat, FastCholesky, IrrationalConstants and LinearAlgebra: no distribution package and no engine. Julia 1.10 or later. MIT licence. diff --git a/docs/Project.toml b/docs/Project.toml index 7e758ee..d266a74 100644 --- a/docs/Project.toml +++ b/docs/Project.toml @@ -4,6 +4,7 @@ Distributions = "31c24e10-a181-5473-b8eb-7969acd0382f" Documenter = "e30172f5-a6a5-5a46-863b-614d45cd2de4" DocumenterInterLinks = "d12716ef-a0f6-4df4-a9f1-a5a34e75c656" ExponentialFamily = "62312e5e-252a-4322-ace9-a5f4bf9b357b" +LiveServer = "16fef848-5104-11e9-1b77-fb7a48bbb589" MessagePassingRulesBase = "4a5b64c7-30c9-4471-82f2-bbe8464b157d" [sources] diff --git a/docs/make.jl b/docs/make.jl index f9241e4..4789cc0 100644 --- a/docs/make.jl +++ b/docs/make.jl @@ -19,6 +19,7 @@ makedocs( "Your first node" => "tutorials/first-node.md", "A deterministic node with a group" => "tutorials/groups.md", "A node with its own algorithm" => "tutorials/algorithm.md", + "A new message passing scheme" => "tutorials/new-scheme.md", ], "Defining nodes" => "nodes.md", "Defining rules" => "rules.md", diff --git a/docs/src/algorithms.md b/docs/src/algorithms.md index 3ca9b70..215609d 100644 --- a/docs/src/algorithms.md +++ b/docs/src/algorithms.md @@ -56,7 +56,7 @@ MessagePassingRulesBase.default_algorithm ## Extending the default A subtype of [`DefaultAlgorithmExtension`](@ref) overrides some rules, or some dependencies, of -the default, and inherits the rest. Resolution looks for the extension's own rule first. When +the default, and inherits the rest. [Resolution](@ref glossary-resolution) looks for the extension's own rule first. When there is none, it falls back to the default's rule. That rule then runs with `DefaultAlgorithm()` in its `algo` slot, the algorithm it was written for. @@ -122,8 +122,25 @@ A node may name a parametric algorithm as its own, `algorithm = T`. A rule that ## Purity -A rule is pure unless it is declared otherwise. A pure rule mutates neither its inputs nor any -state shared beyond one call. It draws randomness only from `ctx.rng`, which the caller owns. +A rule is pure unless it is declared otherwise. Pure means that a call has no effect anyone +outside it can observe, apart from its result: running the rule leaves nothing changed that its +caller did not hand it to change. It does not mean that the rule writes nothing. + +A pure rule may + +- allocate whatever intermediate arrays it needs; +- write its output into the buffer an [in-place rule](@ref glossary-in-place-rule) is given, + since that buffer is handed over to hold the result; +- reuse its own [scratch](@ref glossary-scratch), working memory that nothing else reads; +- draw random numbers from `ctx.rng`, a generator its caller owns and passes in; +- warn or log, for a degenerate input say, which changes no state another computation reads. + +A rule is impure if it + +- mutates an input, a message or a marginal it reads; +- mutates its algorithm, a cache kept in one of its fields for instance; +- writes a global variable, a file or any other state that later code reads; +- carries its own random number generator in its algorithm. An algorithm that carries state of its own is impure, and it says so with a method of [`ispure`](@ref). A rule overrides its algorithm's purity with `pure = false` or `pure = true`. @@ -199,9 +216,10 @@ MessagePassingRulesBase.free_energy_partition A rule may need one input that the factorisation does not give it, while its other inputs follow the factorisation as usual. `default` among a target's inputs stands for the inputs of the -default scheme, and the inputs beside it are added to them. ContinuousTransition's rule towards -`a`, for example, reads `q(a)`, the expansion point of its transformation, under both of its -factorisations. The node below declares the same: +default scheme, and the inputs beside it are added to them. Take a transition `y ~ f(x; a)` whose +function is linearised around the mean of `a`. Its rule towards `a` needs `q(a)`, the expansion +point, under every factorisation, including one where the default scheme gives it the message +on `a` instead. The node below declares that: ```@example algorithms-extend using MessagePassingRulesBase diff --git a/docs/src/calling.md b/docs/src/calling.md index de73d0b..563d500 100644 --- a/docs/src/calling.md +++ b/docs/src/calling.md @@ -134,7 +134,7 @@ MessagePassingRulesBase.RuleInputError Code that builds its own [`RuleArgs`](@ref) resolves a rule with the `find_*` functions. They never throw: they return a [`RuleNotFound`](@ref) when nothing fits. The code then runs the rule -with the `message_passing_*` functions, which do throw. Julia's dispatch does the resolution, +with the `message_passing_*` functions, which do throw. Julia's dispatch does the [resolution](@ref glossary-resolution), over every loaded package. ```@docs diff --git a/docs/src/fallbacks.md b/docs/src/fallbacks.md index 33234b2..8766e79 100644 --- a/docs/src/fallbacks.md +++ b/docs/src/fallbacks.md @@ -9,7 +9,7 @@ reporting the error. A rule fallback is a callable, `fallback(node, target, args the node, the target and the rule's arguments. It returns a [message](@ref glossary-message), or `nothing` when it has none either. -An engine consults the fallback only when resolution returns a [`RuleNotFound`](@ref). A fallback +An engine consults the fallback only when [resolution](@ref glossary-resolution) returns a [`RuleNotFound`](@ref). A fallback therefore never replaces a rule that exists, and an error inside a rule is never turned into a fallback. ReactiveMP takes a fallback as its activation option `rulefallback`. A message that a fallback computes has an undefined [log scale](@ref glossary-log-scale). diff --git a/docs/src/glossary.md b/docs/src/glossary.md index eb74ced..5eab7cc 100644 --- a/docs/src/glossary.md +++ b/docs/src/glossary.md @@ -64,7 +64,9 @@ are always its output and the joint over its inputs. See [`Deterministic`](@ref) ### [Expectation propagation](@id glossary-expectation-propagation) A message passing scheme that approximates each message by projecting the corresponding -marginal onto a simpler family, usually by matching moments. The Probit node uses it. +marginal onto a simpler family, usually by matching moments. +[A node with its own algorithm](@ref tutorial-algorithm) builds a node whose rule towards its +input is an expectation propagation rule. ### [Factor graph](@id glossary-factor-graph) @@ -110,7 +112,8 @@ convention. An interface may have aliases, other names a model can use for it. The logarithm of a message's normalising constant: a message is ``\exp(\text{log scale}) \cdot p(x)`` with ``p`` a normalised distribution. Summed over a graph, -log scales give the model's evidence. See [Log scales](@ref). +log scales give the model's evidence. Only a message with a finite integral has one: an improper +message, which an exact message can be, has none. See [Log scales](@ref). ### [Marginal](@id glossary-marginal) @@ -137,6 +140,13 @@ enters a rule. `PointMass(2.0)` comes from BayesBase. The distribution of `f(x)` when `x` has a known distribution. A deterministic node's message towards its output is the pushforward of the messages on its inputs. +### [Resolution](@id glossary-resolution) + +Finding the rule that runs for a call: the node, the target, the +[algorithm](@ref glossary-algorithm) and the types of the inputs select one rule. It is Julia's +method dispatch over the methods the definition macros generate, so it needs no list of rules. +When no rule matches, the result is a [`RuleNotFound`](@ref). See [Calling rules](@ref). + ### [Rule](@id glossary-rule) A function that computes a message, a joint marginal or an average energy of a node from its diff --git a/docs/src/index.md b/docs/src/index.md index b30c71e..5f76f9f 100644 --- a/docs/src/index.md +++ b/docs/src/index.md @@ -20,7 +20,7 @@ You use the package for three things: A rule is an ordinary Julia function of its inputs. Julia dispatches it on the node, on the target and on the types of the incoming [messages](@ref glossary-message) and -[marginals](@ref glossary-marginal). Because resolution is Julia's own dispatch, a rule defined +[marginals](@ref glossary-marginal). Because [resolution](@ref glossary-resolution) is Julia's own dispatch, a rule defined in any loaded package is found. The package depends on BayesBase and on small numerical packages (FastCholesky, @@ -38,48 +38,29 @@ MessagePassingRulesBase ## A first node and rule -The node below is a [deterministic node](@ref glossary-deterministic-node), `out = in + c`, for -a known shift `c`. It has one rule, for the message towards `out`. The messages are a small -normal type that the example defines, a mean and a variance. A real rule package uses -ExponentialFamily's distributions instead. +A node, `out = in + 1`, with one rule, for the message towards `out`, called as an engine would +call it: ```jldoctest overview julia> using MessagePassingRulesBase -julia> struct Gauss # a normal, by its mean and variance - m::Float64 - v::Float64 - end +julia> struct Shift end # out = in + 1 -julia> struct Shift end - -julia> @define_factor_node(node = Shift, type = Deterministic, interfaces = [:out, :in, :c]) +julia> @define_factor_node(node = Shift, type = Deterministic, interfaces = [:out, :in]) julia> @define_message_update_rule( - node = Shift, target = :out, - args = (m[:in]::Gauss, m[:c]::Real), - logscale = 0, - body = (args) -> Gauss(args.m[:in].m + args.m[:c], args.m[:in].v), + node = Shift, target = :out, args = (m[:in]::Real,), logscale = 0, + body = (args) -> args.m[:in] + 1, ) -julia> result = @call_message_update_rule(node = Shift, target = :out, m = (in = Gauss(1.0, 2.0), c = 3.0)); - -julia> getresult(result) -Gauss(4.0, 2.0) +julia> result = @call_message_update_rule(node = Shift, target = :out, m = (in = 1.0,)); -julia> getlogscale(result) -0 +julia> getresult(result), getlogscale(result) +(2.0, 0) ``` -The node's declaration names its [interfaces](@ref glossary-interface). The rule names its -node, its target, the inputs it takes with their types, and its body. The call runs the rule as -an engine would. It returns a [`RuleResult`](@ref): the message, its -[log scale](@ref glossary-log-scale) and everything that produced them. An engine such as -ReactiveMP builds the node in a graph from the same declaration. It runs the same rule whenever -the rule's inputs change. - -[Your first node](@ref tutorial-first-node) builds a complete node step by step, with -ExponentialFamily's distributions. +[Your first node](@ref tutorial-first-node) builds a real node step by step, and explains each +part. ## The site @@ -87,6 +68,8 @@ ExponentialFamily's distributions. its belief propagation and variational rules. [A deterministic node with a group](@ref tutorial-groups) writes rules for a sum of any number of inputs. [A node with its own algorithm](@ref tutorial-algorithm) gives a node a parametrised algorithm and declares what its rules take. + [A new message passing scheme](@ref tutorial-new-scheme) writes natural-gradient message + passing as an algorithm and its rules. - [Defining nodes](@ref): [`@define_factor_node`](@ref), what a declaration records, and the queries that read it. - [Defining rules](@ref): the three rule macros, targets, the inputs a rule receives, the slots diff --git a/docs/src/inspecting.md b/docs/src/inspecting.md index bb20ebe..8dc5a3f 100644 --- a/docs/src/inspecting.md +++ b/docs/src/inspecting.md @@ -101,7 +101,7 @@ A rule package's tests check its rules in two ways: - against their nodes' declarations, with [`check_rules`](@ref); - against each other, with [`check_rule_ambiguities`](@ref). Two rules that some call matches - equally well make resolution throw a `MethodError`. + equally well make [resolution](@ref glossary-resolution) throw a `MethodError`. Each check takes the modules to check, every loaded module by default, and returns the problems it finds. Both lists are empty for this page's module: @@ -118,11 +118,53 @@ MessagePassingRulesBase.check_rule_ambiguities MessagePassingRulesBase.duplicate_rules ``` -## The registries +## [What the registry is for](@id inspecting-registry) -Each module that defines nodes or rules keeps a registry of what it defined. The registry is -filled when the module loads. It serves introspection only: listings, coverage, checks, and the -near misses a [`RuleNotFoundError`](@ref) reports. Resolution never reads it. +Each module that defines nodes or rules keeps a [`Registry`](@ref) of what it defined: its +[`RuleSpec`](@ref)s, [`NodeSpec`](@ref)s and dependency declarations. The definition macros create +it in the module, as the constant `__message_passing_registry__`, and fill it as the module loads. +Each module keeps its own so that a package's precompiled image carries its own entries; +[`registries`](@ref) gathers them from every loaded module. + +[Resolution](@ref glossary-resolution) does not use the registry. Julia's dispatch finds the rule +for a call, through the methods the macros define. Dispatch cannot say which rules exist, +though, and that is what the registry is for: everything that lists, counts or checks rules reads +it. Take a module that defines a node and two rules: + +```@example registry +using MessagePassingRulesBase + +module Coins +using MessagePassingRulesBase +struct Coin end # out ~ Bernoulli(p) +@define_factor_node(node = Coin, type = Stochastic, interfaces = [:out, :p]) +@define_message_update_rule(node = Coin, target = :out, args = (m[:p]::Real,), body = (args) -> args.m[:p]) +@define_message_update_rule(node = Coin, target = :p, args = (m[:out]::Real,), body = (args) -> args.m[:out]) +end + +length(MessagePassingRulesBase.registered_rules(Coins)) +``` + +The registry lists the module's two rules. [`list_rules`](@ref) and [`rule_coverage`](@ref) read +it, and so do [`check_rules`](@ref) and a test suite's rule-coverage gate: + +```@example registry +MessagePassingRulesBase.rule_coverage(Coins.Coin) +``` + +A call that no rule takes fails with a [`RuleNotFoundError`](@ref), and its near misses, the +rules that come closest, come from the registry too: + +```@example registry +try + @call_message_update_rule(node = Coins.Coin, target = :out, m = (p = "half",)) +catch err + showerror(stdout, err) +end +``` + +A call that matches never consults it. A rule a registry missed would still run; it just would +not be listed, counted, checked or suggested. ```@docs MessagePassingRulesBase.Registry diff --git a/docs/src/internals.md b/docs/src/internals.md index b6bf9ca..e194089 100644 --- a/docs/src/internals.md +++ b/docs/src/internals.md @@ -38,4 +38,5 @@ fragments have no docstrings of their own. MessagePassingRulesBase.default_inputs_match MessagePassingRulesBase.WithLogScale MessagePassingRulesBase.FromBody +MessagePassingRulesBase.Improper ``` diff --git a/docs/src/keywords.md b/docs/src/keywords.md index 7ba0cdf..9bd6ca5 100644 --- a/docs/src/keywords.md +++ b/docs/src/keywords.md @@ -534,11 +534,12 @@ struct Gain end # out = a ⋅ in `logscale` declares the message's [log scale](@ref glossary-log-scale): the scalar with `message = exp(logscale) · result`, for the normalised `result` the body returns. The message towards `in` is ``\mathcal{N}(a x \mid m, v) = |a|^{-1}\, \mathcal{N}(x \mid m/a, v/a^2)``, so its -log scale is ``-\log|a|``, a function of the inputs. The keyword takes one of four forms: +log scale is ``-\log|a|``, a function of the inputs. The keyword takes one of five forms: - a number, `logscale = 0`, when the constant does not depend on the inputs; - a function of the inputs over the slots `(algo, ctx, args)`, named in that order, as above; - `from_body`, when the body returns [`with_logscale`](@ref)`(result, logscale)`; +- [`improper`](@ref), when the message has no normalising constant, so no log scale exists; - nothing: a rule that omits the keyword gives an [`UndefinedLogScale`](@ref) naming it. ```@example messages diff --git a/docs/src/logscales.md b/docs/src/logscales.md index b1598ba..7f95f30 100644 --- a/docs/src/logscales.md +++ b/docs/src/logscales.md @@ -17,7 +17,9 @@ where ``\hat{p}`` is the distribution the rule returns. [Belief propagation](@ref glossary-belief-propagation) is the common case. A message ``\mu(x) = \int f(x, y, \dots) \prod_i \mu_i(y_i) \,\mathrm{d}y`` is in general not normalised. In an acyclic graph, the log scales of such messages add up to the log model evidence, which an -engine can then read at any edge. A naive [variational](@ref glossary-vmp) message, +engine can then read at any edge. A log scale exists only for a message whose integral is finite: +an exact message may be improper, with no normalising constant at all +([Improper messages](@ref logscales-improper)). A naive [variational](@ref glossary-vmp) message, ``\exp \mathbb{E}_q[\log f]``, has no constant with that meaning, so its log scale is undefined. This page covers the rule's side: what a rule declares, and how it reads the log scales of its @@ -28,7 +30,7 @@ and what the evidence means. The normalised distribution a rule returns does not reveal its log scale. `Beta(2, 1)` looks the same whether or not a constant was divided out to get it. A message rule therefore states its -log scale with the `logscale` keyword of [`@define_message_update_rule`](@ref), in one of four +log scale with the `logscale` keyword of [`@define_message_update_rule`](@ref), in one of five forms: | form | when to use it | @@ -36,6 +38,7 @@ forms: | a constant, `logscale = 0` | the rule's constant does not depend on its inputs | | a function of the inputs, `logscale = (args) -> …` | the constant depends on the inputs | | `logscale = from_body` | the constant shares its work with the result | +| `logscale = improper` | the message is improper: it has no constant, so no log scale exists | | the keyword omitted | the log scale is undefined, as for every variational rule | A marginal rule and an average energy have no log scale, and they do not accept the keyword. @@ -116,6 +119,47 @@ julia> getlogscale(@call_message_update_rule(node = Halve, target = :in, m = (ou UndefinedLogScale: the message rule for Halve towards :in under DefaultAlgorithm declares no `logscale` ``` +### [Improper messages](@id logscales-improper) + +A message's normalising constant is its integral, and the integral may be infinite. The message +is then improper, and no log scale exists to declare. An exact belief-propagation message can be +so. Take a normal node with a known mean, `out ~ N(0, v)`, and `out` observed at ``y``. The +message towards the variance is the likelihood + +```math +v \mapsto \mathcal{N}(y \mid 0, v) = (2\pi v)^{-1/2} \exp\!\big(-y^2 / (2v)\big), +``` + +which decays only like ``v^{-1/2}`` as ``v`` grows, so its integral over ``v > 0`` is infinite. +The rule declares `logscale = improper`. Its message's log scale is an +[`UndefinedLogScale`](@ref) whose reason says so, and it propagates as any undefined one does: + +```jldoctest logscales +julia> using MessagePassingRulesBase + +julia> struct Noise end # out ~ N(0, v) + +julia> @define_factor_node(node = Noise, type = Stochastic, interfaces = [:out, :v]) + +julia> @define_message_update_rule( + node = Noise, target = :v, args = (m[:out]::Real,), + logscale = improper, + body = (args) -> (v -> -(log(2π * v) + args.m[:out]^2 / v) / 2), # the log-likelihood + ) + +julia> getlogscale(@call_message_update_rule(node = Noise, target = :v, m = (out = 1.0,))) +UndefinedLogScale: the message rule for Noise towards :v under DefaultAlgorithm gives an improper message: it has no normalising constant +``` + +Omitting the keyword gives an undefined log scale as well, but its reason is that the rule +declares none: a log scale that may exist and that nobody derived. `improper` says that there is +none to derive. A product of an improper message with a proper one can still be normalised, a +proper prior on ``v`` here, but its log scale is undefined too. + +```@docs +MessagePassingRulesBase.improper +``` + ## Undefined log scales An undefined log scale records why it is undefined, and it propagates. Adding it to a number, or diff --git a/docs/src/rules.md b/docs/src/rules.md index 5dfb528..a8a25fd 100644 --- a/docs/src/rules.md +++ b/docs/src/rules.md @@ -352,7 +352,7 @@ end ``` 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 +check is an error, not a reason to select another rule, since [resolution](@ref glossary-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. diff --git a/docs/src/tutorials/algorithm.md b/docs/src/tutorials/algorithm.md index 356b379..acdecca 100644 --- a/docs/src/tutorials/algorithm.md +++ b/docs/src/tutorials/algorithm.md @@ -278,6 +278,8 @@ The node has no average energy, so an engine cannot compute a free energy with i - [Your first node](@ref tutorial-first-node) and [A deterministic node with a group](@ref tutorial-groups) cover rules under the default algorithm. +- [A new message passing scheme](@ref tutorial-new-scheme) writes a whole scheme, + natural-gradient message passing, the same way. - [Algorithms and dependencies](@ref) describes the default scheme, extensions of the default and every form of a dependency. - [Defining nodes](@ref) and [Defining rules](@ref) cover the other keywords of the node and of diff --git a/docs/src/tutorials/first-node.md b/docs/src/tutorials/first-node.md index 5fecd49..44cd27c 100644 --- a/docs/src/tutorials/first-node.md +++ b/docs/src/tutorials/first-node.md @@ -4,7 +4,8 @@ CurrentModule = MessagePassingRulesBase # [Your first node](@id tutorial-first-node) -This tutorial builds a [factor node](@ref glossary-factor-node) from nothing: a normal +This tutorial builds a [factor node](@ref glossary-factor-node) step by step, starting from an +empty module: a normal distribution with a known variance, ```math diff --git a/docs/src/tutorials/new-scheme.md b/docs/src/tutorials/new-scheme.md new file mode 100644 index 0000000..bbe2ac5 --- /dev/null +++ b/docs/src/tutorials/new-scheme.md @@ -0,0 +1,122 @@ +```@meta +CurrentModule = MessagePassingRulesBase +``` + +# [A new message passing scheme](@id tutorial-new-scheme) + +A message passing scheme decides two things for each rule: which inputs it reads, and how it +turns them into a message. [Belief propagation](@ref glossary-belief-propagation), +[variational message passing](@ref glossary-vmp) and +[expectation propagation](@ref glossary-expectation-propagation) answer them differently. A new +scheme needs no new machinery here. It is an [algorithm](@ref glossary-algorithm), a declaration +of what its rules read, and the rules themselves, written with the same macros as any other. + +This tutorial writes one: natural-gradient message passing, after Lukashchuk, Yemets, Ledbetter +and Şenöz, [*Information Geometry of Message Passing*](https://arxiv.org/abs/2608.15922) (2026). +It assumes you have read [A node with its own algorithm](@ref tutorial-algorithm). + +## The scheme + +Take a count observed through a log rate, ``y \sim \operatorname{Poisson}(e^x)``, with a normal +belief ``q(x) = \mathcal{N}(m, v)``. The log factor is + +```math +\log f(y, x) = y x - e^x - \log y!. +``` + +Its exact message towards ``x``, ``x \mapsto f(y, x)``, is not normal. A variational message, +``\exp \mathbb{E}_q[\log f]``, averages the factor under the belief instead. A natural-gradient +message keeps the part of the factor that a normal belief can represent: the normal whose +precision is the expected curvature of ``\log f`` under ``q``, and whose weighted mean follows +the expected slope, + +```math +w = -\mathbb{E}_q\big[\partial_x^2 \log f\big] = \mathbb{E}_q[e^x] = e^{m + v/2}, \qquad +\xi = \mathbb{E}_q\big[\partial_x \log f\big] + w\, m = y - e^{m + v/2} + w\, m. +``` + +Both expectations are in closed form, since ``\mathbb{E}_q[e^x]`` is the mean of a log-normal. The +message reads the belief ``q(x)`` on its own edge, which neither belief propagation nor the +default scheme gives a rule. + +## The algorithm and what it reads + +The scheme is an algorithm, a type with no fields here: + +```@example new-scheme +using MessagePassingRulesBase, BayesBase, ExponentialFamily + +struct NaturalGradient <: AbstractAlgorithm end + +struct PoissonLog end # out ~ Poisson(exp(in)) + +@define_factor_node(node = PoissonLog, type = Stochastic, interfaces = [:out, :in]) +nothing # hide +``` + +Its rule towards `in` reads the observation and the belief on `in` itself, so the node declares +that for the algorithm, with [`@define_dependencies`](@ref). The target `out` follows the default +scheme: + +```@example new-scheme +@define_dependencies( + node = PoissonLog, algorithm = NaturalGradient, + dependencies = [:out => (default,), :in => (q[:out], q[:in])], +) + +MessagePassingRulesBase.dependencies_spec(PoissonLog, NaturalGradient()) +``` + +## The rule + +The rule is the two formulas above, under the algorithm: + +```@example new-scheme +@define_message_update_rule( + node = PoissonLog, target = :in, algorithm = NaturalGradient, + args = (q[:out]::PointMass, q[:in]::UnivariateNormalDistributionsFamily), + body = (args) -> begin + y = mean(args.q[:out]) + m, v = mean_var(args.q[:in]) + w = exp(m + v / 2) # E_q[e^x] + NormalWeightedMeanPrecision(y - w + w * m, w) + end, +) + +@call_message_update_rule( + node = PoissonLog, target = :in, algorithm = NaturalGradient(), + q = (out = PointMass(3), in = NormalMeanVariance(0.0, 1.0)), +) +``` + +## Iterating to a fixed point + +The message depends on the belief it updates, so the two are iterated: the belief is the prior +times the message, and the message is recomputed from the new belief. With a standard normal +prior and the count ``y = 3``: + +```@example new-scheme +prior = NormalMeanVariance(0.0, 1.0) +belief = prior +for iteration in 1:20 + message = getresult(@call_message_update_rule( + node = PoissonLog, target = :in, algorithm = NaturalGradient(), + q = (out = PointMass(3), in = belief), + )) + global belief = prod(ClosedProd(), prior, message) +end +mean_var(belief) +``` + +The belief settles where the scheme is stationary: its precision is the prior's plus +``e^{m + v/2}``, and its mean is ``y - e^{m + v/2}`` for this prior. That is the condition for the +best normal approximation of the posterior in the variational sense: + +```@example new-scheme +m, v = mean_var(belief) +(precision = 1 / v - (1 + exp(m + v / 2)), mean = m - (3 - exp(m + v / 2))) +``` + +An engine does the same iteration on a whole graph, running each rule whenever the inputs it +declared change. Nothing in the rule refers to the engine: a scheme is the algorithm, what its +rules read, and what they compute. diff --git a/src/MessagePassingRulesBase.jl b/src/MessagePassingRulesBase.jl index 2cc4841..204a85f 100644 --- a/src/MessagePassingRulesBase.jl +++ b/src/MessagePassingRulesBase.jl @@ -58,7 +58,7 @@ export @which_message_update_rule, @which_marginal_update_rule, @which_average_e export call_message_update_rule, call_marginal_update_rule, call_average_energy export which_message_update_rule, which_marginal_update_rule, which_average_energy export UndefinedLogScale, UndefinedLogScaleError, require_logscale, isdefined_logscale -export with_logscale, from_body, getlogscale +export with_logscale, from_body, improper, getlogscale export RuleResult, getresult, getrule, getannotations # Generic names a downstream package may well define for itself: public, not exported. @compat public getalgorithm, getcontext, getscratch, getarguments, gettarget diff --git a/src/algorithms.jl b/src/algorithms.jl index 8a72496..34159be 100644 --- a/src/algorithms.jl +++ b/src/algorithms.jl @@ -72,11 +72,13 @@ Whether rules under this algorithm are pure unless they say otherwise. An impure adds a method, `MessagePassingRulesBase.ispure(::Type{<:MyAlgorithm}) = false`, which covers each `MyAlgorithm{T}` of a parametric one as well. -A pure rule mutates neither its inputs nor any state shared beyond one call, such as fields -of its algorithm. It may write to its own output buffer and to scratch storage it owns, so -in-place rules can be pure. Randomness comes from `ctx.rng`, which the caller owns; an -algorithm that carries its own random number generator, or any other state that persists -between calls, is impure. +Pure means that a call has no effect anyone outside it can observe, apart from its result; it +does not mean that the rule writes nothing. A pure rule may allocate intermediate arrays, write +the output buffer an in-place rule is given (the buffer is handed over to hold the result), reuse +its own scratch, which nothing else reads, draw from `ctx.rng`, which its caller owns, and warn or +log. A rule is impure if it mutates an input, mutates its algorithm (a cache kept in a field, say), +writes a global variable, a file or other state that later code reads, or carries its own random +number generator in its algorithm. Purity is declared, not proved: the default is `true`. A rule may override its algorithm with `pure = false`, and an engine auditing purity must read the rule's own flag, so the diff --git a/src/logscale.jl b/src/logscale.jl index b0eadd8..fa33c45 100644 --- a/src/logscale.jl +++ b/src/logscale.jl @@ -13,6 +13,8 @@ A log scale that is not known, with the reason: `cause` says why, `detail` what The causes an engine and the base package record: - `:no_declaration`: the rule that computed the message declares no `logscale`; `detail` is its [`RuleSpec`](@ref); +- `:improper`: the rule declares `logscale = `[`improper`](@ref): its message has no normalising + constant at all; `detail` is its [`RuleSpec`](@ref); - `:initial`: an initial message, not computed by a rule; - `:fallback`: a message computed by a rule fallback; - `:no_compute_logscale`: a product whose pair of distributions has no `compute_logscale` @@ -52,6 +54,8 @@ function describe_undefined(io::IO, logscale::UndefinedLogScale) cause, detail = logscale.cause, logscale.detail if cause === :no_declaration print(io, "the ", detail isa RuleSpec ? rule_heading(detail, io) : "rule", " declares no `logscale`") + elseif cause === :improper + print(io, "the ", detail isa RuleSpec ? rule_heading(detail, io) : "rule", " gives an improper message: it has no normalising constant") elseif cause === :initial print(io, "the message is an initial one, not computed by a rule") elseif cause === :fallback @@ -93,7 +97,7 @@ struct UndefinedLogScaleError <: Exception end function Base.showerror(io::IO, err::UndefinedLogScaleError) - print(io, "UndefinedLogScaleError: a log scale is needed but not known: ") + print(io, "UndefinedLogScaleError: a log scale is needed but ", err.logscale.cause === :improper ? "none exists: " : "not known: ") describe_undefined(io, err.logscale) return nothing end @@ -184,6 +188,26 @@ The `logscale` declaration of a rule whose body computes its log scale: written """ const from_body = FromBody() +""" + Improper + +The type of [`improper`](@ref), the marker of a rule whose message has no log scale to declare. +""" +struct Improper end + +""" + improper + +The `logscale` declaration of a rule whose message is improper: written `logscale = improper`, it +says that the message has no normalising constant, its integral being infinite, so no log scale +exists to compute. An exact belief-propagation message may be so: the message +`v ↦ N(y | μ, v)` towards a variance from an observed `y` decays only like `v^(-1/2)`. The +message's log scale is then an [`UndefinedLogScale`](@ref) whose cause is `:improper`, which +propagates as any undefined one does and which [`require_logscale`](@ref) reports as improper. A +rule that declares no `logscale` says something else: that its log scale is not known. +""" +const improper = Improper() + """ RuleLogScales(; m = NamedTuple()) @@ -209,12 +233,13 @@ and `nothing` for a marginal or an average energy. """ function getlogscale end -# A `logscale` declaration is a number, a function of the rule's inputs, `from_body`, or -# `nothing`: the rule declares none. -valid_logscale_declaration(declaration) = declaration === nothing || declaration isa Union{Real, Function, FromBody} +# A `logscale` declaration is a number, a function of the rule's inputs, `from_body`, `improper`, +# or `nothing`: the rule declares none. +valid_logscale_declaration(declaration) = declaration === nothing || declaration isa Union{Real, Function, FromBody, Improper} describe_logscale_declaration(::Nothing) = "none" describe_logscale_declaration(::FromBody) = "from the body" +describe_logscale_declaration(::Improper) = "none exists: the message is improper" describe_logscale_declaration(::Function) = "a function of the inputs" describe_logscale_declaration(declaration::Real) = string(declaration) @@ -228,6 +253,7 @@ logscale_constant(name, declaration::Function) = throw( @inline rule_logscale(spec, ::Nothing, raw, algorithm, ctx, args, target) = UndefinedLogScale(:no_declaration, spec) @inline rule_logscale(spec, declaration::Real, raw, algorithm, ctx, args, target) = declaration @inline rule_logscale(spec, declaration::FromBody, raw, algorithm, ctx, args, target) = raw.logscale +@inline rule_logscale(spec, ::Improper, raw, algorithm, ctx, args, target) = UndefinedLogScale(:improper, spec) @inline rule_logscale(spec, declaration::Function, raw, algorithm, ctx, args, target) = declaration(algorithm, ctx, args, target) # A rule declared `from_body` returns a `WithLogScale`; any other returns its result bare. diff --git a/src/registry.jl b/src/registry.jl index 5c00b78..58ff19a 100644 --- a/src/registry.jl +++ b/src/registry.jl @@ -16,7 +16,8 @@ precompile image carries its own entries. A redefinition, at the REPL say, repla with the same key. It is for introspection only: listing, checking and displaying rules, and suggesting candidates -when no rule fits. Resolution never reads it; it is Julia's dispatch. +when no rule fits. Resolution, finding which rule runs, does not use it: Julia's dispatch finds +the rule. See [What the registry is for](@ref inspecting-registry). See also [`registries`](@ref), [`registered_rules`](@ref). """ diff --git a/src/result_show.jl b/src/result_show.jl index c63eb49..0aa379f 100644 --- a/src/result_show.jl +++ b/src/result_show.jl @@ -136,6 +136,7 @@ end logscale_source(::Real) = "declared" logscale_source(::Function) = "computed from the inputs" logscale_source(::FromBody) = "computed by the body" +logscale_source(::Improper) = "none: the message is improper" logscale_source(::Nothing) = "not declared" function logscale_label(r::RuleResult) diff --git a/src/rule_macro.jl b/src/rule_macro.jl index 97dd68c..f3672ea 100644 --- a/src/rule_macro.jl +++ b/src/rule_macro.jl @@ -182,11 +182,14 @@ $(DOC_RULE_ALGORITHM) - 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`](@ref)`(result, logscale)`, for a log scale - computed alongside the result. + computed alongside the result; + - [`improper`](@ref): 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`](@ref) naming the rule, which propagates through products; only [`require_logscale`](@ref) turns it into an - error. + 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 @@ -450,6 +453,8 @@ function define_rule_expr(kind, source, macroargs) logscale_fn = :(($(ls_args...), $ls_target) -> $user_ls($(ls_passed...), $(ls_index...))) elseif declaration === :from_body logscale_fn = from_body + elseif declaration === :improper + logscale_fn = improper else # A number, checked by `RuleSpec`; a function written by its name is refused here, since # only a lambda tells which slots it takes. diff --git a/src/rulespec.jl b/src/rulespec.jl index 3fea3cf..96c597b 100644 --- a/src/rulespec.jl +++ b/src/rulespec.jl @@ -39,7 +39,7 @@ What it declares, read as fields: - `annotates`: whether its body takes the `ann` slot, and so may write annotations on its result; an engine gives a rule that does not an annotation store nobody writes to; - `logscale`: what a message rule declares about its result's log scale: `nothing` for none, a - number, a function, or [`from_body`](@ref); `reads_logscale`: whether it reads its inbound + number, a function, [`from_body`](@ref) or [`improper`](@ref); `reads_logscale`: whether it reads its inbound messages' log scales; - `args_check`: the source of the check its inputs must pass as its body starts, or `nothing` for none (the definition macros' `args_check` keyword; a failure is a [`RuleInputError`](@ref)); @@ -88,7 +88,7 @@ function RuleSpec(; inplace && prealloc === nothing && throw(ArgumentError("an in-place rule needs a `preallocate` function")) valid_logscale_declaration(logscale) || - throw(ArgumentError("a rule's `logscale` is a number, a function of its inputs or `from_body`, got $(repr(logscale))")) + throw(ArgumentError("a rule's `logscale` is a number, a function of its inputs, `from_body` or `improper`, got $(repr(logscale))")) kind === :message || (logscale === nothing && !reads_logscale) || throw(ArgumentError("only a message rule has a log scale; `logscale` and `reads_logscale` are for message rules")) effective = something(pure, algorithm <: AbstractAlgorithm ? ispure(algorithm) : true) diff --git a/test/logscale_tests.jl b/test/logscale_tests.jl index 966e490..50ce0a1 100644 --- a/test/logscale_tests.jl +++ b/test/logscale_tests.jl @@ -17,6 +17,14 @@ ) # None declared. @define_message_update_rule(node = Scale, target = :out, args = (m[:in]::Int,), body = (args) -> 2 * args.m[:in]) + + # An improper message: the likelihood of a variance, whose integral is infinite. + struct Noise end # out ~ N(0, v) + @define_factor_node(node = Noise, type = Stochastic, interfaces = [:out, :v]) + @define_message_update_rule( + node = Noise, target = :v, args = (m[:out]::Real,), logscale = improper, + body = (args) -> (v -> -(log(2π * v) + args.m[:out]^2 / v) / 2), + ) end @testitem "logscale:declarations" tags = [:base] setup = [LogScaleRules] begin @@ -86,6 +94,28 @@ end @test with_logscale(1, 2) === with_logscale(result = 1, logscale = 2) end +@testitem "logscale:improper" tags = [:base] setup = [LogScaleRules] begin + using MessagePassingRulesBase + L = LogScaleRules + + result = call_message_update_rule(L.Noise, :v; m = (out = 1.0,)) + @test getresult(result)(1.0) ≈ -(log(2π) + 1) / 2 + logscale = getlogscale(result) + @test logscale isa UndefinedLogScale && !isdefined_logscale(logscale) + @test logscale.cause === :improper && logscale.detail === getrule(result) + @test getrule(result).logscale === improper + # It propagates as an undefined log scale does, and says that none exists where one is needed. + @test logscale + 1.0 === logscale + @test_throws UndefinedLogScaleError require_logscale(logscale) + text = sprint(showerror, UndefinedLogScaleError(logscale)) + @test contains(text, "a log scale is needed but none exists") + @test contains(text, "Noise towards :v under") && contains(text, "gives an improper message: it has no normalising constant") + @test contains(repr(MIME"text/plain"(), result), "gives an improper message") + @test contains(repr(MIME"text/plain"(), getrule(result)), "logscale: none exists: the message is improper") + # A declaration not to know: an omitted `logscale` still reads as not declared. + @test getlogscale(call_message_update_rule(L.Scale, :out; m = (in = 1,))).cause === :no_declaration +end + @testitem "logscale:display" tags = [:base] setup = [LogScaleRules] begin using MessagePassingRulesBase L = LogScaleRules @@ -121,7 +151,7 @@ end call_message_update_rule(N, :out; m = (in = 1.0,)) ) - @test failure(rule(:(logscale = "zero"), :(body = (args) -> 1))) |> msg -> contains(msg, "a number, a function of its inputs or `from_body`") + @test failure(rule(:(logscale = "zero"), :(body = (args) -> 1))) |> msg -> contains(msg, "a number, a function of its inputs, `from_body` or `improper`") @test failure(rule(:(reads_logscale = 1), :(body = (args) -> 1))) |> msg -> contains(msg, "`reads_logscale` must be `true` or `false`") @test failure(rule(:(logscale = (ann) -> 0), :(body = (args) -> 1))) |> msg -> contains(msg, "unknown logscale slot") @test failure(rule(:(logscale = sin), :(body = (args) -> 1))) |> msg -> contains(msg, "write it as a lambda") From 94f87968a1baaf868ec2d29a74c58023f8eef2e1 Mon Sep 17 00:00:00 2001 From: Bagaev Dmitry Date: Tue, 6 Oct 2026 10:11:03 +0200 Subject: [PATCH 2/2] Drop the unreachable display of an improper log scale's source A RuleResult shows a log scale's source only when the log scale is a number, and an improper rule's is always an UndefinedLogScale, so the method never ran (Codecov's patch check). Co-Authored-By: Claude Opus 5.5 (1M context) --- src/result_show.jl | 1 - 1 file changed, 1 deletion(-) diff --git a/src/result_show.jl b/src/result_show.jl index 0aa379f..c63eb49 100644 --- a/src/result_show.jl +++ b/src/result_show.jl @@ -136,7 +136,6 @@ end logscale_source(::Real) = "declared" logscale_source(::Function) = "computed from the inputs" logscale_source(::FromBody) = "computed by the body" -logscale_source(::Improper) = "none: the message is improper" logscale_source(::Nothing) = "not declared" function logscale_label(r::RuleResult)