File size: 11,525 Bytes
9d6c005
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
"""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()