diff options
Diffstat (limited to 'tests/test_fastpaths.py')
| -rw-r--r-- | tests/test_fastpaths.py | 314 |
1 files changed, 314 insertions, 0 deletions
diff --git a/tests/test_fastpaths.py b/tests/test_fastpaths.py new file mode 100644 index 0000000..893be9b --- /dev/null +++ b/tests/test_fastpaths.py @@ -0,0 +1,314 @@ +"""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() |
