aboutsummaryrefslogtreecommitdiffziptar.gz
path: root/tests/test_fastpaths.py
diff options
context:
space:
mode:
Diffstat (limited to 'tests/test_fastpaths.py')
-rw-r--r--tests/test_fastpaths.py314
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()