aboutsummaryrefslogtreecommitdiffziptar.gz
path: root/tests/test_world.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/test_world.py')
-rw-r--r--tests/test_world.py84
1 files changed, 84 insertions, 0 deletions
diff --git a/tests/test_world.py b/tests/test_world.py
new file mode 100644
index 0000000..a9969f8
--- /dev/null
+++ b/tests/test_world.py
@@ -0,0 +1,84 @@
+import json
+import tempfile
+import unittest
+from pathlib import Path
+
+import numpy as np
+
+from tests.helpers import fields, globe_world, line_world
+from worldhistory.world import FIELDS, load_world
+
+
+class WorldTest(unittest.TestCase):
+ def test_line_neighbours_and_sum(self):
+ w = line_world(4)
+ np.testing.assert_array_equal(w.nb_sum(np.array([1.0, 2.0, 3.0, 4.0])), [2, 4, 6, 3])
+ np.testing.assert_array_equal(w.nb_any(np.array([True, False, False, False])), [False, True, False, False])
+
+ def test_nb_sum_on_stacked_arrays(self):
+ w = line_world(3)
+ x = np.array([[1.0, 0, 0], [0, 0, 5.0]])
+ np.testing.assert_array_equal(w.nb_sum(x), [[0, 1, 0], [0, 5, 0]])
+
+ def test_smooth_keeps_constant(self):
+ w = globe_world(1)
+ np.testing.assert_allclose(w.smooth(np.full(w.n, 3.0), 2), 3.0)
+
+ def test_noise_standardised_and_seeded(self):
+ w = globe_world(1)
+ a, b = w.noise(5), w.noise(5)
+ np.testing.assert_array_equal(a, b)
+ self.assertAlmostEqual(float(a.mean()), 0.0, places=6)
+ self.assertAlmostEqual(float(a.std()), 1.0, places=6)
+ self.assertFalse(np.array_equal(a, w.noise(6)))
+
+ def test_globe_neighbours_are_close(self):
+ w = globe_world(1)
+ valid = w.nb >= 0
+ self.assertTrue((valid.sum(1) >= 5).all()) # hexagons 6, pentagons 5
+ i = np.repeat(np.arange(w.n), 6)[valid.ravel()]
+ j = w.nb.ravel()[valid.ravel()]
+ self.assertLess(float(w.km(i, j).max()), 1000.0) # res-1 cells are ~400-600 km apart on an Earth-size globe
+
+ def test_km_and_within(self):
+ w = line_world(10, spacing_km=100.0)
+ self.assertAlmostEqual(float(w.km(np.array([0]), np.array([3]))[0]), 300.0, places=3)
+ self.assertEqual(sorted(w.within(5, 150.0).tolist()), [4, 5, 6])
+
+ def test_era_switch(self):
+ w = line_world(3)
+ w.eras["later"] = fields(3, ocean=True)
+ self.assertFalse(w.fields["ocean"].any())
+ w.set_era("later")
+ self.assertTrue(w.fields["ocean"].all())
+ with self.assertRaises(KeyError):
+ w.set_era("nope")
+
+ def test_load_world_roundtrip(self):
+ g = globe_world(1)
+ import h3.api.basic_int as h3
+ ids = np.array(sorted(c for r0 in h3.get_res0_cells() for c in h3.cell_to_children(r0, 1)), np.uint64)
+ with tempfile.TemporaryDirectory() as d:
+ base = {f"g_{k}": v for k, v in dict(ids=ids, lat=g.lat, lon=g.lon, xyz=g.xyz * 2.0,
+ area_km2=g.area).items()}
+ np.savez(Path(d) / "cells.npz", **base, **fields(g.n))
+ (Path(d) / "eras" / "late").mkdir(parents=True)
+ np.savez(Path(d) / "eras" / "late" / "cells.npz", **fields(g.n, ocean=True))
+ (Path(d) / "cells_meta.json").write_text(json.dumps({"radius_km": 12742.0}))
+ w = load_world(d, eras=["late"])
+ self.assertEqual(w.radius_km, 12742.0)
+ np.testing.assert_allclose(np.linalg.norm(w.xyz, axis=1), 1.0)
+ np.testing.assert_array_equal(w.nb, g.nb)
+ self.assertTrue(set(FIELDS) <= set(w.base))
+ self.assertTrue(w.eras["late"]["ocean"].all())
+
+ def test_load_world_missing_field(self):
+ g = globe_world(1)
+ import h3.api.basic_int as h3
+ ids = np.array(sorted(c for r0 in h3.get_res0_cells() for c in h3.cell_to_children(r0, 1)), np.uint64)
+ with tempfile.TemporaryDirectory() as d:
+ f = fields(g.n)
+ del f["gravity_g"]
+ np.savez(Path(d) / "cells.npz", g_ids=ids, g_lat=g.lat, g_lon=g.lon, g_xyz=g.xyz, g_area_km2=g.area, **f)
+ with self.assertRaisesRegex(ValueError, "gravity_g"):
+ load_world(d)