Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 7 additions & 7 deletions core/src/main/scala/dimwit/stats/IndependentDistributions.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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(
Expand All @@ -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]:

Expand Down
45 changes: 45 additions & 0 deletions core/src/test/scala/dimwit/stats/DistributionSuite.scala
Original file line number Diff line number Diff line change
Expand Up @@ -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))
Expand Down
Loading