17 | module Spidr.Distribution
22 | import Spidr.Constants
31 | interface Distribution (0 dist : (0 event : Shape) -> (0 dim : Nat) -> Type) where
33 | mean : dist event dim -> Tag $
Tensor (dim :: event) F64
36 | cov : dist event dim -> Tag $
Tensor (dim :: dim :: event) F64
40 | variance : {event : _} -> Distribution dist => dist event 1 -> Tag $
Tensor (1 :: event) F64
41 | variance dist = squeeze {from = (1 :: 1 :: event)} <$> cov dist
51 | interface Distribution dist =>
52 | ClosedFormDistribution (0 event : Shape)
53 | (0 dist : (0 event : Shape) -> (0 dim : Nat) -> Type) where
55 | pdf : dist event (S d) -> Tensor (S d :: event) F64 -> Tag $
Tensor [] F64
59 | cdf : dist event (S d) -> Tensor (S d :: event) F64 -> Tag $
Tensor [] F64
66 | data Gaussian : (0 event : Shape) -> (0 dim : Nat) -> Type where
69 | MkGaussian : {d : Nat} -> (mean : Tensor (S d :: event) F64) ->
70 | (cov : Tensor (S d :: S d :: event) F64) ->
71 | Gaussian event (S d)
74 | Taggable (Gaussian event dim) where
75 | tag (MkGaussian mean cov) = [| MkGaussian (tag mean) (tag cov) |]
78 | Distribution Gaussian where
79 | mean (MkGaussian mean' _) = pure mean'
80 | cov (MkGaussian _ cov') = pure cov'
84 | ClosedFormDistribution [1] Gaussian where
85 | pdf (MkGaussian {d} mean cov) x = do
86 | cholCov <- tag !(cholesky $
squeeze {to = [S d, S d]} cov)
87 | tri <- tag $
cholCov |\ squeeze (x - mean)
88 | let exponent = - tri @@ tri / 2.0
89 | covSqrtDet <- reduce @{Prod} [0] (diag cholCov)
90 | let denominator = fromDouble (pow (2.0 * pi) (cast (S d) / 2.0)) * covSqrtDet
91 | pure (exp exponent / denominator)
93 | cdf (MkGaussian {d = S _} _ _) _ =
94 | assert_total $
idris_crash "CDF not implemented for multivariate Gaussian"
95 | cdf (MkGaussian {d = 0} mean cov) x =
96 | pure $
(1.0 + erf (squeeze (x - mean) / (sqrt (squeeze cov * 2.0)))) / 2.0