aboutsummaryrefslogtreecommitdiffziptar.gz
path: root/tests/test_pipeline.py
blob: adbe3c4d3fef9cbc2115a25d74bef287b5bf4bea (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
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")