Steering a Process

Some processes are random without a sample site: a language model writing a sentence, a simulator taking a step. We cannot score their draws, but we can still steer them — run many copies, and keep the ones a reward favours. foerster.steer turns such a process into a program for SMC whose target is the process's own law tilted by the reward, p(trajectory) · exp(reward). foerster.learn reads the trajectories back as training data.

(ns foerster.steering
  (:require [org.replikativ.foerster.core :as infer]
            [org.replikativ.foerster.learn :as learn]
            [org.replikativ.foerster.measure :as m]
            [org.replikativ.foerster.random :as random]
            [org.replikativ.foerster.steer :as steer]
            [org.replikativ.spindel.core :as sp]
            [org.replikativ.spindel.spin.cps :refer [spin]]
            [scicloj.kindly.v4.kind :as kind]))
(def world (sp/create-execution-context))
(defn run
  "Run the inference `(make)` returns, seeded."
  [seed make]
  (random/set-seed! seed)
  (sp/with-context world @(make)))
(defn mean
  "Posterior mean of `f` of the program's value."
  [measure f]
  (let [ps (m/get-particles measure)
        ws (m/normalize-log-weights (mapv second ps))]
    (reduce + (map (fn [[particle _] w] (* w (double (f (m/get-value particle))))) ps ws))))

A process and a reward

The process tosses a fair coin five times and counts the heads. The coin is drawn with random/uniform01 — no sample site, as a language model's next token would be. The reward λ·heads tilts the law towards heads: the target is Binomial(5, ½) · e^{λ·heads}, which is Binomial(5, p) with p = e^λ / (1 + e^λ), and its evidence is ((1 + e^λ)/2)⁵.

(def lambda 0.8)
(def steps 5)
(defn coin-step [heads]
  (spin (+ heads (if (< (random/uniform01) 0.5) 1 0))))
(def exact
  (let [p (/ (Math/exp lambda) (+ 1 (Math/exp lambda)))]
    {:mean (* steps p)
     :log-evidence (* steps (Math/log (/ (+ 1 (Math/exp lambda)) 2)))}))
exact
{:mean 3.4498724056380627, :log-evidence 2.3897674269391627}

steer/model takes the initial state, the :step to the next state, the :reward at the end and :done?. With a :value — an estimate of the reward to come, here the reward so far — SMC resamples at every step on the change of the estimate (twisted SMC); without one, it weighs whole trajectories by their reward at the end (best-of-N).

(defn process [twisted?]
  (steer/model (cond-> {:init 0
                        :step coin-step
                        :done? (constantly false)
                        :max-steps steps
                        :reward (fn [heads] (* lambda heads))}
                 twisted? (assoc :value (fn [heads] (* lambda heads))))))
(def steered
  (array-map
   "twisted" (run 41 #(infer/smc-infer (process true) 1000 {:resampling :stratified}))
   "best-of-N" (run 41 #(infer/smc-infer (process false) 1000 {:resampling :stratified}))))
(kind/table {:column-names ["steering" "mean heads" "log evidence"]
             :row-vectors (into [["exact" (:mean exact) (:log-evidence exact)]]
                                (for [[label measure] steered]
                                  [label (mean measure identity) (m/log-marginal measure)]))})
steeringmean headslog evidence
exact3.44987240563806272.3897674269391627
twisted3.40471391401774742.3715966606088656
best-of-N3.41705506668343162.3677687367670615

Both reach the tilted law and its evidence. The twist pays off when the value estimate is informative and trajectories are long: SMC drops the unpromising ones early instead of finishing them.

A biased proposal, weighted back

A step may come from another process than the one we target — a learned proposal, a smaller model. It then returns (steer/weighted state log-w) with log-w = log p(state) − log q(state), and the weight restores the target. Here the steps come from a coin with heads 0.8:

(defn biased-step [heads]
  (spin (if (< (random/uniform01) 0.8)
          (steer/weighted (inc heads) (Math/log (/ 0.5 0.8)))
          (steer/weighted heads (Math/log (/ 0.5 0.2))))))
(def biased
  (run 43 #(infer/smc-infer (steer/model {:init 0
                                          :step biased-step
                                          :done? (constantly false)
                                          :max-steps steps
                                          :value (fn [heads] (* lambda heads))
                                          :reward (fn [heads] (* lambda heads))})
                            1000 {:resampling :stratified})))
{:mean-heads (mean biased identity)
 :log-evidence (m/log-marginal biased)}
{:mean-heads 3.45921686640455, :log-evidence 2.380572732389737}

The same target, from draws of a different coin.

Trajectories as training data

Every step records its state under [:steer/state t] and the end its reward under :steer/reward, so each particle carries its trajectory. learn/trajectories reads them back with their normalized weights — the data for learning a value estimate or a proposal that makes the next search cheaper:

(def trajectories (learn/trajectories (get steered "best-of-N")))
(kind/table {:column-names ["states" "reward" "weight"]
             :row-vectors (mapv (juxt :states :reward :weight) (take 5 trajectories))})
statesrewardweight
[1 2 2 3 4]
3.20.0022984414509834867
[0 1 2 2 2]
1.64.6404732576814936E-4
[1 2 2 2 2]
1.64.6404732576814936E-4
[0 1 1 2 2]
1.64.6404732576814936E-4
[1 1 1 2 2]
1.64.6404732576814936E-4

A learner that takes plain, unweighted examples wants draws from the target instead: learn/draws resamples the trajectories by weight.

(def draws (learn/draws (get steered "best-of-N") 2000))
{:mean-final-heads (/ (reduce + (map (comp peek :states) draws)) (double (count draws)))
 :exact (:mean exact)}
{:mean-final-heads 3.453, :exact 3.4498724056380627}