diff --git a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/RNG.scala b/rainier-base/src/main/scala/com/stripe/rainier/RNG.scala similarity index 82% rename from rainier-sampler/src/main/scala/com/stripe/rainier/sampler/RNG.scala rename to rainier-base/src/main/scala/com/stripe/rainier/RNG.scala index c732441ce..04c0453e8 100644 --- a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/RNG.scala +++ b/rainier-base/src/main/scala/com/stripe/rainier/RNG.scala @@ -1,7 +1,6 @@ -package com.stripe.rainier.sampler +package com.stripe.rainier import scala.util.Random -import Log._ trait RNG { def standardUniform: Double @@ -18,8 +17,6 @@ object RNG { } final case class ScalaRNG(seed: Long) extends RNG { - FINE.log("Initializing RNG with seed %d", seed) - val rand: Random = new Random(seed) def standardUniform: Double = rand.nextDouble def standardNormal: Double = rand.nextGaussian diff --git a/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/sbc/SBCBenchmark.scala b/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/sbc/SBCBenchmark.scala index 3b9e6c8ac..4f8e0e115 100644 --- a/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/sbc/SBCBenchmark.scala +++ b/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/sbc/SBCBenchmark.scala @@ -1,11 +1,12 @@ package com.stripe.rainier.bench.sbc +import com.stripe.rainier.RNG import org.openjdk.jmh.annotations._ -import java.util.concurrent.TimeUnit +import java.util.concurrent.TimeUnit import com.stripe.rainier.compute._ import com.stripe.rainier.core._ -import com.stripe.rainier.sampler.{RNG, DensityFunction} +import com.stripe.rainier.sampler.DensityFunction @BenchmarkMode(Array(Mode.SampleTime)) @OutputTimeUnit(TimeUnit.MICROSECONDS) @@ -35,7 +36,7 @@ abstract class SBCBenchmark { @Benchmark def run(): Unit = - df.update(params) + df.update(RNG.default, params) } class NormalBenchmark extends SBCBenchmark { @@ -91,6 +92,7 @@ class BinomialPoissonApproximationBenchmark extends SBCBenchmark { SBC(Uniform(0, 0.04))((x: Real) => Binomial(x, 200)) } +/* class GaussianMixtureBenchmark extends SBCBenchmark { def sbc = SBC(Uniform(0, 1))( @@ -102,3 +104,5 @@ class GaussianMixtureBenchmark extends SBCBenchmark { ) )) } + + */ diff --git a/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/ARK.scala b/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/ARK.scala index 025c688b9..47d86f90e 100644 --- a/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/ARK.scala +++ b/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/ARK.scala @@ -1,5 +1,7 @@ package com.stripe.rainier.bench.stan +import com.stripe.rainier.compute.Real +import com.stripe.rainier.compute.Real.parameter import com.stripe.rainier.core._ //https://github.com/stan-dev/stat_comp_benchmarks/tree/master/benchmarks/arK @@ -12,7 +14,7 @@ class ARK extends ModelBenchmark { val betas = Normal(0, 10).latentVec(5) 5.until(ys.size).foldLeft(Model.empty) { (m, t) => - val mu = 1.to(5).foldLeft(alpha) { + val mu = 1.to(5).foldLeft(alpha: Real) { case (mu, k) => mu + betas(k - 1) * ys(t - k) } diff --git a/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/LowDimGaussMix.scala b/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/LowDimGaussMix.scala index 4fb148f81..f3855a94b 100644 --- a/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/LowDimGaussMix.scala +++ b/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/LowDimGaussMix.scala @@ -6,6 +6,7 @@ import com.stripe.rainier.core._ //https://github.com/stan-dev/stat_comp_benchmarks/tree/master/benchmarks/low_dim_gauss_mix //Stan: Gradient evaluation took 0.000292 seconds //JMH: 649.555 ± 2.572 us/op +/* class LowDimGaussMix extends ModelBenchmark { def model = { val (mu1, sigma1) = muSigma() @@ -275,3 +276,4 @@ class LowDimGaussMix extends ModelBenchmark { 1.0084427471213, -1.6306611574291, -4.35233903314464, 3.40936081830715, 2.75002324260943, 0.760843596839809) } + */ diff --git a/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/ModelBenchmark.scala b/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/ModelBenchmark.scala index f032df2e4..65da3fec6 100644 --- a/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/ModelBenchmark.scala +++ b/rainier-benchmark/src/main/scala/com/stripe/rainier/bench/stan/ModelBenchmark.scala @@ -1,10 +1,11 @@ package com.stripe.rainier.bench.stan +import com.stripe.rainier.RNG import org.openjdk.jmh.annotations._ -import java.util.concurrent.TimeUnit +import java.util.concurrent.TimeUnit import com.stripe.rainier.core._ -import com.stripe.rainier.sampler.{RNG, DensityFunction} +import com.stripe.rainier.sampler.DensityFunction @BenchmarkMode(Array(Mode.SampleTime)) @OutputTimeUnit(TimeUnit.SECONDS) @@ -28,7 +29,7 @@ abstract class ModelBenchmark { @Benchmark def run(): Unit = - df.update(params) + df.update(RNG.default, params) @Benchmark def build(): Unit = @@ -46,7 +47,7 @@ object ModelBenchmarks { def main(args: Array[String]): Unit = { new ARK().main() new EightSchools().main() - new LowDimGaussMix().main() +// new LowDimGaussMix().main() new GLMMPoisson2().main() } } diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Compiler.scala b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Compiler.scala index e95439ffa..3118a06a0 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Compiler.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Compiler.scala @@ -1,14 +1,15 @@ package com.stripe.rainier.compute -import com.stripe.rainier.ir +import com.stripe.rainier.{RNG, ir} final case class Compiler(methodSizeLimit: Int, classSizeLimit: Int) { def compile(parameters: Seq[Parameter], - real: Real): Array[Double] => Double = { + real: Real): (RNG, Array[Double]) => Double = { val cf = compile(parameters.map(_.param), List(("base", real))) - return { array => - val globalBuf = new Array[Double](cf.numGlobals) - ir.CompiledFunction.output(cf, array, globalBuf, 0) + return { + case (rng, array) => + val globalBuf = new Array[Double](cf.numGlobals) + ir.CompiledFunction.output(cf, rng, array, globalBuf, 0) } } def compileTargets(group: TargetGroup): ir.DataFunction = { diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Evaluator.scala b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Evaluator.scala index bba3f2b6a..daab71b35 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Evaluator.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Evaluator.scala @@ -33,7 +33,8 @@ class Evaluator(var cache: Map[Real, Double]) { Math.pow(toDouble(base), toDouble(exponent)) case l: Lookup => toDouble(l.table(toDouble(l.index).toInt - l.low)) - case p: Parameter => sys.error(s"No value provided for $p") + case p: Parameter => sys.error(s"No value provided for $p") + case Latent(value, _) => toDouble(value) } def compare(x: Real, y: Real): Int = toDouble(x).compare(toDouble(y)) diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Gradient.scala b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Gradient.scala index 8056a31f8..497294d30 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Gradient.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Gradient.scala @@ -59,6 +59,10 @@ private object Gradient { //no gradient visit(c.left) visit(c.right) + + case Latent(value, _) => + diff(value).register(diff(real)) + visit(value) } } } diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/compute/PartialEvaluator.scala b/rainier-compute/src/main/scala/com/stripe/rainier/compute/PartialEvaluator.scala index 35fb8721c..f8c4ce9eb 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/compute/PartialEvaluator.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/compute/PartialEvaluator.scala @@ -80,6 +80,8 @@ class PartialEvaluator(var noChange: Set[Real], rowIndex: Int) { (l, false) case p: Parameter => (p, false) + case Latent(value, _) => + apply(value) } } diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Real.scala b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Real.scala index 0799e7f04..7c32da2bf 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Real.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Real.scala @@ -10,41 +10,54 @@ sealed trait Real { def bounds: Bounds def +(other: Real): Real = RealOps.add(this, other) + def *(other: Real): Real = RealOps.multiply(this, other) def unary_- : Real = this * (-1) + def -(other: Real): Real = this + (-other) + def /(other: Real): Real = RealOps.divide(this, other) def min(other: Real): Real = RealOps.min(this, other) + def max(other: Real): Real = RealOps.max(this, other) def pow(exponent: Real): Real = RealOps.pow(this, exponent) def exp: Real = RealOps.unary(this, ir.ExpOp) + def log: Real = RealOps.unary(this, ir.LogOp) def sin: Real = RealOps.unary(this, ir.SinOp) + def cos: Real = RealOps.unary(this, ir.CosOp) + def tan: Real = RealOps.unary(this, ir.TanOp) def asin: Real = RealOps.unary(this, ir.AsinOp) + def acos: Real = RealOps.unary(this, ir.AcosOp) + def atan: Real = RealOps.unary(this, ir.AtanOp) def sinh: Real = (this.exp - (-this).exp) / 2 + def cosh: Real = (this.exp + (-this).exp) / 2 + def tanh: Real = this.sinh / this.cosh def abs: Real = RealOps.unary(this, ir.AbsOp) def logit: Real = -((Real.one / this - 1).log) + def logistic: Real = Real.one / (Real.one + (-this).exp) } object Real { implicit def apply[N](value: N)(implicit toReal: ToReal[N]): Real = toReal(value) + def seq[A](as: Seq[A])(implicit toReal: ToReal[A]): Seq[Real] = as.map(toReal(_)) @@ -61,6 +74,7 @@ object Real { } def parameter(): Parameter = new Parameter(new Prior(Real.zero)) + def parameter(fn: Parameter => Real): Parameter = { val x = parameter() x.prior = new Prior(fn(x)) @@ -78,16 +92,21 @@ object Real { } def doubles(seq: Seq[Double]): Real = new Column(seq.toArray) + def longs(seq: Seq[Long]): Real = doubles(seq.map(_.toDouble)) def eq(left: Real, right: Real, ifTrue: Real, ifFalse: Real): Real = lookupCompare(left, right, ifFalse, ifTrue, ifFalse) + def lt(left: Real, right: Real, ifTrue: Real, ifFalse: Real): Real = lookupCompare(left, right, ifFalse, ifFalse, ifTrue) + def gt(left: Real, right: Real, ifTrue: Real, ifFalse: Real): Real = lookupCompare(left, right, ifTrue, ifFalse, ifFalse) + def lte(left: Real, right: Real, ifTrue: Real, ifFalse: Real): Real = lookupCompare(left, right, ifFalse, ifTrue, ifTrue) + def gte(left: Real, right: Real, ifTrue: Real, ifFalse: Real): Real = lookupCompare(left, right, ifTrue, ifTrue, ifFalse) @@ -110,22 +129,32 @@ object Real { sealed trait Constant extends Real { def isZero: Boolean = bounds.lower == 0.0 && bounds.upper == 0.0 + def isOne: Boolean = bounds.lower == 1.0 && bounds.upper == 1.0 + def isTwo: Boolean = bounds.lower == 2.0 && bounds.upper == 2.0 + def isPosInfinity: Boolean = bounds.lower.isPosInfinity && bounds.upper.isPosInfinity + def isNegInfinity: Boolean = bounds.lower.isNegInfinity && bounds.upper.isNegInfinity + def isPositive: Boolean = bounds.lower >= 0.0 def getDouble: Double + def map(fn: Double => Double): Constant + def mapWith(other: Constant)(fn: (Double, Double) => Double): Constant + def +(other: Constant): Constant = ConstantOps.add(this, other) + def *(other: Constant): Constant = ConstantOps.multiply(this, other) + def /(other: Constant): Constant = ConstantOps.divide(this, other) } @@ -142,8 +171,11 @@ object Constant { final private case class Scalar(value: Double) extends Constant { val bounds = Bounds(value, value) + def getDouble = value + def map(fn: Double => Double) = Scalar(fn(value)) + def mapWith(other: Constant)(fn: (Double, Double) => Double) = other match { case Scalar(v) => Scalar(fn(value, v)) @@ -158,8 +190,11 @@ final private[rainier] class Column(val values: Array[Double]) extends Constant { val param = new ir.Param val bounds = Bounds(values.min, values.max) + def getDouble = sys.error("Not a scalar") + def map(fn: Double => Double) = new Column(values.map(fn)) + def mapWith(other: Constant)(fn: (Double, Double) => Double) = other match { case Scalar(v) => @@ -186,6 +221,17 @@ final private[rainier] class Parameter(var prior: Prior) extends NonConstant { private[rainier] class Prior(val density: Real) +case class FnCall(className: String, + methodName: String, + args: List[Real] = List.empty) + extends NonConstant { + val bounds: Bounds = Bounds(Double.NegativeInfinity, Double.PositiveInfinity) +} + +case class Latent(value: Real, generator: Real) extends NonConstant { + val bounds = value.bounds +} + final private case class Unary(original: NonConstant, op: ir.UnaryOp) extends NonConstant { val bounds = op match { diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Target.scala b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Target.scala index 2bd881f75..38b6aa608 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Target.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Target.scala @@ -120,6 +120,9 @@ object TargetGroup { case l: Lookup => loop(l.index) l.table.foreach(loop) + // NB: this should only happen for guide in variational inference + case Latent(value, _) => + loop(value) } } @@ -193,8 +196,12 @@ object TargetGroup { if (indexState.hasParameter) state.nonlinearOp - else + else { state + } + // NB: for variational inference, should stop here + case Latent(value, _) => + loop(value) } seen += (r -> result) result diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Translator.scala b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Translator.scala index 548a62568..5243a8103 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/compute/Translator.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/compute/Translator.scala @@ -20,7 +20,10 @@ private class Translator { binaryExpr(toExpr(base), toExpr(exponent), PowOp) case Compare(left, right) => binaryExpr(toExpr(left), toExpr(right), CompareOp) - case l: Lookup => lookupExpr(l) + case l: Lookup => lookupExpr(l) + case Latent(value, _) => toExpr(value) + case FnCall(className, methodName, args) => + VarDef(Sym.freshSym, FnIR(className, methodName, args.map(toExpr))) } reals += r -> expr expr diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/ir/CompiledFunction.scala b/rainier-compute/src/main/scala/com/stripe/rainier/ir/CompiledFunction.scala index 86f1b418f..be8408a07 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/ir/CompiledFunction.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/ir/CompiledFunction.scala @@ -1,39 +1,50 @@ package com.stripe.rainier.ir import Log._ +import com.stripe.rainier.RNG trait CompiledFunction { def numInputs: Int def numGlobals: Int def numOutputs: Int - def output0(inputs: Array[Double], + def output0(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double - def output1(inputs: Array[Double], + def output1(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double - def output2(inputs: Array[Double], + def output2(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double - def output3(inputs: Array[Double], + def output3(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double - def output4(inputs: Array[Double], + def output4(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double - def output5(inputs: Array[Double], + def output5(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double - def output6(inputs: Array[Double], + def output6(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double - def output7(inputs: Array[Double], + def output7(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double - def output8(inputs: Array[Double], + def output8(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double - def output9(inputs: Array[Double], + def output9(rng: RNG, + inputs: Array[Double], globals: Array[Double], output: Int): Double } @@ -120,22 +131,23 @@ object CompiledFunction { } def output(cf: CompiledFunction, + rng: RNG, inputs: Array[Double], globals: Array[Double], index: Int): Double = { val i = index % 10 val j = index / 10 i match { - case 0 => cf.output0(inputs, globals, j) - case 1 => cf.output1(inputs, globals, j) - case 2 => cf.output2(inputs, globals, j) - case 3 => cf.output3(inputs, globals, j) - case 4 => cf.output4(inputs, globals, j) - case 5 => cf.output5(inputs, globals, j) - case 6 => cf.output6(inputs, globals, j) - case 7 => cf.output7(inputs, globals, j) - case 8 => cf.output8(inputs, globals, j) - case 9 => cf.output9(inputs, globals, j) + case 0 => cf.output0(rng, inputs, globals, j) + case 1 => cf.output1(rng, inputs, globals, j) + case 2 => cf.output2(rng, inputs, globals, j) + case 3 => cf.output3(rng, inputs, globals, j) + case 4 => cf.output4(rng, inputs, globals, j) + case 5 => cf.output5(rng, inputs, globals, j) + case 6 => cf.output6(rng, inputs, globals, j) + case 7 => cf.output7(rng, inputs, globals, j) + case 8 => cf.output8(rng, inputs, globals, j) + case 9 => cf.output9(rng, inputs, globals, j) } } diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/ir/DataFunction.scala b/rainier-compute/src/main/scala/com/stripe/rainier/ir/DataFunction.scala index eb70f80b2..52d1d1033 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/ir/DataFunction.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/ir/DataFunction.scala @@ -1,5 +1,7 @@ package com.stripe.rainier.ir +import com.stripe.rainier.RNG + /* Input layout: - numParamInputs param inputs @@ -29,7 +31,8 @@ case class DataFunction(cf: CompiledFunction, } require(outputStartIndices(data.size) == cf.numOutputs) - def apply(inputs: Array[Double], + def apply(rng: RNG, + inputs: Array[Double], globals: Array[Double], outputs: Array[Double]): Unit = { var k = 0 @@ -40,12 +43,13 @@ case class DataFunction(cf: CompiledFunction, var i = 0 while (i < data.size) { - compute(inputs, globals, outputs, i) + compute(rng, inputs, globals, outputs, i) i += 1 } } - private def compute(inputs: Array[Double], + private def compute(rng: RNG, + inputs: Array[Double], globals: Array[Double], outputs: Array[Double], i: Int): Unit = { @@ -64,6 +68,7 @@ case class DataFunction(cf: CompiledFunction, var o = 0 while (o < numOutputs) { outputs(o) += CompiledFunction.output(cf, + rng, inputs, globals, outputStartIndex + o) @@ -75,6 +80,7 @@ case class DataFunction(cf: CompiledFunction, var o = 0 while (o < numOutputs) { outputs(o) += CompiledFunction.output(cf, + rng, inputs, globals, outputStartIndex + o) diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/ir/ExprMethodGenerator.scala b/rainier-compute/src/main/scala/com/stripe/rainier/ir/ExprMethodGenerator.scala index 09ff6a747..57c6eda56 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/ir/ExprMethodGenerator.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/ir/ExprMethodGenerator.scala @@ -9,7 +9,7 @@ final private case class ExprMethodGenerator(method: MethodDef, val isStatic: Boolean = true val methodName: String = exprMethodName(method.sym.id) val className: String = classNameForMethod(classPrefix, method.sym.id) - val methodDesc: String = "([D[D)D" + val methodDesc: String = "(Lcom/stripe/rainier/RNG;[D[D)D" private val varIndices = inputs.zipWithIndex.toMap @@ -67,5 +67,9 @@ final private case class ExprMethodGenerator(method: MethodDef, traverse(s.second) case m: MethodRef => callExprMethod(classPrefix, m.sym.id) + case f: FnIR => + loadRNG() + f.args.foreach(traverse) + callFunction(f.className, f.methodName, f.args.size) } } diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/ir/IR.scala b/rainier-compute/src/main/scala/com/stripe/rainier/ir/IR.scala index 588301ef0..e9c9677df 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/ir/IR.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/ir/IR.scala @@ -21,6 +21,8 @@ final case class UnaryIR(original: Expr, op: UnaryOp) extends IR final case class LookupIR(index: Expr, table: List[Ref], low: Int) extends IR final case class MethodRef(sym: Sym) extends IR final case class SeqIR(first: VarDef, second: VarDef) extends IR +final case class FnIR(className: String, methodName: String, args: List[Expr]) + extends IR object SeqIR { def apply(seq: Seq[VarDef]): VarDef = diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/ir/MethodGenerator.scala b/rainier-compute/src/main/scala/com/stripe/rainier/ir/MethodGenerator.scala index 2cbac9027..9bcf4c300 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/ir/MethodGenerator.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/ir/MethodGenerator.scala @@ -103,12 +103,23 @@ private trait MethodGenerator { def exprMethodName(id: Int): String = s"_$id" def callExprMethod(classPrefix: String, id: Int): Unit = { + loadRNG() loadParams() loadGlobalVars() methodNode.visitMethodInsn(INVOKESTATIC, classNameForMethod(classPrefix, id), exprMethodName(id), - "([D[D)D", + "(Lcom/stripe/rainier/RNG;[D[D)D", + false) + } + + def callFunction(className: String, methodName: String, nArgs: Int): Unit = { + val typeDef = + s"(Lcom/stripe/rainier/RNG;${(0 until nArgs).map(_ => "D").mkString("")})D" + methodNode.visitMethodInsn(INVOKESTATIC, + className, + methodName, + typeDef, false) } @@ -169,27 +180,32 @@ private trait MethodGenerator { /** The local var layout is assumed to be: For static methods: - 0: params array - 1: globals array - 2..N: locally allocated doubles (two slots each) + 0: rng RNG + 1: params array + 2: globals array + 3..N: locally allocated doubles (two slots each) for output(): 0: this - 1: params array - 2: globals array - 3: output index + 1: rng RNG + 2: params array + 3: globals array + 4: output index **/ def loadParams(): Unit = - methodNode.visitVarInsn(ALOAD, if (isStatic) 0 else 1) + methodNode.visitVarInsn(ALOAD, if (isStatic) 1 else 2) def loadGlobalVars(): Unit = - methodNode.visitVarInsn(ALOAD, if (isStatic) 1 else 2) + methodNode.visitVarInsn(ALOAD, if (isStatic) 2 else 3) + + def loadRNG(): Unit = + methodNode.visitVarInsn(ALOAD, if (isStatic) 0 else 1) - private def localVarSlot(pos: Int) = 2 + (pos * 2) + private def localVarSlot(pos: Int) = 3 + (pos * 2) def loadThis(): Unit = methodNode.visitVarInsn(ALOAD, 0) def loadOutputIndex(): Unit = - methodNode.visitVarInsn(ILOAD, 3) + methodNode.visitVarInsn(ILOAD, 4) } diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/ir/OutputMethodGenerator.scala b/rainier-compute/src/main/scala/com/stripe/rainier/ir/OutputMethodGenerator.scala index 979d7a2fc..6614219dd 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/ir/OutputMethodGenerator.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/ir/OutputMethodGenerator.scala @@ -6,7 +6,7 @@ final private case class OutputMethodGenerator(methodNum: Int, extends MethodGenerator { val isStatic: Boolean = false val methodName: String = s"output$methodNum" - val methodDesc: String = "([D[DI)D" + val methodDesc: String = "(Lcom/stripe/rainier/RNG;[D[DI)D" if (outputIDs.isEmpty) constant(0.0) diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/ir/Packer.scala b/rainier-compute/src/main/scala/com/stripe/rainier/ir/Packer.scala index 6338d997f..4f9e5eb23 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/ir/Packer.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/ir/Packer.scala @@ -59,6 +59,24 @@ private class Packer(methodSizeLimit: Int) { (SeqIR(firstDef, secondDef), firstSize + secondSize + 1) } } + case f: FnIR => + def handleArgs(args: List[Expr]): TailRec[(List[Expr], Int)] = { + if (args.isEmpty) { + TailCalls.done((List.empty[Expr], 0)) + } else { + traverse(args.head, 0).flatMap { + case (exprDef, exprSize) => + handleArgs(args.tail).map { + case (tailDef, tailSize) => + (exprDef +: tailDef, exprSize + tailSize) + } + } + } + } + handleArgs(f.args).map { + case (argsDef, argsSize) => + (FnIR(f.className, f.methodName, argsDef), argsSize) + } case _: MethodRef => sys.error("there shouldn't be any method refs yet") } diff --git a/rainier-compute/src/main/scala/com/stripe/rainier/ir/VarType.scala b/rainier-compute/src/main/scala/com/stripe/rainier/ir/VarType.scala index 32d9b8cc9..12af82f29 100644 --- a/rainier-compute/src/main/scala/com/stripe/rainier/ir/VarType.scala +++ b/rainier-compute/src/main/scala/com/stripe/rainier/ir/VarType.scala @@ -86,6 +86,8 @@ private object VarTypes { case s: SeqIR => traverse(s.first) traverse(s.second) + case fn: FnIR => + fn.args.foreach(traverse) case _: MethodRef => () } diff --git a/rainier-core/src/main/scala/com/stripe/rainier/core/Continuous.scala b/rainier-core/src/main/scala/com/stripe/rainier/core/Continuous.scala index 6b8c6c02b..843b697d8 100644 --- a/rainier-core/src/main/scala/com/stripe/rainier/core/Continuous.scala +++ b/rainier-core/src/main/scala/com/stripe/rainier/core/Continuous.scala @@ -1,7 +1,8 @@ package com.stripe.rainier.core +import com.stripe.rainier.RNG import com.stripe.rainier.compute._ -import com.stripe.rainier.sampler.RNG + import scala.annotation.tailrec /** @@ -11,35 +12,45 @@ trait Continuous extends Distribution[Double] { private[rainier] val support: Support def logDensity(seq: Seq[Double]) = Vec.from(seq).map(logDensity).columnize + def logDensity(x: Real): Real def scale(a: Real): Continuous = Scale(a).transform(this) + def translate(b: Real): Continuous = Translate(b).transform(this) + def exp: Continuous = Exp.transform(this) - def latent: Real - def latentVec(k: Int) = Vec.from(List.fill(k)(latent)) + def latent: Latent + + def latentVec(k: Int) = Vec.from(List.fill(k)(latent: Real)) } /** * A Continuous Distribution that inherits its transforms from a Support object. */ private[rainier] trait StandardContinuous extends Continuous { - def latent: Real = { + def latent: Latent = { val x = Real.parameter { x => support.logJacobian(x) + logDensity(support.transform(x)) } - support.transform(x) + Latent(support.transform(x), generatorCall) } + + def generatorCall: FnCall } /** * Location-scale family distribution */ -trait LocationScaleFamily { self => +trait LocationScaleFamily { + self => def logDensity(x: Real): Real + def generate(r: RNG): Double + def generatorCall: FnCall + val standard: StandardContinuous = new StandardContinuous { val support: Support = UnboundedSupport @@ -47,8 +58,12 @@ trait LocationScaleFamily { self => Generator.from { (r, _) => generate(r) } - def logDensity(real: Real): Real = + + def logDensity(real: Real): Real = { self.logDensity(real) + } + + def generatorCall: FnCall = self.generatorCall } def apply(location: Real, scale: Real): Continuous = { @@ -63,7 +78,10 @@ trait LocationScaleFamily { self => object Normal extends LocationScaleFamily { def logDensity(x: Real): Real = ((x * x) / -2.0) - 0.5 * Real(2 * math.Pi).log + def generate(r: RNG): Double = r.standardNormal + + def generatorCall: FnCall = FnCall(Normal.getClass.getName, "generate") } /** @@ -72,8 +90,11 @@ object Normal extends LocationScaleFamily { object Cauchy extends LocationScaleFamily { def logDensity(x: Real): Real = (((x * x) + 1) * Math.PI).log * -1 + def generate(r: RNG): Double = r.standardNormal / r.standardNormal + + def generatorCall: FnCall = FnCall(Cauchy.getClass.getName, "generate") } /** @@ -82,10 +103,13 @@ object Cauchy extends LocationScaleFamily { object Laplace extends LocationScaleFamily { def logDensity(x: Real): Real = Real(0.5).log - x.abs + def generate(r: RNG): Double = { val u = r.standardUniform - 0.5 Math.signum(u) * -1 * Math.log(1 - (2 * Math.abs(u))) } + + def generatorCall: FnCall = FnCall(Laplace.getClass.getName, "generate") } /** @@ -103,7 +127,7 @@ object Gamma { def standard(shape: Real): StandardContinuous = { Bounds.check(shape, "k > 0")(_ >= 0.0) new StandardContinuous { - val support = BoundedBelowSupport(Real.zero) + private[rainier] val support = BoundedBelowSupport(Real.zero) def logDensity(real: Real): Real = Bounds.positive(real) { @@ -111,38 +135,45 @@ object Gamma { Combinatorics.gamma(shape) - real } + def generatorCall: FnCall = + FnCall(Gamma.getClass.getName, "generateValue", List(shape)) + def generator: Generator[Double] = Generator.require(Set(shape)) { (r, n) => val a = n.toDouble(shape) - if (a < 1) { - val u = r.standardUniform - generate(a + 1, r) * Math.pow(u, 1.0 / a) - } else - generate(a, r) + generateValue(r, a) } + } + } - @tailrec - private def generate(a: Double, r: RNG): Double = { - val d = a - 1.0 / 3.0 - val c = (1.0 / 3.0) / Math.sqrt(d) - - var x = r.standardNormal - var v = 1.0 + c * x - while (v <= 0) { - x = r.standardNormal - v = 1.0 + c * x - } + def generateValue(r: RNG, a: Double): Double = { + if (a < 1) { + val u = r.standardUniform + generate(a + 1, r) * Math.pow(u, 1.0 / a) + } else + generate(a, r) + } - val v3 = v * v * v - val u = r.standardUniform + @tailrec + def generate(a: Double, r: RNG): Double = { + val d = a - 1.0 / 3.0 + val c = (1.0 / 3.0) / Math.sqrt(d) - if ((u < 1 - 0.0331 * x * x * x * x) || - (Math.log(u) < 0.5 * x * x + d * (1 - v3 + Math.log(v3)))) - d * v3 - else - generate(a, r) - } + var x = r.standardNormal + var v = 1.0 + c * x + while (v <= 0) { + x = r.standardNormal + v = 1.0 + c * x } + + val v3 = v * v * v + val u = r.standardUniform + + if ((u < 1 - 0.0331 * x * x * x * x) || + (Math.log(u) < 0.5 * x * x + d * (1 - v3 + Math.log(v3)))) + d * v3 + else + generate(a, r) } } @@ -151,6 +182,7 @@ object Gamma { */ object Exponential { val standard: Continuous = Gamma.standard(1.0) + def apply(rate: Real): Continuous = { Bounds.check(rate, "λ >= 0")(_ >= 0.0) standard.scale(Real.one / rate) @@ -164,7 +196,7 @@ final case class Beta(a: Real, b: Real) extends StandardContinuous { Bounds.check(a, "α >= 0")(_ >= 0.0) Bounds.check(b, "β >= 0")(_ >= 0.0) - val support = new BoundedSupport(Real.zero, Real.one) + private[rainier] val support = BoundedSupport(Real.zero, Real.one) def logDensity(real: Real): Real = Bounds.zeroToOne(real)(betaDensity(real)) @@ -179,11 +211,21 @@ final case class Beta(a: Real, b: Real) extends StandardContinuous { (a - 1) * u.log + (b - 1) * (1 - u).log - Combinatorics.beta(a, b) + + def generatorCall: FnCall = + FnCall(Beta.getClass.getName, "generateValue", List(a, b)) } object Beta { + def generateValue(r: RNG, a: Double, b: Double): Double = { + val z1 = Gamma.generateValue(r, a) + val z2 = Gamma.generateValue(r, b) + z1 / (z1 + z2) + } + def meanAndPrecision(mean: Real, precision: Real): Beta = Beta(mean * precision, (Real.one - mean) * precision) + def meanAndVariance(mean: Real, variance: Real): Beta = meanAndPrecision(mean, mean * (Real.one - mean) / variance - 1) } @@ -202,19 +244,26 @@ object LogNormal { object Uniform { val beta11: Beta = Beta(1, 1) val standard: Continuous = new StandardContinuous { - val support = beta11.support + private[rainier] val support = beta11.support def logDensity(real: Real): Real = beta11.logDensity(real) + val generator: Generator[Double] = Generator.from { (r, _) => r.standardUniform } + + def generatorCall: FnCall = + FnCall(Uniform.getClass.getName, "generateValue") } + def generateValue(r: RNG): Double = r.standardUniform + def apply(from: Real, to: Real): Continuous = standard.scale(to - from).translate(from) } +/* case class Mixture(components: Map[Continuous, Real]) extends Continuous { components.values.foreach { r => Bounds.check(r, "0 <= p <= 1") { p => @@ -222,15 +271,15 @@ case class Mixture(components: Map[Continuous, Real]) extends Continuous { } } + private[rainier] val support = Support.union(components.keys.map { + _.support + }) + def generator: Generator[Double] = Generator.categorical(components).flatMap { d => d.generator } - val support = Support.union(components.keys.map { - _.support - }) - def logDensity(real: Real): Real = Real .logSumExp(components.map { @@ -239,10 +288,20 @@ case class Mixture(components: Map[Continuous, Real]) extends Continuous { } }) - def latent: Real = { + def latent: Latent = { val x = Real.parameter { x => support.logJacobian(x) + logDensity(support.transform(x)) } - support.transform(x) + val cdf = + components.toList + .scanLeft((Option.empty[Continuous], Real.zero)) { + case ((_, acc), (t, p)) => ((Some(t)), p + acc) + } + .collect { case (Some(t), p) => (t, p) } + + Latent(support.transform(x), FnCall()) } } + + + */ diff --git a/rainier-core/src/main/scala/com/stripe/rainier/core/Discrete.scala b/rainier-core/src/main/scala/com/stripe/rainier/core/Discrete.scala index 6fa43e0c4..c1589bd8c 100644 --- a/rainier-core/src/main/scala/com/stripe/rainier/core/Discrete.scala +++ b/rainier-core/src/main/scala/com/stripe/rainier/core/Discrete.scala @@ -1,7 +1,7 @@ package com.stripe.rainier.core +import com.stripe.rainier.RNG import com.stripe.rainier.compute._ -import com.stripe.rainier.sampler.RNG trait Discrete extends Distribution[Long] { self: Discrete => def logDensity(seq: Seq[Long]) = diff --git a/rainier-core/src/main/scala/com/stripe/rainier/core/Generator.scala b/rainier-core/src/main/scala/com/stripe/rainier/core/Generator.scala index 0664dc159..3e73a25f7 100644 --- a/rainier-core/src/main/scala/com/stripe/rainier/core/Generator.scala +++ b/rainier-core/src/main/scala/com/stripe/rainier/core/Generator.scala @@ -1,8 +1,8 @@ package com.stripe.rainier.core +import com.stripe.rainier.RNG import com.stripe.rainier.ir.CompiledFunction import com.stripe.rainier.compute._ -import com.stripe.rainier.sampler.RNG /** * Generator trait, for posterior predictive distributions to be forwards sampled during sampling @@ -79,7 +79,7 @@ sealed trait Generator[+T] { self => val globalBuf = new Array[Double](cf.numGlobals) val reqValues = new Array[Double](cf.numOutputs) 0.until(cf.numOutputs).foreach { i => - reqValues(i) = CompiledFunction.output(cf, array, globalBuf, i) + reqValues(i) = CompiledFunction.output(cf, r, array, globalBuf, i) } implicit val evaluator: Evaluator = new Evaluator( diff --git a/rainier-core/src/main/scala/com/stripe/rainier/core/Injection.scala b/rainier-core/src/main/scala/com/stripe/rainier/core/Injection.scala index b6b43c736..68e2a5e1f 100644 --- a/rainier-core/src/main/scala/com/stripe/rainier/core/Injection.scala +++ b/rainier-core/src/main/scala/com/stripe/rainier/core/Injection.scala @@ -37,7 +37,10 @@ private[rainier] trait Injection { self => } } - def latent: Real = forwards(dist.latent) + def latent: Latent = { + val upstream = dist.latent + Latent(forwards(upstream.value), forwards(upstream.generator)) + } } } diff --git a/rainier-core/src/main/scala/com/stripe/rainier/core/Model.scala b/rainier-core/src/main/scala/com/stripe/rainier/core/Model.scala index 7a3769325..5c5a4cddf 100644 --- a/rainier-core/src/main/scala/com/stripe/rainier/core/Model.scala +++ b/rainier-core/src/main/scala/com/stripe/rainier/core/Model.scala @@ -1,33 +1,55 @@ package com.stripe.rainier.core +import com.stripe.rainier.RNG import com.stripe.rainier.compute._ import com.stripe.rainier.sampler._ import com.stripe.rainier.optimizer._ -class Model(private[rainier] val likelihoods: List[Real], - val track: Set[Real]) { - def prior: Model = Model.track(track ++ likelihoods) - def merge(other: Model) = - new Model(likelihoods ++ other.likelihoods, track ++ other.track) +class MCMCInference(model: Model, + config: SamplerConfig = SamplerConfig.default, + nChains: Int = 4) { - def sample(config: SamplerConfig = SamplerConfig.default, nChains: Int = 4)( - implicit rng: RNG = RNG.default, - progress: Progress = SilentProgress): Trace = { + def sample(implicit rng: RNG = RNG.default, + progress: Progress = SilentProgress): Trace = { val results = 1 .to(nChains) .toList .map { i => - Driver.sample(i, config, density(), progress) + Driver.sample(i, config, model.density(), progress) } .toList - Trace(results.map(_._1), results.map(_._2), results.map(_._3), this) + Trace(results.map(_._1), results.map(_._2), results.map(_._3), model) } - def optimize[T, U](t: T)(implicit toGen: ToGenerator[T, U], - rng: RNG = RNG.default): U = { - val fn = toGen(t).prepare(parameters) - fn(Optimizer.lbfgs(density())) - } +} + +class VariationalInference(model: Model, + guide: Model, + mapping: Map[Latent, Latent]) { + /* + parameters are those that define the guide distributions + - (stochastic) latent variables are generated by guide + - gradients are stochastic - this is different from MCMC! + + a guide is also a model - without observations + or just a set of distributions? + + alternative variational: + - random variables as data - compiled in + - a few different instances - keep it stochastic + + stochastic gradients is not really hmc thing + - calling into generator might be expensive? + */ + def sample(implicit rng: RNG = RNG.default, + progress: Progress = SilentProgress): Unit = {} +} + +class Model(private[rainier] val likelihoods: List[Real], + val track: Set[Real]) { + def prior: Model = Model.track(track ++ likelihoods) + def merge(other: Model) = + new Model(likelihoods ++ other.likelihoods, track ++ other.track) lazy val targetGroup = TargetGroup(likelihoods, track) lazy val dataFn = @@ -41,13 +63,19 @@ class Model(private[rainier] val likelihoods: List[Real], val inputs = new Array[Double](dataFn.numInputs) val globals = new Array[Double](dataFn.numGlobals) val outputs = new Array[Double](dataFn.numOutputs) - def update(vars: Array[Double]): Unit = { + def update(rng: RNG, vars: Array[Double]): Unit = { System.arraycopy(vars, 0, inputs, 0, nVars) - dataFn(inputs, globals, outputs) + dataFn(rng, inputs, globals, outputs) } def density = outputs(0) def gradient(index: Int) = outputs(index + 1) } + + def optimize[T, U](t: T)(implicit toGen: ToGenerator[T, U], + rng: RNG = RNG.default): U = { + val fn = toGen(t).prepare(parameters) + fn(Optimizer.lbfgs(density())) + } } object Model { @@ -59,7 +87,8 @@ object Model { progress: Progress = SilentProgress): List[U] = { val gen = toGen(t) val model = Model.track(gen.requirements) - model.sample(config).predict(gen) + val inference = new MCMCInference(model, config) + inference.sample.predict(gen) } def apply[T, U](ts: T*)(implicit toGen: ToGenerator[T, U]): Model = diff --git a/rainier-core/src/main/scala/com/stripe/rainier/core/SBC.scala b/rainier-core/src/main/scala/com/stripe/rainier/core/SBC.scala index f94b27493..0fe6ed05d 100644 --- a/rainier-core/src/main/scala/com/stripe/rainier/core/SBC.scala +++ b/rainier-core/src/main/scala/com/stripe/rainier/core/SBC.scala @@ -1,5 +1,6 @@ package com.stripe.rainier.core +import com.stripe.rainier.RNG import com.stripe.rainier.sampler._ import com.stripe.rainier.compute._ @@ -108,8 +109,9 @@ final case class SBC[T](priors: Seq[Continuous], val (syntheticValues, trueOutput) = synthesize(syntheticSamples) val (model, real) = fit(syntheticValues) - val sample = - model.sample(samplerFn(Samples * thin / Chains), Chains).thin(thin) + val inference = + new MCMCInference(model, samplerFn(Samples * thin / Chains), Chains) + val sample = inference.sample.thin(thin) val diag = sample.diagnostics val maxRHat = diag.map(_.rHat).max val minEffectiveSampleSize = diag.map(_.effectiveSampleSize).min diff --git a/rainier-core/src/main/scala/com/stripe/rainier/core/Trace.scala b/rainier-core/src/main/scala/com/stripe/rainier/core/Trace.scala index 85bc0a248..205e35126 100644 --- a/rainier-core/src/main/scala/com/stripe/rainier/core/Trace.scala +++ b/rainier-core/src/main/scala/com/stripe/rainier/core/Trace.scala @@ -1,6 +1,8 @@ package com.stripe.rainier.core +import com.stripe.rainier.RNG import com.stripe.rainier.sampler._ + import scala.annotation.tailrec case class Trace(chains: List[List[Array[Double]]], diff --git a/rainier-sampler/src/main/scala/com/stripe/rainier/optimizer/Optimizer.scala b/rainier-sampler/src/main/scala/com/stripe/rainier/optimizer/Optimizer.scala index a090f980b..e0759e10b 100644 --- a/rainier-sampler/src/main/scala/com/stripe/rainier/optimizer/Optimizer.scala +++ b/rainier-sampler/src/main/scala/com/stripe/rainier/optimizer/Optimizer.scala @@ -1,5 +1,6 @@ package com.stripe.rainier.optimizer +import com.stripe.rainier.RNG import com.stripe.rainier.sampler.DensityFunction object Optimizer { @@ -12,7 +13,7 @@ object Optimizer { val eps = 0.1 val lb = new LBFGS(x, m, eps) while (!complete) { - df.update(x) + df.update(RNG.default, x) var i = 0 while (i < df.nVars) { g(i) = df.gradient(i) * -1 diff --git a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/DensityFunction.scala b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/DensityFunction.scala index 52b0ab5ac..fa54a5bd9 100644 --- a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/DensityFunction.scala +++ b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/DensityFunction.scala @@ -1,8 +1,10 @@ package com.stripe.rainier.sampler +import com.stripe.rainier.RNG + trait DensityFunction { def nVars: Int - def update(vars: Array[Double]): Unit + def update(rng: RNG, vars: Array[Double]): Unit def density: Double def gradient(index: Int): Double } diff --git a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Driver.scala b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Driver.scala index cf64db667..919fed80a 100644 --- a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Driver.scala +++ b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Driver.scala @@ -2,6 +2,7 @@ package com.stripe.rainier.sampler import scala.collection.mutable.ListBuffer import Log._ +import com.stripe.rainier.RNG object Driver { def sample(chain: Int, diff --git a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/EHMC.scala b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/EHMC.scala index 54df96448..c0fec6d26 100644 --- a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/EHMC.scala +++ b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/EHMC.scala @@ -1,5 +1,7 @@ package com.stripe.rainier.sampler +import com.stripe.rainier.RNG + class EHMCSampler(maxSteps: Int, minSteps: Int = 1, bufSize: Int = 100, diff --git a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/HMC.scala b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/HMC.scala index a29a753e2..d707e028d 100644 --- a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/HMC.scala +++ b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/HMC.scala @@ -1,5 +1,7 @@ package com.stripe.rainier.sampler +import com.stripe.rainier.RNG + class HMCSampler(nSteps: Int) extends Sampler { def initialize(params: Array[Double], lf: LeapFrog)(implicit rng: RNG) = () diff --git a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/LeapFrog.scala b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/LeapFrog.scala index 6ec2d6584..12dfdd173 100644 --- a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/LeapFrog.scala +++ b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/LeapFrog.scala @@ -1,5 +1,7 @@ package com.stripe.rainier.sampler +import com.stripe.rainier.RNG + final class LeapFrog(density: DensityFunction, statsWindow: Int) { var stats = new Stats(statsWindow) @@ -194,7 +196,7 @@ final class LeapFrog(density: DensityFunction, statsWindow: Int) { private def copyQsAndUpdateDensity(): Unit = { System.arraycopy(pqBuf, nVars, buf, 0, nVars) val t = System.nanoTime() - density.update(buf) + density.update(RNG.default, buf) stats.gradientTimes.add((System.nanoTime() - t).toDouble) stats.gradientEvaluations += 1 } diff --git a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Sampler.scala b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Sampler.scala index deaab2361..fb62dc0b0 100644 --- a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Sampler.scala +++ b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Sampler.scala @@ -1,5 +1,7 @@ package com.stripe.rainier.sampler +import com.stripe.rainier.RNG + trait SamplerConfig { def iterations: Int def warmupIterations: Int diff --git a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Stats.scala b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Stats.scala index 2a931e7d4..ea6274462 100644 --- a/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Stats.scala +++ b/rainier-sampler/src/main/scala/com/stripe/rainier/sampler/Stats.scala @@ -1,5 +1,7 @@ package com.stripe.rainier.sampler +import com.stripe.rainier.RNG + class Stats(n: Int) { var gradientEvaluations = 0L var iterations = 0 diff --git a/rainier-test/src/main/scala/com/stripe/rainier/core/SBCModel.scala b/rainier-test/src/main/scala/com/stripe/rainier/core/SBCModel.scala index d93bcfbc9..0bd2557c5 100644 --- a/rainier-test/src/main/scala/com/stripe/rainier/core/SBCModel.scala +++ b/rainier-test/src/main/scala/com/stripe/rainier/core/SBCModel.scala @@ -1,5 +1,6 @@ package com.stripe.rainier.core +import com.stripe.rainier.{RNG, ScalaRNG} import com.stripe.rainier.compute._ import com.stripe.rainier.sampler._ @@ -33,8 +34,8 @@ trait SBCModel[T] { implicit val rng: RNG = ScalaRNG(1528673302081L) val (values, trueValue) = sbc.synthesize(syntheticSamples) val (model, real) = sbc.fit(values) - val samples = - model.sample(sampler(goldset.size), 1).predict(real) + val inference = new MCMCInference(model, sampler(goldset.size), 1) + val samples = inference.sample.predict(real) (samples, trueValue) } diff --git a/rainier-test/src/test/scala/com/stripe/rainier/compute/CholeskyTest.scala b/rainier-test/src/test/scala/com/stripe/rainier/compute/CholeskyTest.scala index 1a32b6920..247a26578 100644 --- a/rainier-test/src/test/scala/com/stripe/rainier/compute/CholeskyTest.scala +++ b/rainier-test/src/test/scala/com/stripe/rainier/compute/CholeskyTest.scala @@ -1,6 +1,6 @@ package com.stripe.rainier.compute -import com.stripe.rainier.sampler.RNG +import com.stripe.rainier.RNG class CholeskyTest extends ComputeTest { val rng = RNG.default diff --git a/rainier-test/src/test/scala/com/stripe/rainier/compute/RealTest.scala b/rainier-test/src/test/scala/com/stripe/rainier/compute/RealTest.scala index 9c329816b..dde44ec8d 100644 --- a/rainier-test/src/test/scala/com/stripe/rainier/compute/RealTest.scala +++ b/rainier-test/src/test/scala/com/stripe/rainier/compute/RealTest.scala @@ -1,8 +1,16 @@ package com.stripe.rainier.compute +import com.stripe.rainier.RNG import com.stripe.rainier.core._ -import scala.util.{Try, Success, Failure} -import Double.{PositiveInfinity => Inf, NegativeInfinity => NegInf, NaN} + +import scala.util.{Failure, Success, Try} +import Double.{NaN, NegativeInfinity => NegInf, PositiveInfinity => Inf} + +object RealTest { + def generateValue(rng: RNG, x: Double): Double = { + 2 * x + } +} class RealTest extends ComputeTest { def run(description: String, @@ -33,7 +41,7 @@ class RealTest extends ComputeTest { val eval = new Evaluator(Map(x -> n)) val withVar = eval.toDouble(result) assertWithinEpsilon(constant, withVar, s"[c/ev, n=$n]") - val compiled = c(Array(n)) + val compiled = c(RNG.default, Array(n)) assertWithinEpsilon(withVar, compiled, s"[ev/ir, n=$n]") // derivatives of automated differentiation vs numeric differentiation @@ -45,7 +53,7 @@ class RealTest extends ComputeTest { assertWithinEpsilon(numDiff, diffWithVar, s"[numDiff/diffWithVar, n=$n]") - val diffCompiled = dc(Array(n)) + val diffCompiled = dc(RNG.default, Array(n)) assertWithinEpsilon(diffWithVar, diffCompiled, s"[diffWithVar/diffCompiled, n=$n]") @@ -202,4 +210,15 @@ class RealTest extends ComputeTest { Gamma.standard(x.abs).logDensity(y) }) } + + test("function call") { +// val l = Latent(1.0, FnCall(RealTest.getClass.getName, "generateValue", List(2.0))) + val l = FnCall(RealTest.getClass.getName.stripSuffix("$").replace('.', '/'), + "generateValue", + List(2.0)) + val c = Compiler(200, 100).compile(List.empty, l) + val rng = RNG.default + c(rng, Array(3.0)) + } + } diff --git a/rainier-test/src/test/scala/com/stripe/rainier/optimizer/OptimizerTest.scala b/rainier-test/src/test/scala/com/stripe/rainier/optimizer/OptimizerTest.scala index f6033c575..8f2cfa5e1 100644 --- a/rainier-test/src/test/scala/com/stripe/rainier/optimizer/OptimizerTest.scala +++ b/rainier-test/src/test/scala/com/stripe/rainier/optimizer/OptimizerTest.scala @@ -1,5 +1,6 @@ package com.stripe.rainier.optimizer +import com.stripe.rainier.RNG import com.stripe.rainier.core._ import org.scalatest.FunSuite @@ -18,7 +19,7 @@ class OptimizerTest extends FunSuite { def testLBFGS(model: Model): Unit = { val df = model.density() testLBFGS(df.nVars) { x => - df.update(x) + df.update(RNG.default, x) val f = df.density * -1 val g = 0.until(df.nVars).toArray.map { i => df.gradient(i) * -1 diff --git a/rainier-test/src/test/scala/com/stripe/rainier/sampler/LeapFrogTest.scala b/rainier-test/src/test/scala/com/stripe/rainier/sampler/LeapFrogTest.scala index c0931f84f..934690274 100644 --- a/rainier-test/src/test/scala/com/stripe/rainier/sampler/LeapFrogTest.scala +++ b/rainier-test/src/test/scala/com/stripe/rainier/sampler/LeapFrogTest.scala @@ -1,11 +1,12 @@ package com.stripe.rainier.sampler +import com.stripe.rainier.{RNG, ScalaRNG} import org.scalatest.FunSuite class NormalDensityFunction extends DensityFunction { val nVars = 1 var x = 0.0 - def update(vars: Array[Double]): Unit = { + def update(rng: RNG, vars: Array[Double]): Unit = { x = vars(0) } def density = (x * x) / -2.0