// ebm.mbt -- Energy-Based Model composite (v0.119.0).
//
// An Energy-Based Model (EBM) is parameterised by an energy function
// E_theta(x) and a sampling procedure (here Langevin dynamics). The
// composite ties the two together for forward-time use.
//
// Scope of v0.119.0:
// - EBM struct (EnergyFunction + step_size + fd_eps).
// - ebm_energy: convenience wrapper around energy_function_predict.
// - ebm_score: input gradient via finite differences (== score).
// - ebm_sample: a single Langevin chain (init -> n_steps updates).
// - ebm_sample_batch: independent chains started from `init_batch`.
//
// Reference: LeCun et al. 2006; Du & Mordatch 2019 (implicit
// generation with EBMs).
///|
/// EBM: pairs an EnergyFunction with the sampler hyperparameters.
pub struct EBM {
net : EnergyFunction
step_size : Float
fd_eps : Float
}
///|
/// Build a fresh EBM.
pub fn EBM::new(
net : EnergyFunction,
step_size : Float,
fd_eps : Float,
) -> EBM {
{ net, step_size, fd_eps }
}
///|
/// Convenience: returns the scalar energy of a sample.
pub fn ebm_energy(ebm : EBM, x : Array[Float]) -> Float {
energy_function_predict(ebm.net, x)
}
///|
/// Convenience: returns the score -dE/dx (a vector of length
/// `x.length()`). Approximated via central finite differences.
pub fn ebm_score(ebm : EBM, x : Array[Float]) -> Array[Float] {
let g = langevin_gradient(ebm.net, x, ebm.fd_eps)
// Score is -grad E.
let n = x.length()
let out : Array[Float] = Array::make(n, 0.0F)
for i in 0.. Array[Float] {
langevin_chain(ebm.net, init, n_steps, ebm.step_size, ebm.fd_eps, rng)
}
///|
/// Sample `batch_size` independent chains starting from the rows of
/// `init_batch`. Returns an array of arrays (one per chain).
pub fn ebm_sample_batch(
ebm : EBM,
init_batch : Array[Array[Float]],
n_steps : Int,
rng : Xoshiro,
) -> Array[Array[Float]] {
let n = init_batch.length()
let out : Array[Array[Float]] = Array::make(n, Array::make(0, 0.0F))
for c in 0..