import JaxLean.Stdlib
import examples.randint_monte_carlo.generated.DieEstimate
import examples.randint_monte_carlo.generated.GridEstimate
import examples.randint_monte_carlo.generated.MCDie16
import examples.randint_monte_carlo.generated.MCGridSquare8Two ordinary-JAX Monte Carlo examples, proved at estimator boundaries. The uniform/iid laws below are the explicit ideal randint/split specification. Deterministic estimator kernels have separate checked Jaxpr certificates. No claim about bitwise PRNG uniformity, seed independence, or floating rounding.
namespace JaxLean.MCJaxopen scoped BigOperatorsnoncomputable sectionabbrev dieLaw : Rand ℝ := Rand.uniformInt 1 6 (⊢ 0 < 6 All goals completed! 🐙)abbrev gridIndexLaw : Rand ℝ := Rand.uniformInt 0 4 (⊢ 0 < 4 All goals completed! 🐙)def gridValueLaw : Rand ℝ := Rand.map (fun v => (v / 4) ^ (2 : Nat)) gridIndexLawContracts for the deterministic Python functions.
theorem die_estimate_spec (draws : Tensor ℝ [16]) :
die_estimate draws () = (∑ i, draws (i, ())) / 16 := rfltheorem grid_square_spec (indices : Tensor ℝ [8]) :
grid_square indices = Tensor.map (fun v => (v / 4) ^ (2 : Nat)) indices := indices:Tensor ℝ [8]⊢ grid_square indices = Tensor.map (fun v => (v / 4) ^ 2) indices
indices:Tensor ℝ [8]i:Index [8]⊢ grid_square indices i = Tensor.map (fun v => (v / 4) ^ 2) indices i
indices:Tensor ℝ [8]i:Index [8]⊢ (indices i / 4) ^ 2 = (indices i / 4) ^ 2
All goals completed! 🐙theorem grid_estimate_spec (indices : Tensor ℝ [8]) :
grid_estimate indices () = (∑ i, (indices (i, ()) / 4) ^ (2 : Nat)) / 8 := indices:Tensor ℝ [8]⊢ grid_estimate indices () = (∑ i, (indices (i, ()) / 4) ^ 2) / 8
All goals completed! 🐙Bridges from the emitted random programs to their certified estimator functions.
theorem mc_die_boundary :
mc_die_16 (R := ℝ) =
Rand.map (fun draws => die_estimate (fun i => draws i.1) ()) (Rand.iid 16 dieLaw) := rfltheorem mc_grid_boundary :
mc_grid_square_8 (R := ℝ) =
Rand.map (fun draws => grid_estimate (fun i => draws i.1) ())
(Rand.iid 8 gridIndexLaw) := rflOnly the small, single-draw laws are enumerated; never the product sample space.
theorem die_mean : dieLaw.mean = 7 / 2 := ⊢ dieLaw.mean = 7 / 2
⊢ ∑ ω, 1 / 6 * (↑1 + ↑↑ω) = 7 / 2
All goals completed! 🐙theorem die_variance : dieLaw.variance = 35 / 12 := ⊢ dieLaw.variance = 35 / 12
⊢ ∑ ω, 1 / 6 * (↑1 + ↑↑ω - ∑ ω, 1 / 6 * (↑1 + ↑↑ω)) ^ 2 = 35 / 12
All goals completed! 🐙theorem grid_value_mean : gridValueLaw.mean = 7 / 32 := ⊢ gridValueLaw.mean = 7 / 32
⊢ ∑ ω, 1 / 4 * ((↑0 + ↑↑ω) / 4) ^ 2 = 7 / 32
All goals completed! 🐙theorem grid_value_variance : gridValueLaw.variance = 49 / 1024 := ⊢ gridValueLaw.variance = 49 / 1024
⊢ ∑ ω, 1 / 4 * (((↑0 + ↑↑ω) / 4) ^ 2 - ∑ ω, 1 / 4 * ((↑0 + ↑↑ω) / 4) ^ 2) ^ 2 = 49 / 1024
All goals completed! 🐙Sampling and averaging preserve the one-draw mean.
⊢ (∑ i, (Rand.map (fun v => v) dieLaw).mean) / 16 = 7 / 2
change (∑ _ : Fin 16, dieLaw.mean) / 16 = _ ⊢ (∑ x, dieLaw.mean) / 16 = 7 / 2
norm_num [die_mean] All goals completed! 🐙The 1/n law is obtained by propagating division and independent-sum rules.
theorem mc_die_variance_reduction :
(mc_die_16 (R := ℝ)).variance = dieLaw.variance / 16 := by ⊢ mc_die_16.variance = dieLaw.variance / 16
rw [mc_die_boundary ⊢ (Rand.map (fun draws => die_estimate (fun i => draws i.1) ()) (Rand.iid 16 dieLaw)).variance = dieLaw.variance / 16 ⊢ (Rand.map (fun draws => die_estimate (fun i => draws i.1) ()) (Rand.iid 16 dieLaw)).variance = dieLaw.variance / 16] ⊢ (Rand.map (fun draws => die_estimate (fun i => draws i.1) ()) (Rand.iid 16 dieLaw)).variance = dieLaw.variance / 16
simp only [die_estimate_spec] ⊢ (Rand.map (fun draws => (∑ x, draws x) / 16) (Rand.iid 16 dieLaw)).variance = dieLaw.variance / 16
rw [Rand.variance_map_div, ⊢ (Rand.map (fun draws => ∑ x, draws x) (Rand.iid 16 dieLaw)).variance / 16 ^ 2 = dieLaw.variance / 16 ⊢ (∑ i, (Rand.map (fun v => v) dieLaw).variance) / 16 ^ 2 = dieLaw.variance / 16 Rand.variance_iid_sum (f := fun _ v => v) ⊢ (∑ i, (Rand.map (fun v => v) dieLaw).variance) / 16 ^ 2 = dieLaw.variance / 16 ⊢ (∑ i, (Rand.map (fun v => v) dieLaw).variance) / 16 ^ 2 = dieLaw.variance / 16] ⊢ (∑ i, (Rand.map (fun v => v) dieLaw).variance) / 16 ^ 2 = dieLaw.variance / 16
change (∑ _ : Fin 16, dieLaw.variance) / 16 ^ 2 = _ ⊢ (∑ x, dieLaw.variance) / 16 ^ 2 = dieLaw.variance / 16
simp only [Finset.sum_const, Finset.card_univ, Fintype.card_fin, nsmul_eq_mul] ⊢ ↑16 * dieLaw.variance / 16 ^ 2 = dieLaw.variance / 16
ring All goals completed! 🐙
theorem mc_die_variance : (mc_die_16 (R := ℝ)).variance = 35 / 192 := by ⊢ mc_die_16.variance = 35 / 192
rw [mc_die_variance_reduction, ⊢ dieLaw.variance / 16 = 35 / 192 ⊢ 35 / 12 / 16 = 35 / 192 die_variance ⊢ 35 / 12 / 16 = 35 / 192 ⊢ 35 / 12 / 16 = 35 / 192] ⊢ 35 / 12 / 16 = 35 / 192
norm_num All goals completed! 🐙The same argument works after a nonlinear per-sample transformation.
theorem mc_grid_mean : (mc_grid_square_8 (R := ℝ)).mean = 7 / 32 := by ⊢ mc_grid_square_8.mean = 7 / 32
rw [mc_grid_boundary ⊢ (Rand.map (fun draws => grid_estimate (fun i => draws i.1) ()) (Rand.iid 8 gridIndexLaw)).mean = 7 / 32 ⊢ (Rand.map (fun draws => grid_estimate (fun i => draws i.1) ()) (Rand.iid 8 gridIndexLaw)).mean = 7 / 32] ⊢ (Rand.map (fun draws => grid_estimate (fun i => draws i.1) ()) (Rand.iid 8 gridIndexLaw)).mean = 7 / 32
simp only [grid_estimate_spec] ⊢ (Rand.map (fun draws => (∑ x, (draws x / 4) ^ 2) / 8) (Rand.iid 8 gridIndexLaw)).mean = 7 / 32
rw [Rand.mean_map_div, ⊢ (Rand.map (fun draws => ∑ x, (draws x / 4) ^ 2) (Rand.iid 8 gridIndexLaw)).mean / 8 = 7 / 32 ⊢ (∑ i, (Rand.map (fun v => (v / 4) ^ 2) gridIndexLaw).mean) / 8 = 7 / 32 Rand.mean_iid_sum (f := fun _ v => (v / 4) ^ (2 : Nat)) ⊢ (∑ i, (Rand.map (fun v => (v / 4) ^ 2) gridIndexLaw).mean) / 8 = 7 / 32 ⊢ (∑ i, (Rand.map (fun v => (v / 4) ^ 2) gridIndexLaw).mean) / 8 = 7 / 32] ⊢ (∑ i, (Rand.map (fun v => (v / 4) ^ 2) gridIndexLaw).mean) / 8 = 7 / 32
change (∑ _ : Fin 8, gridValueLaw.mean) / 8 = _ ⊢ (∑ x, gridValueLaw.mean) / 8 = 7 / 32
norm_num [grid_value_mean] All goals completed! 🐙
theorem mc_grid_variance_reduction :
(mc_grid_square_8 (R := ℝ)).variance = gridValueLaw.variance / 8 := by ⊢ mc_grid_square_8.variance = gridValueLaw.variance / 8
rw [mc_grid_boundary ⊢ (Rand.map (fun draws => grid_estimate (fun i => draws i.1) ()) (Rand.iid 8 gridIndexLaw)).variance =
gridValueLaw.variance / 8 ⊢ (Rand.map (fun draws => grid_estimate (fun i => draws i.1) ()) (Rand.iid 8 gridIndexLaw)).variance =
gridValueLaw.variance / 8] ⊢ (Rand.map (fun draws => grid_estimate (fun i => draws i.1) ()) (Rand.iid 8 gridIndexLaw)).variance =
gridValueLaw.variance / 8
simp only [grid_estimate_spec] ⊢ (Rand.map (fun draws => (∑ x, (draws x / 4) ^ 2) / 8) (Rand.iid 8 gridIndexLaw)).variance = gridValueLaw.variance / 8
rw [Rand.variance_map_div, ⊢ (Rand.map (fun draws => ∑ x, (draws x / 4) ^ 2) (Rand.iid 8 gridIndexLaw)).variance / 8 ^ 2 = gridValueLaw.variance / 8 ⊢ (∑ i, (Rand.map (fun v => (v / 4) ^ 2) gridIndexLaw).variance) / 8 ^ 2 = gridValueLaw.variance / 8 Rand.variance_iid_sum (f := fun _ v => (v / 4) ^ (2 : Nat)) ⊢ (∑ i, (Rand.map (fun v => (v / 4) ^ 2) gridIndexLaw).variance) / 8 ^ 2 = gridValueLaw.variance / 8 ⊢ (∑ i, (Rand.map (fun v => (v / 4) ^ 2) gridIndexLaw).variance) / 8 ^ 2 = gridValueLaw.variance / 8] ⊢ (∑ i, (Rand.map (fun v => (v / 4) ^ 2) gridIndexLaw).variance) / 8 ^ 2 = gridValueLaw.variance / 8
change (∑ _ : Fin 8, gridValueLaw.variance) / 8 ^ 2 = _ ⊢ (∑ x, gridValueLaw.variance) / 8 ^ 2 = gridValueLaw.variance / 8
simp only [Finset.sum_const, Finset.card_univ, Fintype.card_fin, nsmul_eq_mul] ⊢ ↑8 * gridValueLaw.variance / 8 ^ 2 = gridValueLaw.variance / 8
ring All goals completed! 🐙
theorem mc_grid_variance : (mc_grid_square_8 (R := ℝ)).variance = 49 / 8192 := by ⊢ mc_grid_square_8.variance = 49 / 8192
rw [mc_grid_variance_reduction, ⊢ gridValueLaw.variance / 8 = 49 / 8192 ⊢ 49 / 1024 / 8 = 49 / 8192 grid_value_variance ⊢ 49 / 1024 / 8 = 49 / 8192 ⊢ 49 / 1024 / 8 = 49 / 8192] ⊢ 49 / 1024 / 8 = 49 / 8192
norm_num All goals completed! 🐙The same numeric result for the independent evaluator of the estimator's Jaxpr. Only the sampling law is supplied as a specification.
theorem certified_die_variance :
(Rand.map (fun draws =>
Jaxpr.Program.eval (.cons (fun i => draws i.1) .nil) die_estimate_ir ())
(Rand.iid 16 dieLaw)).variance = 35 / 192 := by ⊢ (Rand.map (fun draws => Jaxpr.Program.eval (Jaxpr.Env.cons (fun i => draws i.1) Jaxpr.Env.nil) die_estimate_ir ())
(Rand.iid 16 dieLaw)).variance =
35 / 192
simp only [die_estimate_translation_correct] ⊢ (Rand.map (fun draws => die_estimate (fun i => draws i.1) ()) (Rand.iid 16 dieLaw)).variance = 35 / 192
rw [← mc_die_boundary ⊢ mc_die_16.variance = 35 / 192 ⊢ mc_die_16.variance = 35 / 192] ⊢ mc_die_16.variance = 35 / 192
exact mc_die_variance All goals completed! 🐙
theorem certified_grid_variance :
(Rand.map (fun draws =>
Jaxpr.Program.eval (.cons (fun i => draws i.1) .nil) grid_estimate_ir ())
(Rand.iid 8 gridIndexLaw)).variance = 49 / 8192 := by ⊢ (Rand.map (fun draws => Jaxpr.Program.eval (Jaxpr.Env.cons (fun i => draws i.1) Jaxpr.Env.nil) grid_estimate_ir ())
(Rand.iid 8 gridIndexLaw)).variance =
49 / 8192
simp only [grid_estimate_translation_correct] ⊢ (Rand.map (fun draws => grid_estimate (fun i => draws i.1) ()) (Rand.iid 8 gridIndexLaw)).variance = 49 / 8192
rw [← mc_grid_boundary ⊢ mc_grid_square_8.variance = 49 / 8192 ⊢ mc_grid_square_8.variance = 49 / 8192] ⊢ mc_grid_square_8.variance = 49 / 8192
exact mc_grid_variance All goals completed! 🐙endend JaxLean.MCJax