worldgen

git clone https://git.godosa.eu/worldgen

master

raw ยท 8405 bytes

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")