diff --git a/core/src/main/scala/dimwit/stats/IndependentDistributions.scala b/core/src/main/scala/dimwit/stats/IndependentDistributions.scala index 4fde99b..2ccb465 100644 --- a/core/src/main/scala/dimwit/stats/IndependentDistributions.scala +++ b/core/src/main/scala/dimwit/stats/IndependentDistributions.scala @@ -71,11 +71,16 @@ class Uniform[T <: Tuple: Labels, V: IsFloating](val low: Tensor[T, V], val high Jax.jrandom.uniform(key.jaxKey, shape = low.shape.dimensions.toPythonProxy, minval = low.jaxValue, maxval = high.jaxValue) ) -/** Uniform distribution */ +/** Discrete uniform distribution over the integers in the half-open interval `[min, max)`, + * i.e. min is inclusive, max exclusive. + */ class DiscreteUniform[T <: Tuple: Labels](val min: Tensor[T, Int32], val max: Tensor[T, Int32]) extends IndependentDistribution[T, Int32]: override def elementWiseLogProb(x: Tensor[T, Int32]): Tensor[T, LogProb] = - liftPyTensor(jstats.randint.logpmf(x.jaxValue, low = min.jaxValue, high = max.jaxValue)) + val logProb = LogProb(-(max - min).asFloat32.log) + val inSupport = (x >= min) and (x < max) + val negInf = LogProb(Tensor(x.shape).fill(Float.NegativeInfinity)) + where(inSupport, logProb, negInf) override def sample(key: Random.Key): Tensor[T, Int32] = liftPyTensor( @@ -89,11 +94,6 @@ object Uniform: require(low.shape.dimensions == high.shape.dimensions, "Low and high must have the same dimensions") new Uniform(low, high) - /** Create a discrete Uniform distribution from low and high int tensors */ - def apply[T <: Tuple: Labels](min: Tensor[T, Int32], max: Tensor[T, Int32]): DiscreteUniform[T] = - require(min.shape.dimensions == max.shape.dimensions, "min and max must have the same dimensions") - new DiscreteUniform(min, max) - /** Bernoulli distribution */ class Bernoulli[T <: Tuple: Labels](val probs: Tensor[T, Prob]) extends IndependentDistribution[T, Bool]: diff --git a/core/src/test/scala/dimwit/stats/DistributionSuite.scala b/core/src/test/scala/dimwit/stats/DistributionSuite.scala index c2982fb..fc7e6b9 100644 --- a/core/src/test/scala/dimwit/stats/DistributionSuite.scala +++ b/core/src/test/scala/dimwit/stats/DistributionSuite.scala @@ -110,6 +110,51 @@ class DistributionSuite extends DimwitTest: val expectedMeans = (uniform.low + uniform.high) *! 0.5f sampleMeans should approxEqual(expectedMeans, 0.2f) + describe("DiscreteUniform Distribution"): + it("logProbs is -log(max - min) inside [min, max)"): + val min = Tensor(Shape(Axis[A] -> 3)).fromArray(Array(0, -3, 2)) + val max = Tensor(Shape(Axis[A] -> 3)).fromArray(Array(4, 3, 10)) + val x = Tensor(Shape(Axis[A] -> 3)).fromArray(Array(1, 0, 9)) + + val dist = DiscreteUniform(min, max) + val scalaLogProbs = dist.elementWiseLogProb(x) + val expectedLogProbs = Tensor(Shape(Axis[A] -> 3)).fromArray( + Array(-math.log(4).toFloat, -math.log(6).toFloat, -math.log(8).toFloat) + ) + scalaLogProbs.asFloat should approxEqual(expectedLogProbs) + + it("logProb is -inf outside [min, max)"): + val min = Tensor(Shape(Axis[A] -> 2)).fromArray(Array(0, 0)) + val max = Tensor(Shape(Axis[A] -> 2)).fromArray(Array(4, 4)) + val x = Tensor(Shape(Axis[A] -> 2)).fromArray(Array(-1, 4)) + + val dist = DiscreteUniform(min, max) + val logProbs = dist.elementWiseLogProb(x) + logProbs.asFloat.toArray.foreach(v => v should be(Float.NegativeInfinity)) + + it("sample means approximates means"): + val discreteUniform = DiscreteUniform( + Tensor(Shape(Axis[A] -> 2)).fromArray(Array(-2, 0)), + Tensor(Shape(Axis[A] -> 2)).fromArray(Array(3, 10)) + ) + val key = Random.Key(42) + val samples = key.splitvmap(Axis[Samples] -> 10000)(k => discreteUniform.sample(k)) + val sampleMeans = samples.asFloat32.mean(Axis[Samples]) + // max is exclusive, so the mean is (min + max - 1) / 2 + val expectedMeans = Tensor(Shape(Axis[A] -> 2)).fromArray(Array(0.0f, 4.5f)) + sampleMeans should approxEqual(expectedMeans, 0.2f) + + it("samples strictly respect bounds"): + val dist = DiscreteUniform( + Tensor(Shape(Axis[A] -> 2)).fromArray(Array(-2, 0)), + Tensor(Shape(Axis[A] -> 2)).fromArray(Array(3, 10)) + ) + val key = Random.Key(42) + val samples = key.splitvmap(Axis[Samples] -> 10000)(k => dist.sample(k)) + + (samples.min(Axis[Samples]) >= dist.min).all.item shouldBe true + (samples.max(Axis[Samples]) < dist.max).all.item shouldBe true + describe("Bernoulli"): it("logProbs matches JAX"): val probs = Tensor(Shape(Axis[A] -> 3)).fromArray(Array(0.3f, 0.5f, 0.8f))