diff options
| author | godosa <godosa@godosa.eu> | 2026-10-07 00:01:14 +0200 |
|---|---|---|
| committer | godosa <godosa@godosa.eu> | 2026-10-07 00:01:14 +0200 |
| commit | ed1dea2639b1191421de3986483aedcc14067a12 (patch) | |
| tree | 0118c6e119a84f1a9433ed0a304c8be1c9098463 /tests/test_state.py | |
| download | worldhistory-ed1dea2639b1191421de3986483aedcc14067a12.tar.gz worldhistory-ed1dea2639b1191421de3986483aedcc14067a12.zip | |
worldhistory: initial public history
Diffstat (limited to 'tests/test_state.py')
| -rw-r--r-- | tests/test_state.py | 63 |
1 files changed, 63 insertions, 0 deletions
diff --git a/tests/test_state.py b/tests/test_state.py new file mode 100644 index 0000000..6ed8351 --- /dev/null +++ b/tests/test_state.py @@ -0,0 +1,63 @@ +import unittest + +import numpy as np + +from worldhistory.state import convert, move, new_state + + +def st3(): + st = new_state(2, 3, np.array([[0.0, 1.0, 0, 0, 0, 0], [5.0, 1.0, 0, 0, 0, 0]]), seed=0) + st.P[0] = [10.0, 0.0, 30.0] + st.O[0, 0] = [100.0, 999.0, 300.0] # cell 1 is empty with a stale value + return st + + +class StateTest(unittest.TestCase): + def test_natives_need_one_column_per_condition(self): + from worldhistory.config import CONDITIONS + with self.assertRaisesRegex(ValueError, "conditions"): + new_state(1, 3, np.zeros((1, len(CONDITIONS) - 1)), seed=0) + + def test_shapes(self): + st = new_state(2, 3, np.zeros((2, 6)), seed=0) + self.assertEqual(st.P.shape, (2, 3)) + self.assertEqual(st.O.shape, (2, 6, 3)) + self.assertFalse(hasattr(st, "E")) + self.assertEqual(st.T.shape, (2, 6, 3)) + + def test_move_conserves_and_mixes(self): + st = st3() + move(st, 0, [0], [2], [10.0]) + np.testing.assert_allclose(st.P[0], [0, 0, 40]) + self.assertAlmostEqual(float(st.O[0, 0, 2]), (30 * 300 + 10 * 100) / 40) + + def test_move_into_empty_takes_arrivals_attrs(self): + st = st3() + move(st, 0, [0], [1], [4.0]) + self.assertAlmostEqual(float(st.O[0, 0, 1]), 100.0) + + def test_move_caps_at_available(self): + st = st3() + move(st, 0, [0, 0], [1, 2], [15.0, 5.0]) # asks 20 of 10: scaled to 7.5 + 2.5 + np.testing.assert_allclose(st.P[0], [0, 7.5, 32.5]) + self.assertAlmostEqual(float(st.P[0].sum()), 40.0) + + def test_convert_between_slots(self): + st = st3() + st.P[1, 2] = 10.0 + st.O[1, 0, 2] = 0.0 + convert(st, 0, 1, np.array([2]), np.array([10.0])) + self.assertAlmostEqual(float(st.P[0, 2]), 20.0) + self.assertAlmostEqual(float(st.P[1, 2]), 20.0) + self.assertAlmostEqual(float(st.O[1, 0, 2]), 150.0) # (10*0 + 10*300) / 20 + + def test_convert_with_new_optima(self): + st = new_state(2, 3, np.zeros((2, 6)), seed=0) + st.P[0, 2] = 30.0 + st.P[1, 2] = 10.0 + st.O[1, 0, 2] = 100.0 + O_new = np.zeros((6, 1)) + O_new[0, 0] = 500.0 + convert(st, 0, 1, [2], [30.0], O_new=O_new) + self.assertAlmostEqual(float(st.O[1, 0, 2]), (10 * 100 + 30 * 500) / 40) + self.assertEqual(float(st.O[0, 0, 2]), 0.0) # the parents' optima are untouched |
