AUREOLE-R-v3 / tests /test_contract.py
PureOne's picture
AUREOLE-R 3.0.0-hf.1: standalone public research release
9d6c005 verified
Raw
History Blame Contribute Delete
11.5 kB
"""Algebraic and adversarial tests for the physical-correction contract."""
import itertools,tempfile,unittest
from pathlib import Path
import numpy as np
from aureole import freeze,draw,correct,sample,exact_mse,residual_metric,WorldMemory
from aureole.core import hoeffding_radius,proposal_from_bound
from aureole.renderer import Scene,receiver_grid,light_grid,unoccluded,physical_table,VisibilityPrior
class EstimatorContract(unittest.TestCase):
def test_enumerated_unbiasedness_arbitrarily_wrong_memory(self):
rng=np.random.default_rng(4)
f=rng.random((3,7,2));h=rng.normal(size=f.shape)*15
q=rng.random((3,7));q/=q.sum(1,keepdims=True)
s=freeze(h,q);mean=np.zeros((3,2))
for j in range(7):
ids=np.full((3,1),j)
mean+=q[:,j,None]*correct(s,ids,f[:,j:j+1,:])
np.testing.assert_allclose(mean,f.sum(1),atol=1e-12)
def test_two_sample_variance_matches_enumeration(self):
f=np.array([[[.2],[1.1],[.5]]]);h=np.array([[[.9],[.1],[.3]]]);q=np.array([[.2,.3,.5]])
s=freeze(h,q);mse=0
for a,b in itertools.product(range(3),repeat=2):
j=np.array([[a,b]]);result=correct(s,j,f[np.array([[0,0]]),j])
mse+=q[0,a]*q[0,b]*float((result[0,0]-f.sum())**2)
self.assertAlmostEqual(mse,float(exact_mse(f,s,2)[0]),places=12)
def test_frozen_snapshot_is_not_aliased_to_updated_memory(self):
h=np.ones((1,3,1));q=np.full((1,3),1/3);s=freeze(h,q)
h[:]=8;q[:]=0
np.testing.assert_allclose(s.integral,3)
with self.assertRaises(ValueError):s.control[0,0,0]=9
def test_same_sample_refitting_counterexample(self):
# Memorize the one sampled term, then pretend that refitted function was
# fixed before the draw. The residual is zero and mean becomes I/K.
f=np.array([1.,3.]);outputs=[]
for j in range(2):
h=np.zeros((1,2,1));h[0,j,0]=f[j]
outputs.append(correct(freeze(h,np.full((1,2),.5)),np.array([[j]]),np.array([[[f[j]]]]))[0,0])
self.assertAlmostEqual(np.mean(outputs),f.sum()/2)
self.assertNotEqual(np.mean(outputs),f.sum())
def test_approximate_control_integral_bias_counterexample(self):
f=np.array([.3,.8]);h=np.array([.1,.1]);wrong_integral=.7
expectation=wrong_integral+np.sum(f-h)
self.assertAlmostEqual(expectation-f.sum(),wrong_integral-h.sum())
def test_zero_support_is_rejected(self):
with self.assertRaises(ValueError):freeze(np.zeros((1,2,1)),np.array([[1.,0.]]))
def test_perfect_control_zero_variance(self):
f=np.arange(12,dtype=float).reshape(2,3,2)
s=freeze(f,np.full((2,3),1/3))
np.testing.assert_allclose(exact_mse(f,s,1),0)
j=draw(s,1,np.random.default_rng(1))
np.testing.assert_allclose(correct(s,j,f[np.arange(2)[:,None],j]),f.sum(1))
def test_signed_estimator_clipping_changes_mean(self):
f=np.array([[[0.],[1.]]]);h=np.array([[[1.],[0.]]]);s=freeze(h,np.array([[.5,.5]]))
outputs=[float(correct(s,np.array([[j]]),f[:,j:j+1,:])[0,0]) for j in range(2)]
self.assertEqual(outputs,[-1.,3.])
self.assertEqual(np.mean(outputs),1.)
self.assertGreater(np.maximum(outputs,0).mean(),1.)
def test_exactly_learning_one_term_can_increase_realized_variance(self):
f=np.ones((1,2,1));q=np.array([[.5,.5]])
before=exact_mse(f,freeze(np.zeros_like(f),q))[0]
h=np.array([[[1.],[0.]]]);after=exact_mse(f,freeze(h,q))[0]
self.assertEqual(before,0)
self.assertGreater(after,before)
def test_metric_psd_and_posterior_identity(self):
rng=np.random.default_rng(88);c=rng.normal(size=(3,7));q=rng.random(7);q/=q.sum()
g=residual_metric(c,q,n=2)
self.assertGreaterEqual(np.linalg.eigvalsh(g).min(),-1e-10)
x=rng.normal(size=(7,7));p=x@x.T
# Sum over independent spectral sources gives exact expectation.
vals,vecs=np.linalg.eigh(p);actual=0.
for i in range(7):
residual=c*vecs[:,i][None,:]*np.sqrt(max(vals[i],0))
actual+=(np.sum(residual**2/q)-np.sum(residual.sum(1)**2))/2
self.assertAlmostEqual(actual,float(np.trace(g@p)),places=9)
def test_future_query_risk_reduction(self):
rng=np.random.default_rng(90);a=rng.normal(size=(5,5));p=a@a.T
c=rng.normal(size=(3,5));g=residual_metric(c,np.full(5,.2));h=rng.normal(size=5);r=.3
ph=p@h;newp=p-np.outer(ph,ph)/(r+h@ph)
improvement=np.trace(g@(p-newp));formula=ph@g@ph/(r+h@ph)
self.assertAlmostEqual(float(improvement),float(formula),places=10)
def test_control_gauge_preserves_every_sample_output(self):
rng=np.random.default_rng(55);h=rng.normal(size=(3,6,2));q=rng.random((3,6));q/=q.sum(1,keepdims=True)
f=rng.random(h.shape);a=rng.normal(size=(3,2));h2=h+q[...,None]*a[:,None,:]
s=freeze(h,q);s2=freeze(h2,q)
for index in range(6):
j=np.full((3,1),index);physical=f[np.arange(3)[:,None],j]
np.testing.assert_allclose(correct(s,j,physical),correct(s2,j,physical),atol=1e-12)
zero_sum=h-q[...,None]*h.sum(1)[:,None,:]
np.testing.assert_allclose(zero_sum.sum(1),0,atol=1e-12)
def test_gauge_is_proposal_dependent(self):
h=np.array([[[.2],[.9]]]);q1=np.array([[.5,.5]]);q2=np.array([[.3,.7]])
h2=h+q1[...,None]*2
f=np.array([[[1.]]]);j=np.array([[0]])
self.assertNotAlmostEqual(correct(freeze(h,q2),j,f)[0,0],correct(freeze(h2,q2),j,f)[0,0])
def test_hoeffding_covers_enumerated_small_case(self):
f=np.array([[[0.],[1.]]]);h=np.array([[[.8],[.3]]]);s=freeze(h,np.array([[.5,.5]]))
radius=hoeffding_radius(s,np.ones_like(f),8,delta=.1)[0,0]
violation=0.
for seq in itertools.product(range(2),repeat=8):
ids=np.array([seq]);result=correct(s,ids,f[np.zeros_like(ids),ids])[0,0]
violation+=2**-8*(abs(result-1)>radius)
self.assertLessEqual(violation,.1)
def test_sampling_oracle_api_and_bad_inputs(self):
s=freeze(np.zeros((2,3,1)),np.full((2,3),1/3))
estimate,j,f=sample(s,lambda j:np.ones(j.shape+(1,)),3,np.random.default_rng(3))
np.testing.assert_allclose(estimate,3)
for n in (0,-1,True,1.2):
with self.assertRaises(ValueError):draw(s,n,np.random.default_rng(1))
with self.assertRaises(ValueError):correct(s,np.zeros((2,1),int),np.full((2,1,1),np.nan))
with self.assertRaises(ValueError):residual_metric(np.ones((2,3)),np.full(3,1/3),channel_metric=np.diag([1,-1]))
def test_proposal_full_support_in_all_trusted_case(self):
b=np.zeros((4,9,3));p=np.ones((4,9));trusted=np.ones((4,9),bool)
q=proposal_from_bound(b,p,trusted,active=True)
self.assertTrue((q>0).all());np.testing.assert_allclose(q.sum(1),1)
class MemoryAndOracle(unittest.TestCase):
def test_invalid_canonical_addresses_and_scene_inputs_rejected(self):
m=WorldMemory(2,3)
for bad in ([-1],[2],[.5]):
with self.assertRaises(ValueError):m.predict(bad,np.full((1,3),.5))
with self.assertRaises(ValueError):m.trusted(bad)
with self.assertRaises(ValueError):m.retain_only(bad)
with self.assertRaises(ValueError):Scene(0,np.zeros((3,4)))
scene=Scene.create(1)
with self.assertRaises(ValueError):scene.visibility(np.array([np.nan,0,0]),np.array([0,0,2]))
def test_exact_memory_survives_500_ticks(self):
m=WorldMemory(3,4);m.commit(np.array([1]),np.array([[0,3]]),np.array([[0,1]]));m.advance(500)
result=m.predict(np.array([1]),np.full((1,4),.5))
np.testing.assert_allclose(result,[[0,.5,.5,1]]);self.assertEqual(m.tick,500)
def test_repeated_ray_does_not_create_precision(self):
m=WorldMemory(2,3);rows=np.array([0]);ids=np.array([[1,1,1]])
m.commit(rows,ids,np.ones_like(ids));self.assertEqual(np.sum(m.trusted(rows)),1)
def test_geometry_revision_revokes_trust_preserves_fallible_value(self):
m=WorldMemory(2,3);m.commit(np.array([0]),np.array([[1]]),np.array([[1]]));m.notify_geometry_change()
self.assertFalse(m.trusted(np.array([0])).any())
self.assertEqual(m.predict(np.array([0]),np.full((1,3),.5))[0,1],1)
def test_observed_conflict_revokes_global_trust_after_batch(self):
m=WorldMemory(2,3);m.commit(np.array([0,1]),np.array([[1],[2]]),np.array([[1],[0]]))
conflicts=m.commit(np.array([0]),np.array([[1]]),np.array([[0]]),revise_on_conflict=True)
self.assertEqual(conflicts,1);self.assertEqual(m.epoch,1)
self.assertFalse(m.trusted(np.array([1])).any());self.assertTrue(m.trusted(np.array([0]))[0,1])
self.assertEqual(m.values[1,2],0)
def test_epoch_rollover_drops_unrepresentable_provenance(self):
m=WorldMemory(1,1);m.commit(np.array([0]),np.array([[0]]),np.array([[1]]))
m.epoch=np.iinfo(np.int32).max;m.notify_geometry_change()
self.assertTrue(np.isnan(m.values).all());self.assertEqual(m.epoch,0)
def test_screen_eviction_discards_hidden_evidence(self):
m=WorldMemory(3,2);m.commit(np.array([0]),np.array([[1]]),np.array([[1]]));m.retain_only(np.array([1]))
self.assertTrue(np.isnan(m.values[0]).all())
def test_roundtrip_and_namespace_guard(self):
m=WorldMemory(3,4,"room-v1");m.commit(np.array([0]),np.array([[2]]),np.array([[0]]));m.advance(500)
with tempfile.TemporaryDirectory() as temp:
path=Path(temp)/"memory.npz";m.save(path);new=WorldMemory.load(path,"room-v1")
np.testing.assert_equal(new.values,m.values);self.assertEqual(new.tick,500)
with self.assertRaises(ValueError):WorldMemory.load(path,"room-v2")
def test_malformed_checkpoint_rejected(self):
with tempfile.TemporaryDirectory() as temp:
path=Path(temp)/"bad.npz"
np.savez(path,namespace=np.array("a"),values=np.array([[.4]],np.float32),epochs=np.array([[0]],np.int32),epoch=np.int64(0),tick=np.int64(1))
with self.assertRaises(ValueError):WorldMemory.load(path,"a")
def test_physical_segment_hit_miss_and_beyond_emitter(self):
spheres=np.array([[0,0,1,.2],[10,10,1,.1],[-10,-10,1,.1]])
scene=Scene(0,spheres)
p=np.array([[0,0,0],[1,1,0],[0,0,0]])
l=np.array([[0,0,2],[1,1,2],[0,0,.5]])
np.testing.assert_equal(scene.visibility(p,l),[0,1,1])
def test_physical_lighting_linearity(self):
scene=Scene.create(201);p=receiver_grid(4,4);l=light_grid(3)
b=unoccluded(p,l);f=physical_table(scene,p,l,b)
np.testing.assert_allclose(physical_table(scene,p,l,2.3*b),2.3*f)
self.assertTrue((f>=0).all());self.assertTrue((f<=b+1e-15).all())
def test_portable_model_bounds_and_split(self):
root=Path(__file__).resolve().parents[1]
model=VisibilityPrior(root/"models/visibility_prior.npz")
p=receiver_grid(3,3);l=light_grid(2);scene=Scene.create(777)
result=model(scene.features(p[:,None,:],l[None,:,:]))
self.assertEqual(result.shape,(9,4));self.assertTrue(((result>=0)&(result<=1)).all())
import json
r=json.loads((root/"results/training.json").read_text())
self.assertFalse(set(r["training_scene_ids"])&set(r["test_scene_ids"]))
self.assertFalse(set(r["validation_scene_ids"])&set(r["test_scene_ids"]))
self.assertLess(r["numpy_torch_max_abs_error"],1e-5)
if __name__=="__main__":unittest.main()