raw · 7702 bytes
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 | """Stage runner with an input-hash cache.""" from __future__ import annotations import hashlib import importlib import os import time from dataclasses import dataclass, field from pathlib import Path import numpy as np from . import config as C from .grid import Grid STAGES = ["grid", "sketch", "plates", "crust", "elevation", "erosion", "climate", "hydrology", "seabed", "environment", "ice", "fields", "render"] class StageError(RuntimeError): pass @dataclass class Ctx: root: Path cfg: dict tect: dict res: int data: dict = field(default_factory=dict) grid: Grid | None = None out_dir: Path | None = None # render: out/r<res> unless set (an era writes out/r<res>/eras/<era>) preview_dir: Path | None = None low_memory: bool = False # trade speed for a lower memory peak; never changes the results @property def seed(self) -> int: return int(self.cfg["build"]["seed"]) def need(self, *keys): missing = [k for k in keys if k not in self.data] if missing: raise StageError(f"missing inputs {missing}: run the stage that produces them first") return [self.data[k] for k in keys] def inputs_key(root: Path, res: int) -> str: h = hashlib.sha256(str(res).encode()) files = [] for sub, pat in (("config", "*.toml"), ("masks", "*.png"), ("sketch", "*.png")): files += sorted(p for p in (root / sub).glob(pat) if p.name != C.ERAS_FILE) # eras: their own key files += sorted(Path(__file__).resolve().parent.glob("*.py")) for p in files: h.update(p.name.encode()) h.update(p.read_bytes()) return h.hexdigest()[:16] ALIGN = 64 def save_npz_aligned(path, /, **arrays) -> None: """np.savez(path, **arrays), but each array's data starts at a multiple of ALIGN bytes in the file (the local zip header gets a padding extra field, as zipalign does): an ordinary .npz for np.load, whose arrays npz_maps can map in place. Alignment matters for more than speed: numpy sums 8-byte-misaligned data in buffered chunks, i.e. in another order — mapped misaligned inputs would change the last bits of results.""" import io import struct import zipfile from numpy.lib import format as F with zipfile.ZipFile(path, "w", compression=zipfile.ZIP_STORED, allowZip64=True) as zf: for name, v in arrays.items(): v = np.asarray(v) if v.dtype.hasobject: raise ValueError(f"save_npz_aligned: {name} has object dtype") head = io.BytesIO() d = F.header_data_from_array_1_0(v) try: F.write_array_header_1_0(head, d) except ValueError: head = io.BytesIO() F.write_array_header_2_0(head, d) info = zipfile.ZipInfo(f"{name}.npy", date_time=(1980, 1, 1, 0, 0, 0)) info.compress_type = zipfile.ZIP_STORED start = zf.fp.tell() + 30 + len(info.filename.encode()) + 20 + len(head.getvalue()) # 20: zip64 field pad = -start % ALIGN if 0 < pad < 4: # an extra field is at least its 4-byte header pad += ALIGN if pad: info.extra = struct.pack("<HH", 0xA1A1, pad - 4) + bytes(pad - 4) with zf.open(info, "w", force_zip64=True) as m: m.write(head.getvalue()) _write_data(m, v) def _write_data(m, v) -> None: """The array's bytes (C order, or Fortran order for an F-contiguous array, as np.save) in 16 MB pieces.""" flat = v.T.reshape(-1) if (v.flags.f_contiguous and not v.flags.c_contiguous) else np.ascontiguousarray(v).reshape(-1) step = max(1, (16 << 20) // max(v.itemsize, 1)) for i in range(0, flat.size, step): m.write(flat[i:i + step].tobytes()) def npz_maps(f: Path) -> dict: """The arrays of an uncompressed .npz (np.savez) mapped from the file, copy-on-write: plain writable arrays with the same values, whose pages the OS reads on use and can drop again (writes stay private, the file never changes). Members that can't be mapped (compressed, object dtype) are read into memory as np.load would.""" import mmap import struct import zipfile out = {} with open(f, "rb") as fh, zipfile.ZipFile(fh) as zf: mm = mmap.mmap(fh.fileno(), 0, access=mmap.ACCESS_COPY) for info in zf.infolist(): name = info.filename[:-4] if info.filename.endswith(".npy") else info.filename ok = info.compress_type == zipfile.ZIP_STORED if ok: fh.seek(info.header_offset) local = fh.read(30) n_name, n_extra = struct.unpack("<HH", local[26:30]) fh.seek(info.header_offset + 30 + n_name + n_extra) version = np.lib.format.read_magic(fh) read_header = {(1, 0): np.lib.format.read_array_header_1_0, (2, 0): np.lib.format.read_array_header_2_0}.get(version) ok = read_header is not None if ok: shape, fortran, dtype = read_header(fh, max_header_size=1 << 20) ok = not dtype.hasobject ok = fh.tell() % ALIGN == 0 # misaligned: read (see save_npz_aligned) if ok: count = int(np.prod(shape, dtype=np.int64)) a = np.frombuffer(mm, dtype=dtype, count=count, offset=fh.tell()) if count else np.empty(0, dtype) out[name] = a.reshape(shape, order="F" if fortran else "C") else: with zf.open(info) as m: out[name] = np.lib.format.read_array(m, allow_pickle=False) return out def _after(ctx: Ctx, name: str) -> None: if name == "grid": ctx.grid = Grid.from_arrays(ctx.data, ctx.res, float(ctx.cfg["planet"]["radius_km"])) def build(root: Path, res: int, start: str | None = None, stop: str | None = None, log=print, low_memory: bool = False) -> Ctx: cfg, tect = C.load(root) ctx = Ctx(root, cfg, tect, res, low_memory=low_memory) cache = root / "out" / "cache" / f"r{res}" cache.mkdir(parents=True, exist_ok=True) key = inputs_key(root, res) stages = STAGES[: STAGES.index(stop) + 1] if stop else STAGES forced = False for name in stages: forced = forced or name == start f = cache / f"{name}.npz" t0 = time.time() if not forced and f.exists(): with np.load(f, allow_pickle=False) as z: hit = str(z["_key"]) == key if hit and not ctx.low_memory: ctx.data.update({k: z[k] for k in z.files if k != "_key"}) if hit: if ctx.low_memory: ctx.data.update({k: v for k, v in npz_maps(f).items() if k != "_key"}) _after(ctx, name) log(f"{name}: cached") continue forced = True out = importlib.import_module(f"mapgen.{name}").run(ctx) tmp = f.with_name(f"{f.stem}.tmp-{os.getpid()}.npz") save_npz_aligned(tmp, _key=np.array(key), **out) os.replace(tmp, f) # never truncated in place: maps of the old file stay valid if ctx.low_memory: # fields kept on disk from here on (the cache file just written) maps = npz_maps(f) out = {k: maps.get(k, v) if isinstance(v, np.ndarray) else v for k, v in out.items()} ctx.data.update(out) _after(ctx, name) log(f"{name}: {time.time() - t0:.1f}s") return ctx |