worldhistory

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

master

raw · 14672 bytes

"""The speed-ups that compute only where values can change give the same floats, bit for bit, as the full-array
versions they replaced (kept here as oracles)."""
import unittest

import numpy as np

from worldhistory import state as S
from worldhistory.state import ATTRS, new_state


def move_full(st, r, src, dst, amt):
    """The original move: every cell updated."""
    src, dst, amt = np.asarray(src, np.int64), np.asarray(dst, np.int64), np.asarray(amt, float)
    keep = amt > 0
    src, dst, amt = src[keep], dst[keep], amt[keep]
    if not len(amt):
        return
    n = st.P.shape[1]
    P = st.P[r]
    out = np.bincount(src, amt, n)
    scale = np.where(out > P, P / np.maximum(out, 1e-12), 1.0)
    amt = amt * scale[src]
    stay = np.maximum(P - np.bincount(src, amt, n), 0.0)
    inflow = np.bincount(dst, amt, n)
    tot = stay + inflow
    for name in ATTRS:
        A = getattr(st, name)[r]
        if A.shape[0] == 0:
            continue
        Sm = np.stack([np.bincount(dst, amt * A[x, src], n) for x in range(A.shape[0])])
        A[:] = np.where(tot > 0, (stay * A + Sm) / np.maximum(tot, 1e-12), A)
    st.P[r] = tot


def random_state(n=4000, seed=0):
    rng = np.random.default_rng(seed)
    st = new_state(2, n, np.full((2, 6), 10.0, np.float32), seed)
    st.P[:] = np.where(rng.random((2, n)) < 0.4, rng.random((2, n)) * 1e4, 0.0)
    st.P[0, :20] = rng.random(20) * 1e-13               # thin groups (below the 1e-12 floor)
    st.O[:] = (rng.normal(10, 5, st.O.shape)).astype(np.float32)
    st.T[:] = rng.random(st.T.shape).astype(np.float32)
    st.C[:, 0] = np.where(rng.random((2, n)) < 0.1, rng.random((2, n)) * 0.3, 0.0)   # curses: float64, nonzero
    return st, rng


class MoveTest(unittest.TestCase):
    def test_move_matches_full_update(self):
        for seed in range(4):
            a, rng = random_state(seed=seed)
            b, _ = random_state(seed=seed)
            n = a.P.shape[1]
            for _ in range(5):
                k = int(rng.integers(1, 3000))
                src, dst = rng.integers(0, n, k), rng.integers(0, n, k)
                amt = rng.random(k) * 5e3 * (rng.random(k) < 0.9)
                S.move(a, 0, src, dst, amt)
                move_full(b, 0, src, dst, amt)
            for name in ("P", *ATTRS):
                with self.subTest(seed=seed, attr=name):
                    self.assertTrue(np.array_equal(getattr(a, name), getattr(b, name)))


class MoveSumOrderTest(unittest.TestCase):
    def test_many_moves_into_few_cells_sum_in_move_order(self):
        """Thousands of moves into a few cells, float64 curses: any other summation order shows in the last bits."""
        for seed in range(3):
            a, rng = random_state(seed=seed)
            b, _ = random_state(seed=seed)
            n = a.P.shape[1]
            c = rng.random(n) * 0.4
            for st in (a, b):
                st.C[0, 0] = c
            k = 20000
            src, dst = rng.integers(0, n, k), rng.integers(0, 50, k)
            amt = np.exp(rng.normal(0, 3, k))
            S.move(a, 0, src, dst, amt)
            move_full(b, 0, src, dst, amt)
            for name in ("P", *ATTRS):
                with self.subTest(seed=seed, attr=name):
                    self.assertTrue(np.array_equal(getattr(a, name), getattr(b, name)))


class ConflictTest(unittest.TestCase):
    def test_conflict_matches_full_arrays_many_races(self):
        """≥ 8 races: numpy sums a column-major block pairwise — the subset must still add race by race."""
        from tests.helpers import make_race
        from worldhistory.conflict import conflict
        rng = np.random.default_rng(9)
        races = [make_race(id=f"r{k}", conflict={"aggression": rng.random(), "power": 0.5 + rng.random(),
                                                "dread": rng.random() * 0.5, "defend": 0.5, "border": 0.1,
                                                "curse": 0.3 * (k % 3 == 0)}, family=f"f{k % 5}") for k in range(9)]
        n = 20000
        P = np.where(rng.random((9, n)) < 0.25, np.exp(rng.normal(5, 3, (9, n))), 0.0)
        q, crowd = rng.random((9, n)), rng.random((9, n)) * 3
        blame, want_blame = np.zeros_like(P), np.zeros_like(P)
        loss, press = conflict(P, q, crowd, races, blame=blame)
        want_loss, want_press = conflict_full(P, q, crowd, races, blame=want_blame)
        self.assertTrue(np.array_equal(loss, want_loss))
        self.assertTrue(np.array_equal(press, want_press))
        self.assertTrue(np.array_equal(blame, want_blame))

    def test_conflict_matches_full_arrays(self):
        from tests.helpers import make_race
        from worldhistory.conflict import conflict
        rng = np.random.default_rng(5)
        races = [make_race(id=k, conflict={"aggression": a, "power": p, "dread": d, "defend": 0.5, "border": 0.1,
                                          "curse": c}, family=f)
                 for k, a, p, d, c, f in (("a", 0.4, 1.0, 0.2, 0.0, "x"), ("b", 0.9, 1.5, 0.0, 0.5, "y"),
                                          ("c", 0.2, 0.7, 0.6, 0.0, "z"))]
        n = 5000
        P = np.where(rng.random((3, n)) < 0.3, rng.random((3, n)) * 100, 0.0)
        P[2] = 0.0                                       # an absent race
        q, crowd = rng.random((3, n)), rng.random((3, n)) * 3
        blame = np.zeros_like(P)
        loss, press = conflict(P, q, crowd, races, blame=blame)
        want_blame = np.zeros_like(P)
        want_loss, want_press = conflict_full(P, q, crowd, races, blame=want_blame)
        self.assertTrue(np.array_equal(loss, want_loss))
        self.assertTrue(np.array_equal(press, want_press))
        self.assertTrue(np.array_equal(blame, want_blame))


def conflict_full(P, q_eff, crowd, races, core_q=0.6, blame=None):
    R = len(races)
    tot = P.sum(0)
    share = np.where(tot > 0, P / np.maximum(tot, 1e-12), 0.0)
    fam = [r.family for r in races]
    C = [r.conflict for r in races]
    loss, press = np.zeros_like(P), np.zeros_like(P)
    for i in range(R):
        ci = C[i]
        loss[i] += ci["internal"] * np.minimum(crowd[i], 2.0) * P[i]
        for j in range(R):
            if fam[j] == fam[i]:
                continue
            cj = C[j]
            attacked = cj["aggression"] * (1 - ci["dread"]) * cj["power"] / ci["power"] * share[j]
            attacking = ci["aggression"] * cj["defend"] * (q_eff[j] >= core_q) * cj["power"] / ci["power"] * share[j]
            press[i] += attacked + cj["dread"] * share[j]
            loss[i] += ci["border"] * (attacked + attacking) * P[i]
            if blame is not None and ci["curse"] > 0:
                blame[j] += ci["curse"] * ci["border"] * attacked
    return np.minimum(loss, 0.9 * P), press


def prospective_full(st, world):
    """The original prospective: neighbour sums over every cell."""
    for r in range(st.P.shape[0]):
        pos = st.P[r] > 0
        if not pos.any():
            continue
        P = np.where(pos, st.P[r], 0.0)
        wsum = world.nb_sum(P)
        empty = np.flatnonzero((st.P[r] <= 0) & (wsum > 0))
        if not len(empty):
            continue
        st.O[r][:, empty] = world.nb_sum(P * st.O[r])[:, empty] / np.maximum(wsum[empty], 1e-12)


def step_tech_full(st, world, races, tp, steps=1.0):
    """The original step_tech: neighbour sums over every cell."""
    from worldhistory.config import DOMAINS
    from worldhistory.tech import neigh_pop
    for r, race in enumerate(races):
        P = st.P[r]
        occ = P >= 1
        if not occ.any():
            continue
        oc = np.flatnonzero(occ)
        N = neigh_pop(world, P, tp["passes"])[oc]
        s = N / (N + tp["n_half"])
        T = st.T[r]
        for d, name in enumerate(DOMAINS):
            Td = T[d, oc]
            gain = tp["rate"] * steps * race.tech.get(name, 1.0) * s * (1 - Td)
            loss = np.where(N < tp["loss_below"], tp["loss_rate"] * steps * Td, 0.0)
            T[d, oc] = np.clip(Td + gain - loss, 0, 1)
        if tp["diffuse"] > 0:
            PT = P * T
            Pc, Tc = P[oc], T[:, oc]
            M = (PT[:, oc] + world.nb_sum(PT)[:, oc]) / np.maximum(Pc + world.nb_sum(P)[oc], 1e-12)
            T[:, oc] = Tc + tp["diffuse"] * np.clip(M - Tc, 0, None)


def sparse_state(n, seed, frac):
    """Groups on a few patches (as in a run: most cells empty), thin and fractional groups among them."""
    st, rng = random_state(n, seed)
    st.P[:] = np.where(rng.random((2, n)) < frac, rng.random((2, n)) * 3e3, 0.0)
    st.P[0, :5] = rng.random(5) * 0.5
    return st


class NeighbourhoodTest(unittest.TestCase):
    def setUp(self):
        from tests.helpers import globe_world
        self.w = globe_world(2)                                         # 5882 cells

    def test_adjacency_is_symmetric(self):
        a = self.w._adj
        self.assertEqual((a != a.T).nnz, 0)

    def test_nb_sum_at_matches_full_sum(self):
        rng = np.random.default_rng(3)
        x = rng.normal(size=(3, self.w.n))
        rows = np.sort(rng.choice(self.w.n, 400, replace=False))
        self.assertTrue(np.array_equal(self.w.nb_sum_at(x, rows), self.w.nb_sum(x)[:, rows]))
        self.assertTrue(np.array_equal(self.w.nb_sum_at(x[0], rows), self.w.nb_sum(x[0])[rows]))

    def test_prospective_matches_full(self):
        from worldhistory.adaptation import prospective
        for seed, frac in ((0, 0.02), (1, 0.2), (2, 0.0005), (3, 0.9)):
            a, b = sparse_state(self.w.n, seed, frac), sparse_state(self.w.n, seed, frac)
            prospective(a, self.w)
            prospective_full(b, self.w)
            with self.subTest(seed=seed):
                self.assertTrue(np.array_equal(a.O, b.O))
                self.assertFalse(np.array_equal(a.O, sparse_state(self.w.n, seed, frac).O))   # it did change O

    def test_step_tech_matches_full(self):
        from tests.helpers import make_history, make_race
        from worldhistory.tech import step_tech
        tp = make_history().tech
        races = [make_race("a"), make_race("b", tech={"farming": 2.0})]
        for seed, frac in ((0, 0.02), (1, 0.3), (2, 0.9)):
            for passes in (1, 2):
                a, b = sparse_state(self.w.n, seed, frac), sparse_state(self.w.n, seed, frac)
                for _ in range(3):
                    step_tech(a, self.w, races, {**tp, "passes": passes}, 1.5)
                    step_tech_full(b, self.w, races, {**tp, "passes": passes}, 1.5)
                with self.subTest(seed=seed, passes=passes):
                    self.assertTrue(np.array_equal(a.T, b.T))


def inherit_full(st, world, races, rules, births, regions, ok):
    """The original inherit: neighbour sums and gates over every cell."""
    from worldhistory.lineage import _born, convert, native, progress, smoothstep
    for rule in _born(rules):
        s, p = rule["s"], rule["p"]
        spec = races[s].emerge
        Ps, Pp = st.P[s], st.P[p]
        near_s = Ps + world.nb_sum(Ps)
        near = near_s + Pp + world.nb_sum(Pp)
        share = np.where(near > 0, near_s / np.maximum(near, 1e-12), 0.0)
        where = (Pp > 0) & (births[p] > 0) & (near_s > 0) & ok[s]
        if spec["spread"] == "region" and spec["region"]:
            where &= regions(spec["region"])
        if spec["mode"] == "ritual":
            gate = np.ones(world.n)
        else:
            a = progress(st.O[p], races[p], races[s])
            gate = smoothstep((a - spec["mix_min"]) / max(spec["birth_sure"] - spec["mix_min"], 1e-12))
        amt = np.where(where, spec["dominance"] * gate * share * births[p], 0.0)
        cells = np.flatnonzero(amt > 0)
        if len(cells):
            convert(st, p, s, cells, amt[cells], O_new=np.repeat(native(races[s])[:, None], len(cells), 1))


class InheritFastTest(unittest.TestCase):
    def test_inherit_matches_full(self):
        from tests.helpers import globe_world, make_history, make_race
        from tests.test_lineage import D, PARENT_TOL, SUB_TOL
        from worldhistory.config import link
        from worldhistory.habitat import native_optima
        from worldhistory.lineage import inherit, init_rules
        from worldhistory.regions import Regions
        w = globe_world(2)
        for seed, emerge in enumerate(({}, {"mode": "ritual"}, {"spread": "region", "region": "north"})):
            races = link(make_history(regions={"north": {"kind": "box", "lat": [0, 90], "lon": [-180, 180]}}), [make_race("p", tolerance=PARENT_TOL),
                                          make_race("s", tolerance=SUB_TOL, emerge={"parent": "p", **emerge})])
            regs = Regions(w, {"north": {"kind": "box", "lat": [0, 90], "lon": [-180, 180]}})
            out = []
            for fn in (inherit, inherit_full):
                rng = np.random.default_rng(seed)
                st = new_state(2, w.n, native_optima(races), seed=0)
                st.P[:] = np.where(rng.random((2, w.n)) < 0.1, rng.random((2, w.n)) * 500, 0.0)
                st.O[0, D] = rng.uniform(0, 2500, w.n).astype(st.O.dtype)
                births = np.where(rng.random((2, w.n)) < 0.7, rng.random((2, w.n)) * 20, 0.0)
                ok = rng.random((2, w.n)) < 0.8
                rules = init_rules(races)
                rules[0]["origin"] = 0
                fn(st, w, races, rules, births, regs, ok)
                out.append(st)
            with self.subTest(emerge=emerge):
                self.assertTrue(np.array_equal(out[0].P, out[1].P))
                self.assertTrue(np.array_equal(out[0].O, out[1].O))
                self.assertGreater(out[0].P[1].sum(), 0)


class ChangeCellsTest(unittest.TestCase):
    def test_compiled_scan_matches_numpy(self):
        if S._change_cells_jit is None:
            self.skipTest("numba not installed")
        rng = np.random.default_rng(11)
        n = 20000
        for seed in range(5):
            P = np.where(rng.random(n) < 0.3, rng.random(n) * 1e3, 0.0)
            k = rng.choice(n, 400, replace=False)
            P[k[:100]] = -0.0
            P[k[100:200]] = rng.random(100) * 1e-12
            P[k[200:250]] = -rng.random(50)
            P[k[250:270]] = np.nan
            P[k[270:280]] = -np.nan
            P[k[280:290]] = 1e-12
            C = np.where(rng.random((1, n)) < 0.05, rng.random((1, n)), 0.0)
            C[0, k[290:300]] = np.nan
            C[0, k[300:310]] = -0.0
            src, dst = rng.integers(0, n, 300), rng.integers(0, n, 300)
            for wide in ([C], []):
                with self.subTest(seed=seed, wide=len(wide)):
                    self.assertTrue(np.array_equal(S.change_cells(P, wide, src, dst), S._change_cells_np(P, wide, src, dst)))


if __name__ == "__main__":
    unittest.main()