Skip to content

Probabilistic Programming ​

Reactant.jl compiles ordinary Julia code through MLIR, running it on CPU, GPU, or TPU without rewriting. You write a normal Julia function, wrap a call to it in @compile, and Reactant traces the function, lowers it through MLIR, and hands you back a callable compiled program. Array inputs are staged through ConcreteRArray (created with Reactant.to_rarray) so they live on the target device.

Reactant.ProbProg is the Julia front-end for the impulse dialect, implemented across Enzyme (dialect definition and inference materialization passes) and Enzyme-JAX (backend-specific lowering). The impulse dialect provides high-level MLIR ops for describing probabilistic modeling and inference, materializes inference computation through compiler passes, and applies general-purpose and probabilistic-programming-specific optimizations during lowering.

optimize = :probprog opt-in required

For now, @compile needs an explicit optimize = :probprog argument on probabilistic programs to enable the impulse-specific MLIR passes (you'll see this in every @compile call below). Merging those passes into the default @compile pipeline is work in progress; once it lands, the explicit opt-in will no longer be required.

Next, we walk through two operating modes of Reactant.ProbProg: a trace-based mode built around a generative function, and a custom log-density mode that takes a custom log-density function.

Trace-based mode ​

We describe a Bayesian linear regression question:

slope∼N(0,2)intercept∼N(0,10)yi∣slope,intercept∼N(slope⋅xi+intercept,1)

Both regression coefficients are given Gaussian priors, tighter on slope (standard deviation 2) and looser on intercept (standard deviation 10). Each observation y_i is then drawn from a Gaussian centered on the fitted value slope · x_i + intercept with fixed noise (standard deviation 1).

The data ​

We start by drawing synthetic data using a known slope and intercept of -2 and 10 respectively.

julia
true_slope, true_intercept = -2.0, 10.0

xs = collect(Float64, 1:10)
ys = (true_slope .* xs) .+ true_intercept .+ randn(length(xs))
(xs, ys)
([1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0, 8.0, 9.0, 10.0], [7.684507339560489, 5.771278937873204, 3.9842073592681, 2.4379120343492175, -1.056488593162767, -1.8620400010248728, -2.4374529385010475, -6.602463585066136, -8.537803479443754, -11.275209899547658])

Describing the Model ​

We describe this model in Reactant.ProbProg as follows:

julia
using Reactant: ProbProg

function model(rng, xs)
    _, slope = ProbProg.sample(
        rng, ProbProg.Normal(0.0, 2.0, (1,)); symbol=:slope,
    )
    _, intercept = ProbProg.sample(
        rng, ProbProg.Normal(0.0, 10.0, (1,)); symbol=:intercept,
    )
    _, ys = ProbProg.sample(
        rng,
        ProbProg.Normal(slope .* xs .+ intercept, 1.0, (length(xs),));
        symbol=:ys,
    )
    return ys
end
model (generic function with 1 method)

Each random choice is introduced by a ProbProg.sample(rng, dist; symbol=...) call that takes a random number generator (RNG) and a distribution function. The symbol keyword names the sample site used for conditioning and specifying parameters to infer.

As a calling convention, ProbProg.sample returns (rng, value); the first element (omitted with _ above) is the updated RNG. In the current implementation rng is a traced ReactantRNG whose state corresponds to a tensor<2xui64> RNG state in the generated MLIR. We don't thread it through manually because Reactant tracing handles the input/output threading at the IR level, and ReactantRNG's internal state is updated via Julia mutability (see here for details).

Describing Inference ​

We condition on the observed ys with a Constraint object:

julia
obs = ProbProg.Constraint(:ys => ys)
Reactant.ProbProg.Constraint with 1 entry:
  Address([:ys]) => [7.68451, 5.77128, 3.98421, 2.43791, -1.05649, -1.86204, -2…

The current implementation requires a bit of boilerplate to flatten the Constraint into a tensor representation and to extract its address set before passing them to the @compile'd function below (see Traces and constrained inference for details):

julia
obs_tensor = ProbProg.flatten_constraint(obs)
1×10 ConcretePJRTArray{Float64,2}:
 7.68451  5.77128  3.98421  2.43791  -1.05649  …  -6.60246  -8.5378  -11.2752
julia
constrained_addresses = ProbProg.extract_addresses(obs)
Set{Reactant.ProbProg.Address} with 1 element:
  Reactant.ProbProg.Address([:ys])

We then specify what parameters to infer:

julia
selection = ProbProg.select(
    ProbProg.Address(:slope),
    ProbProg.Address(:intercept),
)
OrderedCollections.OrderedSet{Reactant.ProbProg.Address} with 2 elements:
  Reactant.ProbProg.Address([:slope])
  Reactant.ProbProg.Address([:intercept])

We express inference in a single function that conditions the model on the constraint with generate and then runs NUTS over the selected sites with mcmc.

julia
function infer(rng, xs, obs_tensor, step_size, inverse_mass_matrix)
    trace, = ProbProg.generate(
        rng, obs_tensor, model, xs; constrained_addresses,
    )
    trace, = ProbProg.mcmc(
        rng, trace, model, xs;
        selection, algorithm=:NUTS,
        step_size, inverse_mass_matrix,
        num_warmup=200, num_samples=500,
    )
    return trace
end
infer (generic function with 1 method)

The returned trace contains the sampling result as a 2D tensor: each row is the concatenation of all selected sites' flattened values for one post-warmup sample. (We will show a possible trace for this example problem below.)

Compiling with @compile ​

We compile infer with Reactant's @compile for compiler-optimized probabilistic inference:

julia
rng                 = ReactantRNG()
step_size           = Reactant.ConcreteRNumber(0.1)
inverse_mass_matrix = Reactant.ConcreteRArray([1.0 0.0; 0.0 1.0])

compiled_fn = @compile optimize=:probprog infer(
    rng, xs, obs_tensor, step_size, inverse_mass_matrix,
)
Reactant compiled function infer (with tag ##infer_reactant#1121)

Defaults

It is often sufficient to start with step_size = 1.0 and an identity inverse_mass_matrix. With the default adapt_step_size = true and adapt_mass_matrix = true, mcmc adaptively selects appropriate values during the warmup iterations.

The compiled_fn is a callable object that takes the same arguments as infer and returns the inference result. We can execute the compiled inference program any number of times by calling it:

julia
trace_tensor = compiled_fn(rng, xs, obs_tensor, step_size, inverse_mass_matrix)
500×2 ConcretePJRTArray{Float64,2}:
 -2.01768   9.13982
 -1.96376   9.78679
 -1.97303   9.80052
 -2.11513  10.7676
 -2.11833  10.2808
 -2.07205  10.2312
 -2.22261  10.7209
 -1.95495   9.7324
 -2.06577  10.0622
 -2.04596  10.373
  ⋮        
 -2.03448  10.2303
 -2.14592  11.0606
 -2.16216  10.6024
 -2.13906  10.1653
 -2.08842   9.44954
 -1.91963   9.52156
 -2.1129   10.003
 -2.299    11.5197
 -1.76598   8.41758

In this array, each row is one post-warmup sample, and each column is a parameter. The columns follow the order that we passed to ProbProg.select above, i.e., :slope first, then :intercept.

text
           :slope   :intercept
sample 1:    ...       ...    
sample 2:    ...       ...    
   ⋮          ⋮         ⋮     
sample N:    ...       ...

The posterior means obtained are:

julia
mean(trace_tensor, dims=1)
1×2 ConcretePJRTArray{Float64,2}:
 -2.04755  10.0591

We can see that NUTS recovers the values that were used to generate the data.

Custom logpdf mode ​

In larger applications, it is often infeasible to express the model in a PPL modeling language as we showed in the trace-based mode above. We can use Reactant.ProbProg to compile and run its inference algorithms directly on a hand-written log-density function via the custom logpdf mode.

For example, we can write the log-density function of the previous Bayesian linear regression model directly:

julia
function logdensity(θ, xs, ys)
    X = hcat(xs, ones(length(xs)))
    residuals = ys .- X * θ
    pr = -0.5 * sum(θ .^ 2 ./ [4.0, 100.0])
    ll = -0.5 * sum(residuals .^ 2)
    return ll + pr
end
logdensity (generic function with 1 method)

We pass logdensity to the mcmc_logpdf interface along with an initial position vector (the parameter values the chain starts from):

julia
function infer_logpdf(rng, θ0, xs, ys, step_size, inverse_mass_matrix)
    trace, = ProbProg.mcmc_logpdf(rng, logdensity, θ0, xs, ys;
        algorithm=:NUTS,
        step_size, inverse_mass_matrix,
        num_warmup=200, num_samples=500,
    )
    return trace
end

θ0 = Reactant.to_rarray(reshape([0.0, 0.0], 1, 2))
compiled_logpdf = @compile optimize=:probprog infer_logpdf(
    rng, θ0, xs, ys, step_size, inverse_mass_matrix,
)
trace = compiled_logpdf(rng, θ0, xs, ys, step_size, inverse_mass_matrix)
500×2 ConcretePJRTArray{Float64,2}:
 -2.04832  10.1233
 -2.1537   10.6291
 -2.17585  10.4862
 -2.12881  10.2779
 -2.22419  11.053
 -2.03044   9.70694
 -2.0363   10.3512
 -1.99145   9.56144
 -2.04283  10.4398
 -2.25429  11.6513
  ⋮        
 -2.19348  10.7226
 -2.00955  10.3428
 -1.88306   9.15419
 -2.15422  10.8365
 -2.15831  11.114
 -2.05893   9.78873
 -1.95016   8.98744
 -2.2208   11.2409
 -2.09346  10.049

We get similar inference results

julia
(
    posterior_mean_slope     = mean(trace[:, 1]),
    posterior_mean_intercept = mean(trace[:, 2]),
)
(posterior_mean_slope = -2.043721628421829, posterior_mean_intercept = 10.053586401159023)

NUTS recovers both posterior means here too.

More Explanations ​