Skip to content

Commit 4f0aeb0

Browse files
committed
Add FFI fma and normal/exponential distribution
1 parent 012ba5e commit 4f0aeb0

18 files changed

Lines changed: 880 additions & 111 deletions

‎.gitignore‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -28,4 +28,5 @@
2828
*.synctex.gz
2929
*.synctex.gz(busy)
3030
*.pdfsync
31-
test_data/
31+
test_data/
32+
__pycache__/

‎FFI.lean‎

Lines changed: 7 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,7 @@
1+
module -- shake: keep-all --deprecated_module: ignore
2+
3+
public import FFI.FMA
4+
public import FFI.Float
5+
public import FFI.Float32
6+
public import FFI.Model.Float
7+
public import FFI.Model.Float32

‎FFI/FMA.lean‎

Lines changed: 57 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,57 @@
1+
/-
2+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Gaëtan Serré
5+
-/
6+
module
7+
8+
prelude
9+
public import Init.Data.Float.Model.Unpacked.Round
10+
11+
-- This file is part of the logical model for floats which authors of float libraries
12+
-- need to rely on.
13+
@[expose] public section
14+
15+
namespace Float.Model.UnpackedFloat
16+
17+
/--
18+
Computes the fused multiply-add `x * y + z` of three floating point numbers and rounds the
19+
result according to the given specification. The product `x * y` is exact and never rounded
20+
on its own; only the final sum is rounded.
21+
-/
22+
def fma (spec : Format) : UnpackedFloat → UnpackedFloat → UnpackedFloat → UnpackedFloat
23+
| .notANumber, _, _ => .notANumber
24+
| _, .notANumber, _ => .notANumber
25+
| _, _, .notANumber => .notANumber
26+
| .zero _, .infinity _, _ => .notANumber
27+
| .infinity _, .zero _, _ => .notANumber
28+
| .infinity sign₁, .infinity sign₂, .infinity sign₃ =>
29+
if sign₁ * sign₂ == sign₃ then .infinity sign₃ else .notANumber
30+
| .infinity sign₁, .finite sign₂ .., .infinity sign₃ =>
31+
if sign₁ * sign₂ == sign₃ then .infinity sign₃ else .notANumber
32+
| .finite sign₁ .., .infinity sign₂, .infinity sign₃ =>
33+
if sign₁ * sign₂ == sign₃ then .infinity sign₃ else .notANumber
34+
| .infinity sign₁, .infinity sign₂, _ => .infinity (sign₁ * sign₂)
35+
| .infinity sign₁, .finite sign₂ .., _ => .infinity (sign₁ * sign₂)
36+
| .finite sign₁ .., .infinity sign₂, _ => .infinity (sign₁ * sign₂)
37+
| _, _, .infinity sign₃ => .infinity sign₃
38+
| .zero sign₁, .zero sign₂, .zero sign₃ =>
39+
if sign₁ * sign₂ == sign₃ then .zero sign₃ else .zero .positive
40+
| .zero sign₁, .finite sign₂ .., .zero sign₃ =>
41+
if sign₁ * sign₂ == sign₃ then .zero sign₃ else .zero .positive
42+
| .finite sign₁ .., .zero sign₂, .zero sign₃ =>
43+
if sign₁ * sign₂ == sign₃ then .zero sign₃ else .zero .positive
44+
| .zero _, _, z => z
45+
| _, .zero _, z => z
46+
| .finite s₁ m₁ e₁ _, .finite s₂ m₂ e₂ _, .zero _ =>
47+
roundWithAccuracy spec (s₁ * s₂) (m₁ * m₂) (e₁ + e₂) .exact
48+
| .finite s₁ m₁ e₁ _, .finite s₂ m₂ e₂ _, .finite s₃ m₃ e₃ _ =>
49+
let productMantissa := m₁ * m₂
50+
let productExponent := e₁ + e₂
51+
let smallerExponent := min productExponent e₃
52+
let (productMantissa, _) := decreaseExponent productMantissa productExponent smallerExponent
53+
let (m₃, _) := decreaseExponent m₃ e₃ smallerExponent
54+
let mantissa := (s₁ * s₂).apply productMantissa + s₃.apply m₃
55+
normalize spec mantissa smallerExponent .positive
56+
57+
end Float.Model.UnpackedFloat

‎FFI/Float.lean‎

Lines changed: 31 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,31 @@
1+
/-
2+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Gaëtan Serré
5+
-/
6+
module
7+
8+
prelude
9+
public import Init.Data.Float.Float
10+
public import FFI.Model.Float
11+
12+
@[expose] public section
13+
14+
namespace Float
15+
16+
/--
17+
Computes the fused multiply-add `x * y + z` of three floating-point numbers. This operation is
18+
performed with a single rounding, which can be more accurate than performing the multiplication and
19+
addition separately.
20+
21+
This function has a logical model in terms of `Float.Model`. It is implemented in compiled code by
22+
the C function `fma`.
23+
-/
24+
@[extern "fma"] def fma : Float → Float → Float → Float :=
25+
fun x y z => .ofModel (x.toModel.fma y.toModel z.toModel)
26+
27+
/-- `log (1 + x)`, the C99 `log1p`, accurate for small `x` where `log (1 + x)` loses the leading
28+
digits of the result to the rounding of `1 + x`. -/
29+
@[extern "log1p"] opaque log1p (x : Float) : Float
30+
31+
end Float

‎FFI/Float32.lean‎

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,29 @@
1+
/-
2+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Gaëtan Serré
5+
-/
6+
module
7+
8+
prelude
9+
public import Init.Data.Float.Float32
10+
public import FFI.Model.Float32
11+
12+
@[expose] public section
13+
14+
namespace Float32
15+
16+
/--
17+
Computes the fused multiply-add `x * y + z` of three floating-point numbers. This operation is performed with a single rounding, which can be more accurate than performing the multiplication and addition separately.
18+
19+
This function has a logical model in terms of `Float32.Model`. It is implemented in compiled code
20+
by the C function `fmaf`.
21+
-/
22+
@[extern "fmaf"] def fma : Float32 → Float32 → Float32 → Float32 :=
23+
fun x y z => .ofModel (x.toModel.fma y.toModel z.toModel)
24+
25+
/-- `log (1 + x)`, the C99 `log1p`, accurate for small `x` where `log (1 + x)` loses the leading
26+
digits of the result to the rounding of `1 + x`. -/
27+
@[extern "log1pf"] opaque log1p (x : Float32) : Float32
28+
29+
end Float32

‎FFI/Model/Float.lean‎

Lines changed: 24 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,24 @@
1+
/-
2+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Gaëtan Serré
5+
-/
6+
module
7+
8+
prelude
9+
public import Init.Data.Float.Model.Float
10+
public import FFI.FMA
11+
12+
-- This file is part of the logical model for floats which authors of float libraries
13+
-- need to rely on.
14+
@[expose] public section
15+
16+
namespace Float.Model
17+
18+
/--
19+
Compute the fused multiply-add `a * b + c` of three `Float.Model`, with a single rounding.
20+
-/
21+
def fma (a b c : Float.Model) : Float.Model :=
22+
pack (UnpackedFloat.fma Format.binary64 a.unpack b.unpack c.unpack)
23+
24+
end Float.Model

‎FFI/Model/Float32.lean‎

Lines changed: 26 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,26 @@
1+
/-
2+
Copyright (c) 2026 Gaëtan Serré. All rights reserved.
3+
Released under Apache 2.0 license as described in the file LICENSE.
4+
Authors: Gaëtan Serré
5+
-/
6+
module
7+
8+
prelude
9+
public import Init.Data.Float.Model.Float32
10+
public import FFI.FMA
11+
12+
-- This file is part of the logical model for floats which authors of float libraries
13+
-- need to rely on.
14+
@[expose] public section
15+
16+
namespace Float32.Model
17+
18+
open Float.Model (Format UnpackedFloat)
19+
20+
/--
21+
Compute the fused multiply-add `a * b + c` of three `Float32.Model`, with a single rounding.
22+
-/
23+
def fma (a b c : Float32.Model) : Float32.Model :=
24+
pack (UnpackedFloat.fma Format.binary32 a.unpack b.unpack c.unpack)
25+
26+
end Float32.Model

‎RandomDo.lean‎

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -9,6 +9,7 @@ public import RandomDo.Monad.Notation
99
public import RandomDo.NumLean.Distributions
1010
public import RandomDo.NumLean.PCG64
1111
public import RandomDo.NumLean.SeedSequence
12+
public import RandomDo.NumLean.Ziggurat
1213
public import RandomDo.Tactic.Deriving
1314
public import RandomDo.Tactic.Elab
1415
public import RandomDo.Tactic.ForInStep

‎RandomDo/NumLean/Distributions.lean‎

Lines changed: 78 additions & 14 deletions
Original file line numberDiff line numberDiff line change
@@ -7,7 +7,8 @@ module
77

88
public import RandomDo.NumLean.PCG64
99
public meta import RandomDo.NumLean.PCG64
10-
public import Batteries.Data.Float.Basic
10+
public import FFI.Float
11+
public import RandomDo.NumLean.Ziggurat
1112

1213
/-!
1314
# Sample from specific distributions using the PCG-64 generator.
@@ -20,6 +21,11 @@ seeded alike produce the same values.
2021
* `randUInt64` / `randUInt32`: sample a `UInt64`, or a `UInt32` as numpy's `next_uint32` does.
2122
* `random`: sample a `Float` in `[0, 1)`.
2223
* `randInt`: sample an integer in `[low, high)` or `[low, high]`, as numpy's `Generator.integers`.
24+
* `uniform`: sample a `Float` in `[low, high)`, as numpy's `Generator.uniform`.
25+
* `standardNormal` / `normal`: sample a normal deviate, as numpy's `Generator.normal`.
26+
* `standardExponential` / `exponential`: sample an exponential deviate, as numpy's
27+
`Generator.exponential`.
28+
* `ziggurat`: the sampler shape both of those share.
2329
-/
2430

2531
@[expose] public section
@@ -94,22 +100,80 @@ def random : RandPCG IO Float := do
94100
let x ← randUInt64
95101
return (x >>> 11).toFloat * (Float.ofBits <| 0x3CA <<< (52 : UInt64))
96102

97-
/-- `x * y + z`, rounded once, as the C `fma`: the product and the sum are computed exactly, as
98-
integers scaled by a power of two, and only the final value is rounded to a `Float`. -/
99-
def fma (x y z : Float) : Float :=
100-
match x.toRatParts, y.toRatParts, z.toRatParts with
101-
| some (vx, ex), some (vy, ey), some (vz, ez) =>
102-
-- the product is `vx * vy * 2 ^ (ex + ey)`, so the sum is exact over the common exponent `e`
103-
let ep := ex + ey
104-
let e := min ep ez
105-
let n := vx * vy * 2 ^ (ep - e).toNat + vz * 2 ^ (ez - e).toNat
106-
if e ≥ 0 then Int.divFloat (n * 2 ^ e.toNat) 1 else Int.divFloat n (2 ^ (-e).toNat)
107-
| _, _, _ => x * y + z
108-
109103
/-- Sample a `Float` uniformly in `[low, high)`. -/
110104
def uniform (low high : Float) : RandPCG IO Float := do
111105
if low > high then throw <| IO.userError "low > high"
112106
let x ← random
113-
return fma x (high - low) low
107+
return Float.fma x (high - low) low
108+
109+
/-- The rejection test on a strip that sticks out of the curve: the point drawn at height `u`
110+
between the density at the strip's two edges lies under the curve. -/
111+
@[inline] def zigguratWedge (f : Array Float) (idx : Nat) (u density : Float) : Bool :=
112+
Float.fma (f[idx - 1]! - f[idx]!) u f[idx]! < density
113+
114+
/-- The sampler shape shared by numpy's `random_standard_normal` and
115+
`random_standard_exponential`, the ziggurat of Marsaglia and Tsang: the density is covered by 256
116+
strips of equal area, and one 64-bit output supplies at once the strip `idx` and the integer `ri`
117+
that `w` scales to an abscissa, `split` saying how those bits are laid out and whether the deviate
118+
comes out negated.
119+
120+
The draw is returned as it stands when `ri` falls below `k[idx]`, which is where about 99% of the
121+
draws end. Otherwise the strip either is the base one, whose unbounded part `tail` samples from the
122+
integer drawn, or sticks out of the curve, and then the point is tested against `density` and the
123+
whole draw is started over on rejection. -/
124+
@[specialize] partial def ziggurat (k : Array UInt64) (w f : Array Float)
125+
(split : UInt64 → Nat × UInt64 × Bool) (density : Float → Float)
126+
(tail : UInt64 → RandPCG IO Float) : RandPCG IO Float := do
127+
let (idx, ri, negate) := split (← randUInt64)
128+
let x := ri.toFloat * w[idx]!
129+
let x := if negate then -x else x
130+
if ri < k[idx]! then return x
131+
if idx == 0 then tail ri
132+
else if zigguratWedge f idx (← random) (density x) then return x
133+
else ziggurat k w f split density tail
134+
135+
/-- Sample the tail of the standard normal beyond `Ziggurat.norR`, as the `idx == 0` branch of
136+
numpy's `random_standard_normal`: draw from an exponential tail until the pair of draws falls under
137+
the normal's, which is Marsaglia's method for the tail. `negate` carries the sign numpy reads off
138+
the integer already drawn, which the loop does not redraw. -/
139+
partial def normalTail (negate : Bool) : RandPCG IO Float := do
140+
let xx := -Ziggurat.norInvR * Float.log1p (-(← random))
141+
let yy := -Float.log1p (-(← random))
142+
if yy + yy > xx * xx then
143+
return if negate then -(Ziggurat.norR + xx) else Ziggurat.norR + xx
144+
normalTail negate
145+
146+
/-- Sample from the standard normal, as numpy's `random_standard_normal`. The 64-bit output gives
147+
the strip in its low byte, then the sign, then a 52-bit abscissa; the tail takes its sign from a
148+
further bit of that same abscissa, as numpy does. -/
149+
def standardNormal : RandPCG IO Float :=
150+
ziggurat Ziggurat.ki Ziggurat.wi Ziggurat.fi
151+
(fun r =>
152+
let idx := (r &&& 0xFF).toNat
153+
let r := r >>> 8
154+
(idx, (r >>> 1) &&& 0x000FFFFFFFFFFFFF, (r &&& 1) == 1))
155+
(fun x => Float.exp ((-0.5) * x * x))
156+
(fun rabs => normalTail (((rabs >>> 8) &&& 1) == 1))
157+
158+
/-- Sample from the standard exponential, as numpy's `random_standard_exponential`. The 64-bit
159+
output is first shifted by three, then gives the strip in its low byte and a 53-bit abscissa; the
160+
tail is the exponential's own, memoryless, so one draw beyond `Ziggurat.expR` suffices. -/
161+
def standardExponential : RandPCG IO Float :=
162+
ziggurat Ziggurat.ke Ziggurat.we Ziggurat.fe
163+
(fun r =>
164+
let r := r >>> 3
165+
((r &&& 0xFF).toNat, r >>> 8, false))
166+
(fun x => Float.exp (-x))
167+
(fun _ => do return Ziggurat.expR - Float.log1p (-(← random)))
168+
169+
/-- Draw random samples from a normal (Gaussian) distribution. -/
170+
def normal (loc : Float := 0) (scale : Float := 1) : RandPCG IO Float := do
171+
if scale < 0 then throw <| IO.userError "scale < 0"
172+
return Float.fma scale (← standardNormal) loc
173+
174+
/-- Draw samples from an exponential distribution. -/
175+
def exponential (scale : Float := 1) : RandPCG IO Float := do
176+
if scale < 0 then throw <| IO.userError "scale < 0"
177+
return scale * (← standardExponential)
114178

115179
end NumLean

0 commit comments

Comments
 (0)