aboutsummaryrefslogtreecommitdiffziptar.gz
path: root/tests/test_state.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/test_state.py')
-rw-r--r--tests/test_state.py63
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