aboutsummaryrefslogtreecommitdiffziptar.gz
path: root/tests/test_pipeline.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/test_pipeline.py')
-rw-r--r--tests/test_pipeline.py185
1 files changed, 185 insertions, 0 deletions
diff --git a/tests/test_pipeline.py b/tests/test_pipeline.py
new file mode 100644
index 0000000..adbe3c4
--- /dev/null
+++ b/tests/test_pipeline.py
@@ -0,0 +1,185 @@
+import shutil
+import tempfile
+import unittest
+from pathlib import Path
+
+import numpy as np
+
+from mapgen import pipeline as P
+from mapgen.testing import FIXTURE_TOML
+from tests.helpers import ROOT, make_ctx
+
+TECT = """
+[[plate]]
+id = "a"
+seed = [0.0, 0.0]
+kind = "continental"
+motion = [90.0, 3.0]
+[[plate]]
+id = "b"
+seed = [0.0, 90.0]
+kind = "oceanic"
+motion = [270.0, 3.0]
+"""
+
+
+class PipelineTest(unittest.TestCase):
+ def setUp(self):
+ self.tmp = Path(tempfile.mkdtemp())
+ (self.tmp / "config").mkdir()
+ shutil.copy(FIXTURE_TOML, self.tmp / "config" / "world.toml")
+ (self.tmp / "config" / "tectonics.toml").write_text(TECT)
+ self.log = []
+
+ def tearDown(self):
+ shutil.rmtree(self.tmp)
+
+ def _build(self, **kw):
+ self.log = []
+ return P.build(self.tmp, 1, stop="grid", log=self.log.append, **kw)
+
+ def test_cache_hit(self):
+ ctx = self._build()
+ self.assertEqual(ctx.grid.n, 842)
+ self.assertFalse(any("cached" in m for m in self.log))
+ ctx = self._build()
+ self.assertTrue(any("grid: cached" in m for m in self.log))
+ self.assertEqual(ctx.grid.n, 842)
+
+ def test_config_change_invalidates(self):
+ self._build()
+ p = self.tmp / "config" / "world.toml"
+ p.write_text(p.read_text().replace("seed = 1296", "seed = 1297"))
+ self._build()
+ self.assertFalse(any("cached" in m for m in self.log))
+
+ def test_from_forces_rerun(self):
+ self._build()
+ self._build(start="grid")
+ self.assertFalse(any("cached" in m for m in self.log))
+
+ def test_need_reports_missing(self):
+ ctx = make_ctx(1)
+ with self.assertRaisesRegex(P.StageError, "sk_land"):
+ ctx.need("sk_land")
+
+
+class OceanExportTest(unittest.TestCase):
+ def test_cells_and_fields_carry_the_ocean(self):
+ import json
+ from mapgen.testing import built_world
+ root = built_world()
+ outs = sorted((root / "out").glob("r*/cells.npz"))
+ self.assertTrue(outs)
+ for cells in outs + sorted((root / "out").glob("r*/eras/*/cells.npz")): # every era re-runs climate
+ z = np.load(cells)
+ for k in ("current", "current_speed", "sst", "upwelling", "productivity"):
+ self.assertIn(k, z.files, f"{cells}: {k}")
+ meta = json.loads((outs[0].parent / "fields.json").read_text())
+ for k in ("sst", "productivity", "current_speed", "upwelling"):
+ self.assertIn(k, meta["continuous"])
+ layers = [L["id"] for L in json.loads((outs[0].parent / "viewer" / "layers.json").read_text())]
+ self.assertIn("currents", layers)
+
+
+class NpzMapsTest(unittest.TestCase):
+ def test_same_arrays_as_np_load_writable_and_file_untouched(self):
+ import hashlib
+ import tempfile
+ import numpy as np
+ from pathlib import Path
+ from mapgen import pipeline as P
+ rng = np.random.default_rng(0)
+ arrays = {"f": rng.normal(size=1000), "i": rng.integers(0, 9, (40, 3)).astype(np.int16),
+ "b": rng.random(77) < 0.5, "F": np.asfortranarray(rng.random((30, 4))), "s": np.array(2.5),
+ "e": np.zeros((0, 3), np.float32), "_key": np.array("abc"), "u": np.arange(5, dtype=np.uint64),
+ "big": rng.normal(size=(3_000_000,)), "nc": rng.random((50, 8))[:, ::3], "x" * 200: np.arange(7)}
+ with tempfile.TemporaryDirectory() as t:
+ for save in (np.savez, np.savez_compressed, P.save_npz_aligned):
+ f = Path(t) / f"{save.__name__}.npz"
+ save(f, **arrays)
+ before = hashlib.sha256(f.read_bytes()).hexdigest()
+ got = P.npz_maps(f)
+ if save is P.save_npz_aligned: # every member mapped in place, at aligned addresses
+ for k, v in got.items():
+ if v.size:
+ self.assertEqual(v.ctypes.data % P.ALIGN, 0, k)
+ b = v
+ while isinstance(b, (np.ndarray, memoryview)):
+ b = b.base if isinstance(b, np.ndarray) else b.obj
+ self.assertIsInstance(b, __import__("mmap").mmap, k)
+ with np.load(f) as z:
+ self.assertEqual(sorted(got), sorted(z.files))
+ for k in z.files:
+ self.assertEqual(got[k].dtype, z[k].dtype, k)
+ self.assertEqual(got[k].shape, z[k].shape, k)
+ self.assertTrue(np.array_equal(got[k], z[k]), k)
+ self.assertTrue(np.array_equal(got[k], arrays[k]), k)
+ for k, v in got.items(): # whatever is mapped is aligned (else: read)
+ b = v
+ while isinstance(b, (np.ndarray, memoryview)):
+ b = b.base if isinstance(b, np.ndarray) else b.obj
+ if isinstance(b, __import__("mmap").mmap):
+ self.assertEqual(v.ctypes.data % P.ALIGN, 0, (save.__name__, k))
+ got["f"][:] = -1.0 # private copy-on-write: allowed, file unchanged
+ got["F"][0, 0] = 9.0
+ self.assertEqual(float(got["F"][0, 0]), 9.0)
+ self.assertEqual(hashlib.sha256(f.read_bytes()).hexdigest(), before, save.__name__)
+ self.assertTrue(np.array_equal(np.load(f)["f"], arrays["f"]))
+
+ def test_low_memory_rebuild_from_cache_identical(self):
+ import shutil
+ import tempfile
+ from pathlib import Path
+ from mapgen import pipeline as P
+ from mapgen.testing import small_world
+ outs = []
+ tmp = Path(tempfile.mkdtemp())
+ try:
+ small_world(tmp)
+ for low in (False, True, True): # normal; low (fresh cache hit); low again
+ ctx = P.build(tmp, 2, log=lambda m: None, low_memory=low)
+ outs.append({k: v for k, v in ctx.data.items()})
+ for k, v in outs[0].items():
+ for o in outs[1:]:
+ self.assertTrue(np.array_equal(np.asarray(o[k]), np.asarray(v), equal_nan=np.asarray(v).dtype.kind == "f"), k)
+ self.assertFalse(list((tmp / "out" / "cache" / "r2").glob("*.tmp-*")), "no temporary files left")
+ import mmap
+ ctx = P.build(tmp, 2, start="climate", log=lambda m: None, low_memory=True) # freshly computed
+ for k in ("T_jun", "P_ann", "z_surface_m", "plate"):
+ b = ctx.data[k]
+ while isinstance(b, (np.ndarray, memoryview)):
+ b = b.base if isinstance(b, np.ndarray) else b.obj
+ self.assertIsInstance(b, mmap.mmap, f"{k}: kept on disk, not in RAM")
+ self.assertTrue(np.array_equal(ctx.data[k], outs[0][k], equal_nan=True), k)
+ finally:
+ shutil.rmtree(tmp)
+
+
+class AlignedSumTest(unittest.TestCase):
+ def test_misaligned_data_sums_differently(self):
+ """Why the cache is aligned: numpy sums misaligned float64 data in another order (else drop the alignment)."""
+ import numpy as np
+ x = np.random.default_rng(0).normal(size=2_000_000)
+ buf = bytearray(8 * len(x) + 8)
+ a = np.frombuffer(memoryview(buf)[4:4 + 8 * len(x)], np.float64)
+ np.copyto(a, x)
+ self.assertFalse(a.flags.aligned)
+ self.assertTrue(np.array_equal(a, x))
+ if a.sum() == x.sum():
+ self.skipTest("this numpy sums misaligned data the same way")
+ self.assertNotEqual(a.sum(), x.sum())
+
+
+class BlasSpinTest(unittest.TestCase):
+ def test_import_sets_a_short_blas_spin_unless_given(self):
+ import os
+ import subprocess
+ import sys
+ code = "import os, mapgen, numpy; print(os.environ['OPENBLAS_THREAD_TIMEOUT'])"
+ env = {k: v for k, v in os.environ.items() if k != "OPENBLAS_THREAD_TIMEOUT"}
+ out = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True, env=env, cwd=os.getcwd())
+ self.assertEqual(out.stdout.strip(), "4")
+ out = subprocess.run([sys.executable, "-c", code], capture_output=True, text=True,
+ env={**env, "OPENBLAS_THREAD_TIMEOUT": "9"}, cwd=os.getcwd())
+ self.assertEqual(out.stdout.strip(), "9")