From 43a3a381c845420fc3da10cddf682bc91877fa08 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Wed, 2 Sep 2026 14:45:01 +0100 Subject: [PATCH 01/19] Add a benchmark for genetic_value over the number of causal sites The genetic value computation is slow for a trait with a large number of weakly causal sites, so add a benchmark that measures how it scales. The script simulates 100k sample nodes over 5Mb, caching the tree sequence, and times genetic_value for 1, 100, 1000, 10,000 and 100,000 causal sites at each of the individual, node and edge levels. The mutation rate is raised to ten times the human rate so that there are enough sites in a genome short enough to simulate quickly; this leaves the allele frequency spectrum, and so the weakly causal regime, intact. sim_trait is timed separately because it has a per-site Python loop of its own, the numba kernels are compiled by an untimed warm up call, and --max-seconds gives up on the larger points of the grid once a single call goes over budget. On a tree sequence with 217k nodes the cost per causal site is flat from 1000 sites upwards, at 444us for individuals, 384us for nodes and 1807us for edges, so 100,000 causal sites take 44s, 38s and 181s respectively. Holding the number of causal sites fixed and varying the sample size, the cost per site normalised by num_nodes is constant across a sixteen fold range, so the computation is O(num_causal * num_nodes): every causal site pays several passes over every node in the tree sequence. The median causal site is carried by 0.4% of the nodes. --- benchmarks/README.md | 40 ++++ benchmarks/benchmark_genetic_value.py | 260 ++++++++++++++++++++++++++ 2 files changed, 300 insertions(+) create mode 100644 benchmarks/README.md create mode 100644 benchmarks/benchmark_genetic_value.py diff --git a/benchmarks/README.md b/benchmarks/README.md new file mode 100644 index 0000000..16796f4 --- /dev/null +++ b/benchmarks/README.md @@ -0,0 +1,40 @@ +# Benchmarks + +Performance benchmarks for tstrait. These are not run in CI; they are here so +that performance work can be repeated and compared. + +## `benchmark_genetic_value.py` + +Measures how `tstrait.genetic_value` scales with the number of causal sites in +a trait, which is the regime where the current implementation is slow: a trait +with a large number of weakly causal sites. + +``` +uv run --group test benchmarks/benchmark_genetic_value.py +``` + +By default this simulates 100,000 sample nodes (50,000 diploid individuals) +over 5Mb, and times `genetic_value` for 1, 100, 1000, 10,000 and 100,000 causal +sites at each of the `individual`, `node` and `edge` levels, taking the minimum +of three replicates. + +The mutation rate defaults to 1e-7, ten times the human rate, so that there are +enough sites in a genome short enough to simulate quickly. The mutation rate +does not affect the allele frequency spectrum, so the causal sites are as +weakly causal as they would be under a realistic rate; only the number of sites +per tree is inflated. + +`sim_trait` is timed separately, because it has a per-site Python loop of its +own that we do not want folded into the `genetic_value` numbers. The numba +kernels are compiled by a warm up call that is not timed. + +The simulation parameters, the grid of causal site counts, the levels and the +number of replicates are all command line arguments; see `--help`. + +### Output + +The simulated tree sequence is cached in `_output/`, keyed by the simulation +parameters, so that repeated runs do not resimulate it. Results are written to +`_output/genetic_value.csv` in long format, one row per replicate, together +with the dimensions of the tree sequence they were measured on. `_output/` is +gitignored. diff --git a/benchmarks/benchmark_genetic_value.py b/benchmarks/benchmark_genetic_value.py new file mode 100644 index 0000000..9ffb5c8 --- /dev/null +++ b/benchmarks/benchmark_genetic_value.py @@ -0,0 +1,260 @@ +""" +Benchmark tstrait.genetic_value as the number of causal sites grows. + +The interesting regime is a trait with a large number of weakly causal sites, +i.e. many causal sites that are each carried by a small number of samples. +Causal sites are chosen uniformly at random over all sites, and because the +site frequency spectrum is dominated by rare variants this puts us in that +regime by default. + +Run with ``uv run --group test benchmarks/benchmark_genetic_value.py``. +""" + +import argparse +import csv +import functools +import pathlib +import sys +import time + +import msprime +import tskit + +import tstrait + +DEFAULT_NUM_CAUSAL = [1, 100, 1000, 10_000, 100_000] +LEVELS = ["individual", "node", "edge"] + + +def cached_simulation(args): + """ + Return the tree sequence defined by the simulation arguments, simulating it + and caching it to disk if we have not done so already. + """ + name = ( + f"n={args.samples}_L={args.length:.0f}_r={args.recombination_rate:g}" + f"_mu={args.mutation_rate:g}_Ne={args.population_size:g}_seed={args.seed}" + ) + path = args.cache_dir / f"{name}.trees" + if path.exists(): + print(f"Loading cached {path}") + return tskit.load(path) + + print(f"Simulating {name}") + before = time.perf_counter() + ts = msprime.sim_ancestry( + args.samples // 2, + ploidy=2, + sequence_length=args.length, + recombination_rate=args.recombination_rate, + population_size=args.population_size, + random_seed=args.seed, + ) + ancestry_time = time.perf_counter() - before + before = time.perf_counter() + ts = msprime.sim_mutations(ts, rate=args.mutation_rate, random_seed=args.seed) + mutation_time = time.perf_counter() - before + print(f" ancestry {ancestry_time:.1f}s, mutations {mutation_time:.1f}s") + + args.cache_dir.mkdir(parents=True, exist_ok=True) + ts.dump(path) + print(f" cached to {path}") + return ts + + +def describe(ts): + return ( + f"samples={ts.num_samples} individuals={ts.num_individuals} " + f"nodes={ts.num_nodes} edges={ts.num_edges} trees={ts.num_trees} " + f"sites={ts.num_sites}" + ) + + +def time_call(func, replicates): + """ + Call func the given number of times, returning the result of the last call + and the elapsed time of each call. + """ + times = [] + for _ in range(replicates): + before = time.perf_counter() + result = func() + times.append(time.perf_counter() - before) + return result, times + + +def warm_up(ts, model, levels): + """ + Run the smallest possible simulation through each code path so that the + numba kernels are compiled before we start timing. + """ + before = time.perf_counter() + trait_df = tstrait.sim_trait(ts, model=model, num_causal=1, random_seed=1) + for level in levels: + tstrait.genetic_value(ts, trait_df, level=level) + print(f"Warm up (includes numba compilation): {time.perf_counter() - before:.1f}s") + + +def run_benchmark(ts, args): + """ + Time each cell of the grid, returning the timings and the causal site + counts that we got through before running out of time budget. + """ + model = tstrait.trait_model(distribution="normal", mean=0, var=1) + warm_up(ts, model, args.levels) + + rows = [] + completed = [] + for num_causal in args.num_causal: + trait_df, times = time_call( + functools.partial( + tstrait.sim_trait, + ts, + model=model, + num_causal=num_causal, + random_seed=args.seed, + ), + args.replicates, + ) + for replicate, seconds in enumerate(times): + rows.append(("sim_trait", num_causal, "", replicate, seconds)) + report(rows[-1]) + slowest = max(times) + for level in args.levels: + _, times = time_call( + functools.partial(tstrait.genetic_value, ts, trait_df, level=level), + args.replicates, + ) + for replicate, seconds in enumerate(times): + rows.append(("genetic_value", num_causal, level, replicate, seconds)) + report(rows[-1]) + slowest = max(slowest, max(times)) + completed.append(num_causal) + # The cost is superlinear in num_causal, so once a single call is over + # budget the next point on the grid is not worth waiting for. + if slowest > args.max_seconds: + skipped = args.num_causal[len(completed) :] + if len(skipped) > 0: + print( + f" a call took {slowest:.1f}s, over the {args.max_seconds:g}s " + f"budget: skipping {skipped}" + ) + break + return rows, completed + + +def report(row): + phase, num_causal, level, _, seconds = row + print(f" {phase:<14} {num_causal:>7} {level:<11} {seconds:8.3f}s") + + +def summarise(rows, completed, ts, args): + """ + Print the minimum time over the replicates of each cell of the grid, along + with the time per causal site. + """ + best = {} + for phase, num_causal, level, _, seconds in rows: + key = (phase, level, num_causal) + best[key] = min(best.get(key, seconds), seconds) + + print(f"\n{describe(ts)}") + print(f"Minimum of {args.replicates} replicates\n") + header = ( + f"{'phase':<14} {'level':<11} {'num_causal':>10} {'seconds':>10} {'us/site':>10}" + ) + print(header) + print("-" * len(header)) + for phase, level in [("sim_trait", "")] + [ + ("genetic_value", x) for x in args.levels + ]: + for num_causal in completed: + seconds = best[(phase, level, num_causal)] + print( + f"{phase:<14} {level:<11} {num_causal:>10} {seconds:>10.3f} " + f"{seconds / num_causal * 1e6:>10.1f}" + ) + + +def write_csv(rows, ts, args): + args.output.parent.mkdir(parents=True, exist_ok=True) + with open(args.output, "w", newline="") as f: + writer = csv.writer(f) + writer.writerow( + [ + "phase", + "num_causal", + "level", + "replicate", + "seconds", + "num_samples", + "num_individuals", + "num_nodes", + "num_edges", + "num_trees", + "num_sites", + ] + ) + for row in rows: + writer.writerow( + [ + *row, + ts.num_samples, + ts.num_individuals, + ts.num_nodes, + ts.num_edges, + ts.num_trees, + ts.num_sites, + ] + ) + print(f"\nWrote {args.output}") + + +def parse_args(): + default_output = pathlib.Path(__file__).parent / "_output" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument( + "--samples", type=int, default=100_000, help="Number of sample nodes" + ) + parser.add_argument("--length", type=float, default=5e6, help="Sequence length") + parser.add_argument("--recombination-rate", type=float, default=1e-8) + parser.add_argument("--mutation-rate", type=float, default=1e-7) + parser.add_argument("--population-size", type=float, default=10_000) + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--num-causal", type=int, nargs="+", default=DEFAULT_NUM_CAUSAL) + parser.add_argument("--levels", nargs="+", choices=LEVELS, default=LEVELS) + parser.add_argument("--replicates", type=int, default=3) + parser.add_argument( + "--max-seconds", + type=float, + default=60, + help=( + "Stop before the next number of causal sites once a single call has " + "taken longer than this" + ), + ) + parser.add_argument("--cache-dir", type=pathlib.Path, default=default_output) + parser.add_argument( + "--output", type=pathlib.Path, default=default_output / "genetic_value.csv" + ) + return parser.parse_args() + + +def main(): + sys.stdout.reconfigure(line_buffering=True) + args = parse_args() + ts = cached_simulation(args) + print(describe(ts)) + max_causal = max(args.num_causal) + if ts.num_sites < max_causal: + raise ValueError( + f"Only {ts.num_sites} sites in the tree sequence, but {max_causal} " + "causal sites were requested. Increase --length or --mutation-rate." + ) + rows, completed = run_benchmark(ts, args) + summarise(rows, completed, ts, args) + write_csv(rows, ts, args) + + +if __name__ == "__main__": + main() From 1cde68dd524b977c455f64e783e78cf9d2a0b80a Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Wed, 2 Sep 2026 16:03:50 +0100 Subject: [PATCH 02/19] Accumulate genetic values in a single descent of the ARG Replace the tree by tree loop in genetic_value with one descent over the ARG that accumulates every causal site at once, using the tskit child index to find the child edges of a node and carrying a range of causal site indexes down so that a subtree spanning no causal site is pruned. The causal allele state change of a mutation telescopes down a path, so values are additive from root to node and nested and back mutations need no special handling; the has_mutation pruning is gone, and so is the separate accumulation pass for edges, which now fall out of the descent by crediting the edge that a value arrives through. The value carried down is a sum over the sites in the current range, so when an edge narrows that range it is not the value the child should inherit: the mutations above at the sites that dropped out do not apply below. The path back to the root is therefore kept, and the value is recomputed over the narrowed range when that happens. It is rare, a hundredth of a percent of descents on a large tree sequence. This is slower than what it replaces and is committed as a checkpoint only. On 100k samples at level="node", 1000 causal sites go from 0.436s to 14.4s and 10,000 from 3.761s to 51.7s. Cost is independent of allele frequency as intended, rare and uniform causal sites taking the same time, but the per visit constant is around 40 times worse than a prototype that only counted visits, once the value lookups, the per edge mutation searches and the output writes are in the loop. Add naive_genetic_value, the tree by tree implementation kept as a reference oracle, and TestGeneticValueReference comparing against it over the tests/data.py tree sequences, all_trees_ts(2..5) with recurrent and back mutations, and simulations with and without recombination, at each of the individual, node and edge levels and with multiple traits, along with isolated samples, multiple roots and mutations above a root. Give the benchmark a causal site selection axis, since drawing sites uniformly is dominated by the common variants in the tail of the frequency spectrum and the rare ones behave quite differently. --- benchmarks/benchmark_genetic_value.py | 106 +++++++-- tests/test_genetic_value.py | 181 +++++++++++++++ tests/test_individual_node_edge.py | 17 -- tests/test_jit.py | 323 ++++++++++++++++++-------- tstrait/genetic_value.py | 314 ++++++++++++++++++------- tstrait/jit.py | 259 ++++++++++++++++++--- 6 files changed, 938 insertions(+), 262 deletions(-) diff --git a/benchmarks/benchmark_genetic_value.py b/benchmarks/benchmark_genetic_value.py index 9ffb5c8..237645a 100644 --- a/benchmarks/benchmark_genetic_value.py +++ b/benchmarks/benchmark_genetic_value.py @@ -18,12 +18,14 @@ import time import msprime +import numpy as np import tskit import tstrait DEFAULT_NUM_CAUSAL = [1, 100, 1000, 10_000, 100_000] LEVELS = ["individual", "node", "edge"] +SELECTIONS = ["uniform", "rare"] def cached_simulation(args): @@ -83,6 +85,34 @@ def time_call(func, replicates): return result, times +def causal_site_pool(ts, model, seed): + """ + Simulate effect sizes for every site, so that causal sites can be drawn + from the pool by allele frequency without resimulating. + """ + before = time.perf_counter() + pool = tstrait.sim_trait(ts, model=model, num_causal=ts.num_sites, random_seed=seed) + print( + f"Effect sizes for all {ts.num_sites} sites: {time.perf_counter() - before:.1f}s" + ) + return pool + + +def select_causal(pool, selection, num_causal, rare_threshold, rng): + """ + Draw num_causal rows from the pool, either uniformly over the sites or + restricted to the rare ones. Uniform selection is dominated by the common + variants in the tail of the frequency spectrum, which behave quite + differently from the weakly causal sites the ARG descent is aimed at. + """ + if selection == "rare": + pool = pool[pool["allele_freq"] < rare_threshold] + if len(pool) < num_causal: + return None + keep = rng.choice(len(pool), size=num_causal, replace=False) + return pool.iloc[np.sort(keep)] + + def warm_up(ts, model, levels): """ Run the smallest possible simulation through each code path so that the @@ -103,10 +133,11 @@ def run_benchmark(ts, args): model = tstrait.trait_model(distribution="normal", mean=0, var=1) warm_up(ts, model, args.levels) + pool = causal_site_pool(ts, model, args.seed) rows = [] completed = [] for num_causal in args.num_causal: - trait_df, times = time_call( + _, times = time_call( functools.partial( tstrait.sim_trait, ts, @@ -117,18 +148,35 @@ def run_benchmark(ts, args): args.replicates, ) for replicate, seconds in enumerate(times): - rows.append(("sim_trait", num_causal, "", replicate, seconds)) + rows.append(("sim_trait", num_causal, "uniform", "", replicate, seconds)) report(rows[-1]) slowest = max(times) - for level in args.levels: - _, times = time_call( - functools.partial(tstrait.genetic_value, ts, trait_df, level=level), - args.replicates, + for selection in args.selections: + rng = np.random.default_rng(args.seed) + trait_df = select_causal( + pool, selection, num_causal, args.rare_threshold, rng ) - for replicate, seconds in enumerate(times): - rows.append(("genetic_value", num_causal, level, replicate, seconds)) - report(rows[-1]) - slowest = max(slowest, max(times)) + if trait_df is None: + print(f" too few {selection} sites for num_causal={num_causal}") + continue + for level in args.levels: + _, times = time_call( + functools.partial(tstrait.genetic_value, ts, trait_df, level=level), + args.replicates, + ) + for replicate, seconds in enumerate(times): + rows.append( + ( + "genetic_value", + num_causal, + selection, + level, + replicate, + seconds, + ) + ) + report(rows[-1]) + slowest = max(slowest, max(times)) completed.append(num_causal) # The cost is superlinear in num_causal, so once a single call is over # budget the next point on the grid is not worth waiting for. @@ -144,8 +192,8 @@ def run_benchmark(ts, args): def report(row): - phase, num_causal, level, _, seconds = row - print(f" {phase:<14} {num_causal:>7} {level:<11} {seconds:8.3f}s") + phase, num_causal, selection, level, _, seconds = row + print(f" {phase:<14} {num_causal:>7} {selection:<8} {level:<11} {seconds:8.3f}s") def summarise(rows, completed, ts, args): @@ -154,25 +202,31 @@ def summarise(rows, completed, ts, args): with the time per causal site. """ best = {} - for phase, num_causal, level, _, seconds in rows: - key = (phase, level, num_causal) + for phase, num_causal, selection, level, _, seconds in rows: + key = (phase, selection, level, num_causal) best[key] = min(best.get(key, seconds), seconds) print(f"\n{describe(ts)}") print(f"Minimum of {args.replicates} replicates\n") header = ( - f"{'phase':<14} {'level':<11} {'num_causal':>10} {'seconds':>10} {'us/site':>10}" + f"{'phase':<14} {'selection':<10} {'level':<11} {'num_causal':>10} " + f"{'seconds':>10} {'us/site':>10}" ) print(header) print("-" * len(header)) - for phase, level in [("sim_trait", "")] + [ - ("genetic_value", x) for x in args.levels - ]: + combinations = [("sim_trait", "uniform", "")] + combinations += [ + ("genetic_value", s, x) for s in args.selections for x in args.levels + ] + for phase, selection, level in combinations: for num_causal in completed: - seconds = best[(phase, level, num_causal)] + key = (phase, selection, level, num_causal) + if key not in best: + continue + seconds = best[key] print( - f"{phase:<14} {level:<11} {num_causal:>10} {seconds:>10.3f} " - f"{seconds / num_causal * 1e6:>10.1f}" + f"{phase:<14} {selection:<10} {level:<11} {num_causal:>10} " + f"{seconds:>10.3f} {seconds / num_causal * 1e6:>10.1f}" ) @@ -184,6 +238,7 @@ def write_csv(rows, ts, args): [ "phase", "num_causal", + "selection", "level", "replicate", "seconds", @@ -223,6 +278,15 @@ def parse_args(): parser.add_argument("--seed", type=int, default=42) parser.add_argument("--num-causal", type=int, nargs="+", default=DEFAULT_NUM_CAUSAL) parser.add_argument("--levels", nargs="+", choices=LEVELS, default=LEVELS) + parser.add_argument( + "--selections", nargs="+", choices=SELECTIONS, default=SELECTIONS + ) + parser.add_argument( + "--rare-threshold", + type=float, + default=0.001, + help="Allele frequency below which a causal site counts as rare", + ) parser.add_argument("--replicates", type=int, default=3) parser.add_argument( "--max-seconds", diff --git a/tests/test_genetic_value.py b/tests/test_genetic_value.py index 9a205d1..05075ea 100644 --- a/tests/test_genetic_value.py +++ b/tests/test_genetic_value.py @@ -1076,3 +1076,184 @@ def test_pleiotropy(self, sample_ts): var_array = grouped.var().values.T[0] np.testing.assert_almost_equal(mean_array, np.zeros(2), decimal=2) np.testing.assert_almost_equal(var_array, np.ones(2), decimal=2) + + +def naive_genetic_value(ts, trait_df, level): + """Reference implementation of genetic_value, computed tree by tree. + + For each row of the trait dataframe this seeks to the tree at the causal + site and propagates the causal allele down from the virtual root, stopping + wherever a mutation changes the state, which is how genetic values were + computed before the ARG descent. + """ + num_trait = int(np.max(trait_df["trait_id"])) + 1 + size = { + "individual": ts.num_individuals, + "node": ts.num_nodes, + "edge": ts.num_edges, + }[level] + genetic_value_table = np.zeros((num_trait, size)) + tree = tskit.Tree(ts) + for data in trait_df.itertuples(): + site = ts.site(data.site_id) + tree.seek(site.position) + state = {tree.virtual_root: site.ancestral_state} + has_mutation = set() + for m in site.mutations: + state[m.node] = m.derived_state + has_mutation.add(m.node) + # One extra entry, so that the virtual root can be a causal node. + nodes_value = np.zeros(ts.num_nodes + 1) + stack = [u for u, allele in state.items() if allele == data.causal_allele] + while len(stack) > 0: + u = stack.pop() + nodes_value[u] = data.effect_size + for child in tree.children(u): + if child not in has_mutation: + stack.append(child) + nodes_value = nodes_value[: ts.num_nodes] + + if level == "individual": + value = np.zeros(ts.num_individuals) + for u, individual in enumerate(ts.nodes_individual): + if individual != tskit.NULL: + value[individual] += nodes_value[u] + elif level == "edge": + value = np.zeros(ts.num_edges) + for u in range(ts.num_nodes): + edge = tree.edge(u) + if edge != tskit.NULL: + value[edge] += nodes_value[u] + else: + value = nodes_value + genetic_value_table[data.trait_id] += value + + return pd.DataFrame( + { + "trait_id": np.repeat(np.arange(num_trait), size), + f"{level}_id": np.tile(np.arange(size), num_trait), + "genetic_value": genetic_value_table.flatten(), + } + ) + + +class TestGeneticValueReference: + """Compare genetic_value against a tree by tree reference implementation + over a range of tree topologies and allele configurations. + + The causal alleles drawn by random_trait_df include the ancestral state of + a site, which is the case the ARG descent folds into the value it seeds the + roots with, so it is covered here rather than separately. + """ + + def verify(self, ts, num_trait, seed=1, levels=("node", "edge")): + trait_df = random_trait_df(ts, num_trait, seed) + for level in levels: + if level == "individual" and ts.num_individuals == 0: + continue + expected = naive_genetic_value(ts, trait_df, level) + result = tstrait.genetic_value(ts=ts, trait_df=trait_df, level=level) + pd.testing.assert_frame_equal(result, expected, check_dtype=False) + + @pytest.mark.parametrize( + "ts_func", + [ + binary_tree, + diff_ind_tree, + non_binary_tree, + triploid_tree, + binary_tree_seq, + simple_tree_seq, + allele_freq_one, + ], + ) + @pytest.mark.parametrize("num_trait", [1, 3]) + def test_data_tree_sequence(self, ts_func, num_trait): + self.verify(ts_func(), num_trait, levels=("individual", "node", "edge")) + + @pytest.mark.parametrize("n", [2, 3, 4, 5]) + @pytest.mark.parametrize("rate", [2.0, 10.0]) + def test_all_trees(self, n, rate): + ts = multi_allelic_mutations(all_trees_ts(n), rate=rate, seed=n) + self.verify(ts, num_trait=2) + + @pytest.mark.parametrize("recombination_rate", [0, 1e-7]) + @pytest.mark.parametrize("seed", [1, 2, 3]) + def test_simulated(self, recombination_rate, seed): + ts = msprime.sim_ancestry( + 8, + sequence_length=10_000, + recombination_rate=recombination_rate, + population_size=1000, + random_seed=seed, + ) + ts = multi_allelic_mutations(ts, rate=1e-4, seed=seed) + self.verify(ts, num_trait=3, seed=seed, levels=("individual", "node", "edge")) + + def test_isolated_samples(self): + """Isolated samples are roots, so they carry the ancestral state.""" + tables = tskit.TableCollection(sequence_length=10) + for _ in range(4): + tables.nodes.add_row(flags=tskit.NODE_IS_SAMPLE, time=0) + site = tables.sites.add_row(position=0, ancestral_state="A") + tables.mutations.add_row(site=site, node=0, derived_state="T") + ts = tables.tree_sequence() + for causal_allele in ["A", "T"]: + trait_df = pd.DataFrame( + { + "site_id": [0], + "effect_size": [2.5], + "trait_id": [0], + "causal_allele": [causal_allele], + } + ) + for level in ("node", "edge"): + expected = naive_genetic_value(ts, trait_df, level) + result = tstrait.genetic_value(ts=ts, trait_df=trait_df, level=level) + pd.testing.assert_frame_equal(result, expected, check_dtype=False) + + def test_multiple_roots(self): + """A forest has several roots, each carrying the ancestral state.""" + tables = tskit.Tree.generate_balanced(4).tree_sequence.dump_tables() + tables.edges.replace_with(tables.edges[tables.edges.parent != 6]) + site = tables.sites.add_row(position=0, ancestral_state="A") + tables.mutations.add_row(site=site, node=4, derived_state="T") + tables.sort() + ts = tables.tree_sequence() + for causal_allele in ["A", "T"]: + trait_df = pd.DataFrame( + { + "site_id": [0], + "effect_size": [3.0], + "trait_id": [0], + "causal_allele": [causal_allele], + } + ) + for level in ("node", "edge"): + expected = naive_genetic_value(ts, trait_df, level) + result = tstrait.genetic_value(ts=ts, trait_df=trait_df, level=level) + pd.testing.assert_frame_equal(result, expected, check_dtype=False) + + def test_mutation_above_a_root(self): + """A mutation with no edge applies to the root it sits above.""" + tables = tskit.Tree.generate_balanced(4).tree_sequence.dump_tables() + site = tables.sites.add_row(position=0, ancestral_state="A") + tables.mutations.add_row(site=site, node=6, derived_state="T") + tables.mutations.add_row(site=site, node=4, derived_state="G") + tables.sort() + tables.build_index() + tables.compute_mutation_parents() + ts = tables.tree_sequence() + for causal_allele in ["A", "T", "G"]: + trait_df = pd.DataFrame( + { + "site_id": [0], + "effect_size": [1.5], + "trait_id": [0], + "causal_allele": [causal_allele], + } + ) + for level in ("node", "edge"): + expected = naive_genetic_value(ts, trait_df, level) + result = tstrait.genetic_value(ts=ts, trait_df=trait_df, level=level) + pd.testing.assert_frame_equal(result, expected, check_dtype=False) diff --git a/tests/test_individual_node_edge.py b/tests/test_individual_node_edge.py index c61ee52..b8257cb 100644 --- a/tests/test_individual_node_edge.py +++ b/tests/test_individual_node_edge.py @@ -13,23 +13,6 @@ import tskit import tstrait -from tstrait.genetic_value import _accumulate_edge_values - - -def test_accumulate_edge_values(): - nodes_genetic_value = np.array([10, 2, 0, 0, -3, 3, 5, 20, -1]) - node_edges = np.array( - [tskit.NULL, 0, tskit.NULL, 1, 1, 1, 2, tskit.NULL, 2], dtype=int - ) - - observed = _accumulate_edge_values( - nodes_genetic_value=nodes_genetic_value, - nodes_edge=node_edges, - num_nodes=len(nodes_genetic_value), - num_edges=3, - ) - - np.testing.assert_array_equal(observed, [2, 0, 4]) @pytest.fixture(scope="module") diff --git a/tests/test_jit.py b/tests/test_jit.py index 9726122..0b517de 100644 --- a/tests/test_jit.py +++ b/tests/test_jit.py @@ -9,13 +9,17 @@ The examples are small trees from the tskit generators, drawn in the class docstrings so that the expected values can be checked against the topology. +Each is turned into a tree sequence with a single causal site, since the +descent works from the mutations of the ARG rather than from a tree. """ import numpy as np +import pandas as pd import pytest import tskit from tstrait import jit +from tstrait.genetic_value import _GeneticValue from .data import ( binary_tree, @@ -23,6 +27,8 @@ triploid_tree, ) # noreorder +ANCESTRAL_STATE = "A" + def kernel(func, param): """ @@ -31,27 +37,54 @@ def kernel(func, param): return func if param == "jit" else func.py_func +def one_site(tree, mutations): + """ + Return the tree's tree sequence carrying a single site at position 0, with + the given ``(node, derived_state)`` mutations. Mutations must be listed + from the oldest node down, so that a mutation follows its parent. + """ + tables = tree.tree_sequence.dump_tables() + tables.sites.clear() + tables.mutations.clear() + site = tables.sites.add_row(position=0, ancestral_state=ANCESTRAL_STATE) + for node, derived_state in mutations: + tables.mutations.add_row(site=site, node=node, derived_state=derived_state) + tables.sort() + tables.build_index() + tables.compute_mutation_parents() + return tables.tree_sequence() + + @pytest.fixture(params=["jit", "nojit"]) def node_genetic_value(request): """ - Return a function computing the node genetic values for a tskit tree. + Return a function computing the node genetic values of a tree carrying one + causal site. - The ``causal_nodes`` are the nodes at which the causal allele appears, - which includes the virtual root when the ancestral state is causal, and the - ``mutation_nodes`` are the nodes at which a mutation occurs. + The ``mutations`` are ``(node, derived_state)`` pairs against an ancestral + state of "A", and a node's value is ``effect_size`` when the allele it + inherits is ``causal_allele``. """ - func = kernel(jit._compute_nodes_genetic_value, request.param) - - def f(tree, causal_nodes, mutation_nodes=(), effect_size=1): - has_mutation = np.zeros(len(tree.left_child_array), dtype=bool) - has_mutation[list(mutation_nodes)] = True - return func( - left_child_array=tree.left_child_array, - right_sib_array=tree.right_sib_array, - causal_nodes=np.array(causal_nodes, dtype=np.int32), - has_mutation=has_mutation, - effect_size=effect_size, + func = kernel(jit._descend_arg, request.param) + + def f( + tree, + mutations=(), + causal_allele=ANCESTRAL_STATE, + effect_size=1, + level="node", + ): + ts = tree if isinstance(tree, tskit.TreeSequence) else one_site(tree, mutations) + trait_df = pd.DataFrame( + { + "site_id": [0], + "effect_size": [effect_size], + "trait_id": [0], + "causal_allele": [causal_allele], + } ) + genetic = _GeneticValue(ts, trait_df) + return func(**genetic._descent_arguments(level, 0)) return f @@ -74,14 +107,18 @@ def f(ts, nodes_genetic_value): return f -def isolated_samples_tree(n): +def isolated_samples_ts(n, mutations=()): """ - Return a tree of n isolated sample nodes, which are therefore all roots. + Return a tree sequence of n isolated sample nodes, which are all roots, + carrying one causal site. """ tables = tskit.TableCollection(sequence_length=1) for _ in range(n): tables.nodes.add_row(flags=tskit.NODE_IS_SAMPLE, time=0) - return tables.tree_sequence().first() + site = tables.sites.add_row(position=0, ancestral_state=ANCESTRAL_STATE) + for node, derived_state in mutations: + tables.mutations.add_row(site=site, node=node, derived_state=derived_state) + return tables.tree_sequence() def multiroot_tree(): @@ -94,16 +131,84 @@ def multiroot_tree(): return tables.tree_sequence().first() -def empty_tree(): +def empty_ts(): + """ + Return a tree sequence that has no nodes, but does have a causal site. + """ + tables = tskit.TableCollection(sequence_length=1) + tables.sites.add_row(position=0, ancestral_state=ANCESTRAL_STATE) + return tables.tree_sequence() + + +class TestSearchSorted: + """ + The binary search the descent uses to find the mutations of an edge that + lie in a range of causal sites. + """ + + @pytest.fixture(params=["jit", "nojit"]) + def search_sorted(self, request): + return kernel(jit._search_sorted, request.param) + + @pytest.mark.parametrize( + ("key", "expected"), [(-1, 0), (0, 0), (1, 1), (3, 1), (4, 1), (5, 3), (9, 4), (10, 5)] + ) + def test_search(self, search_sorted, key, expected): + values = np.array([0, 4, 4, 7, 9], dtype=np.int32) + assert search_sorted(values, 0, len(values), key) == expected + + def test_restricted_range(self, search_sorted): + # Only the values in [1, 4) are searched, so a key below the range + # comes back as the start of it. + values = np.array([0, 4, 4, 7, 9], dtype=np.int32) + assert search_sorted(values, 1, 4, -1) == 1 + assert search_sorted(values, 1, 4, 8) == 4 + + def test_empty_range(self, search_sorted): + values = np.array([0, 4, 4, 7, 9], dtype=np.int32) + assert search_sorted(values, 2, 2, 4) == 2 + + +class TestGroupWeight: """ - Return the tree of a tree sequence that has no nodes. + The total weight of the entries of one group that fall in a range of + causal sites. Group 0 holds sites 1 and 3, group 1 is empty and group 2 + holds sites 0, 2 and 2. """ - return tskit.TableCollection(sequence_length=1).tree_sequence().first() + + @pytest.fixture(params=["jit", "nojit"]) + def group_weight(self, request): + func = kernel(jit._group_weight, request.param) + offset = np.array([0, 2, 2, 5], dtype=np.int32) + site = np.array([1, 3, 0, 2, 2], dtype=np.int32) + weight_sum = np.array([0.0, 1.0, 3.0, 7.0, 15.0, 31.0]) + + def f(group, start, stop): + return func(offset, site, weight_sum, group, start, stop) + + return f + + def test_whole_group(self, group_weight): + assert group_weight(0, 0, 4) == 3.0 + assert group_weight(2, 0, 4) == 28.0 + + def test_empty_group(self, group_weight): + assert group_weight(1, 0, 4) == 0.0 + + def test_partial_range(self, group_weight): + assert group_weight(0, 0, 2) == 1.0 + assert group_weight(0, 2, 4) == 2.0 + # Both entries at site 2 are inside the range or outside it together. + assert group_weight(2, 2, 3) == 24.0 + assert group_weight(2, 0, 1) == 4.0 + + def test_range_outside_group(self, group_weight): + assert group_weight(0, 4, 8) == 0.0 class TestBalancedBinaryTree: """ - tskit.Tree.generate_balanced(4), in which node 7 is the virtual root:: + tskit.Tree.generate_balanced(4):: 6 +-+-+ @@ -117,63 +222,75 @@ def tree(self): return tskit.Tree.generate_balanced(4) def test_no_causal_nodes(self, tree, node_genetic_value): + # The causal allele occurs nowhere in the tree. np.testing.assert_array_equal( - node_genetic_value(tree, []), [0, 0, 0, 0, 0, 0, 0] + node_genetic_value(tree, [], "T"), [0, 0, 0, 0, 0, 0, 0] ) def test_leaf(self, tree, node_genetic_value): - # Node 0 has no children, so the sibling walk is never entered. + # Node 0 has no children, so the descent stops as soon as it is + # reached. np.testing.assert_array_equal( - node_genetic_value(tree, [0]), [1, 0, 0, 0, 0, 0, 0] + node_genetic_value(tree, [(0, "T")], "T"), [1, 0, 0, 0, 0, 0, 0] ) def test_internal_node(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [4]), [1, 1, 0, 0, 1, 0, 0] + node_genetic_value(tree, [(4, "T")], "T"), [1, 1, 0, 0, 1, 0, 0] ) def test_root(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [6]), [1, 1, 1, 1, 1, 1, 1] + node_genetic_value(tree, [(6, "T")], "T"), [1, 1, 1, 1, 1, 1, 1] ) - def test_virtual_root(self, tree, node_genetic_value): - # The ancestral state is the causal allele, so the virtual root is - # causal. Its own value is computed but is not part of the output. - value = node_genetic_value(tree, [7]) - assert len(value) == 7 - np.testing.assert_array_equal(value, [1, 1, 1, 1, 1, 1, 1]) + def test_ancestral_state_is_causal(self, tree, node_genetic_value): + # Every node carries the ancestral state, so the value is seeded at + # the root and never changes. + np.testing.assert_array_equal(node_genetic_value(tree), [1, 1, 1, 1, 1, 1, 1]) def test_mutation_on_internal_node(self, tree, node_genetic_value): - # Node 4 and its children carry whatever allele the mutation on node 4 - # introduced, so the traversal does not descend into them. + # Node 4 and its children carry the allele the mutation on node 4 + # introduced, which is not the causal one. np.testing.assert_array_equal( - node_genetic_value(tree, [6], [4]), [0, 0, 1, 1, 0, 1, 1] + node_genetic_value(tree, [(6, "T"), (4, "G")], "T"), [0, 0, 1, 1, 0, 1, 1] ) def test_mutation_on_leaf(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [6], [0]), [0, 1, 1, 1, 1, 1, 1] + node_genetic_value(tree, [(6, "T"), (0, "G")], "T"), [0, 1, 1, 1, 1, 1, 1] ) - def test_two_causal_nodes(self, tree, node_genetic_value): - # Both mutations are back to the causal allele, so both nodes are - # causal and each starts its own traversal. + def test_back_mutations(self, tree, node_genetic_value): + # The mutations on nodes 1 and 5 are back to the causal allele, which + # the state changes telescope to without any special handling. np.testing.assert_array_equal( - node_genetic_value(tree, [1, 5], [1, 5]), [0, 1, 1, 1, 0, 1, 0] + node_genetic_value(tree, [(1, "T"), (5, "T")], "T"), + [0, 1, 1, 1, 0, 1, 0], ) def test_effect_size(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [4], effect_size=-2.5), + node_genetic_value(tree, [(4, "T")], "T", effect_size=-2.5), [-2.5, -2.5, 0, 0, -2.5, 0, 0], ) + def test_edge_level(self, tree, node_genetic_value): + # The value arriving at a node is credited to the edge above it, and + # the root has no edge. generate_balanced(4) has edges 0..3 into the + # leaves and edges 4 and 5 into nodes 4 and 5. + ts = one_site(tree, [(6, "T")]) + value = node_genetic_value(ts, causal_allele="T", level="edge") + expected = np.zeros(ts.num_edges) + for u in range(6): + expected[ts.first().edge(u)] = 1 + np.testing.assert_array_equal(value, expected) + class TestStarTree: """ - tskit.Tree.generate_star(5), in which node 6 is the virtual root. The - polytomy at the root makes the sibling walk iterate five times:: + tskit.Tree.generate_star(5). The polytomy at the root makes the descent + iterate over five child edges:: 5 +-+-+-+-+ @@ -185,29 +302,32 @@ def tree(self): return tskit.Tree.generate_star(5) def test_root(self, tree, node_genetic_value): - np.testing.assert_array_equal(node_genetic_value(tree, [5]), [1, 1, 1, 1, 1, 1]) + np.testing.assert_array_equal( + node_genetic_value(tree, [(5, "T")], "T"), [1, 1, 1, 1, 1, 1] + ) - def test_virtual_root(self, tree, node_genetic_value): - np.testing.assert_array_equal(node_genetic_value(tree, [6]), [1, 1, 1, 1, 1, 1]) + def test_ancestral_state_is_causal(self, tree, node_genetic_value): + np.testing.assert_array_equal(node_genetic_value(tree), [1, 1, 1, 1, 1, 1]) def test_mutation_on_root(self, tree, node_genetic_value): - # The root is the virtual root's only child, so a mutation there stops - # the traversal immediately. + # The root is the only child of the virtual root, so a mutation there + # takes the causal allele away from the whole tree. np.testing.assert_array_equal( - node_genetic_value(tree, [6], [5]), [0, 0, 0, 0, 0, 0] + node_genetic_value(tree, [(5, "G")]), [0, 0, 0, 0, 0, 0] ) - def test_mutations_along_sibling_chain(self, tree, node_genetic_value): - # Mutations at both ends of the chain of children and in the middle. + def test_mutations_across_the_polytomy(self, tree, node_genetic_value): + # Mutations at both ends of the run of children and in the middle. np.testing.assert_array_equal( - node_genetic_value(tree, [5], [0, 2, 4]), [0, 1, 0, 1, 0, 1] + node_genetic_value(tree, [(5, "T"), (0, "G"), (2, "G"), (4, "G")], "T"), + [0, 1, 0, 1, 0, 1], ) class TestNonBinaryTree: """ - tskit.Tree.generate_balanced(6, arity=3), in which node 10 is the virtual - root. The polytomy is below the root:: + tskit.Tree.generate_balanced(6, arity=3), with the polytomy below the + root:: 9 +---+---+ @@ -220,31 +340,34 @@ class TestNonBinaryTree: def tree(self): return tskit.Tree.generate_balanced(6, arity=3) - def test_virtual_root(self, tree, node_genetic_value): + def test_ancestral_state_is_causal(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [10]), [1, 1, 1, 1, 1, 1, 1, 1, 1, 1] + node_genetic_value(tree), [1, 1, 1, 1, 1, 1, 1, 1, 1, 1] ) def test_middle_subtree(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [7]), [0, 0, 1, 1, 0, 0, 0, 1, 0, 0] + node_genetic_value(tree, [(7, "T")], "T"), + [0, 0, 1, 1, 0, 0, 0, 1, 0, 0], ) def test_mutation_on_middle_subtree(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [10], [7]), [1, 1, 0, 0, 1, 1, 1, 0, 1, 1] + node_genetic_value(tree, [(7, "G")]), + [1, 1, 0, 0, 1, 1, 1, 0, 1, 1], ) def test_mutations_on_outer_subtrees(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [10], [6, 8]), [0, 0, 1, 1, 0, 0, 0, 1, 0, 1] + node_genetic_value(tree, [(6, "G"), (8, "G")]), + [0, 0, 1, 1, 0, 0, 0, 1, 0, 1], ) class TestCombTree: """ - tskit.Tree.generate_comb(5), in which node 9 is the virtual root. A ladder - is the deepest traversal for a given number of leaves:: + tskit.Tree.generate_comb(5). A ladder is the deepest descent for a given + number of leaves:: 8 +-+-+ @@ -263,65 +386,64 @@ def tree(self): def test_root(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [8]), [1, 1, 1, 1, 1, 1, 1, 1, 1] + node_genetic_value(tree, [(8, "T")], "T"), [1, 1, 1, 1, 1, 1, 1, 1, 1] ) - def test_virtual_root(self, tree, node_genetic_value): + def test_ancestral_state_is_causal(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [9]), [1, 1, 1, 1, 1, 1, 1, 1, 1] + node_genetic_value(tree), [1, 1, 1, 1, 1, 1, 1, 1, 1] ) def test_second_rung(self, tree, node_genetic_value): # Everything below node 7, which is all of the tree except leaf 0 and # the root. np.testing.assert_array_equal( - node_genetic_value(tree, [7]), [0, 1, 1, 1, 1, 1, 1, 1, 0] + node_genetic_value(tree, [(7, "T")], "T"), [0, 1, 1, 1, 1, 1, 1, 1, 0] ) def test_mutation_near_root(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [8], [7]), [1, 0, 0, 0, 0, 0, 0, 0, 1] + node_genetic_value(tree, [(7, "G")]), [1, 0, 0, 0, 0, 0, 0, 0, 1] ) def test_mutation_near_leaves(self, tree, node_genetic_value): np.testing.assert_array_equal( - node_genetic_value(tree, [8], [5]), [1, 1, 1, 0, 0, 0, 1, 1, 1] + node_genetic_value(tree, [(5, "G")]), [1, 1, 1, 0, 0, 0, 1, 1, 1] ) class TestDegenerateTrees: """ - Trees that do not have a single root above a set of internal nodes. + Tree sequences that do not have a single root above a set of internal + nodes. """ def test_single_node(self, node_genetic_value): - # generate_balanced(1) is the single node 0, with virtual root 1. + # generate_balanced(1) is the single node 0, which is its own root. tree = tskit.Tree.generate_balanced(1) - np.testing.assert_array_equal(node_genetic_value(tree, [1]), [1]) - np.testing.assert_array_equal(node_genetic_value(tree, [0]), [1]) + np.testing.assert_array_equal(node_genetic_value(tree), [1]) + np.testing.assert_array_equal(node_genetic_value(tree, [(0, "T")], "T"), [1]) def test_no_nodes(self, node_genetic_value): - # The virtual root is node 0 and there is nothing to return. - tree = empty_tree() - assert tree.virtual_root == 0 - np.testing.assert_array_equal(node_genetic_value(tree, [0]), []) + np.testing.assert_array_equal(node_genetic_value(empty_ts()), []) def test_isolated_samples(self, node_genetic_value): - # Nodes 0, 1 and 2 are all roots, so the virtual root has three - # children: + # Nodes 0, 1 and 2 are all roots: # # 0 1 2 # - tree = isolated_samples_tree(3) - np.testing.assert_array_equal(node_genetic_value(tree, [3]), [1, 1, 1]) + np.testing.assert_array_equal( + node_genetic_value(isolated_samples_ts(3)), [1, 1, 1] + ) def test_isolated_samples_with_mutation(self, node_genetic_value): - tree = isolated_samples_tree(3) - np.testing.assert_array_equal(node_genetic_value(tree, [3], [1]), [1, 0, 1]) + np.testing.assert_array_equal( + node_genetic_value(isolated_samples_ts(3, [(1, "G")])), [1, 0, 1] + ) def test_multiple_roots(self, node_genetic_value): - # Nodes 4 and 5 are both roots and node 6 is isolated, so it is not - # reached from the virtual root: + # Nodes 4 and 5 are both roots and node 6 is isolated, so it is never + # reached: # # 4 5 # +++ +++ @@ -329,45 +451,38 @@ def test_multiple_roots(self, node_genetic_value): # tree = multiroot_tree() assert tree.roots == [4, 5] - np.testing.assert_array_equal( - node_genetic_value(tree, [7]), [1, 1, 1, 1, 1, 1, 0] - ) + np.testing.assert_array_equal(node_genetic_value(tree), [1, 1, 1, 1, 1, 1, 0]) def test_multiple_roots_one_blocked(self, node_genetic_value): - tree = multiroot_tree() np.testing.assert_array_equal( - node_genetic_value(tree, [7], [4]), [0, 0, 1, 1, 0, 1, 0] + node_genetic_value(multiroot_tree(), [(4, "G")]), [0, 0, 1, 1, 0, 1, 0] ) class TestAccumulateIndividualValues: """ Sum the node genetic values over the nodes of each individual, using the - tree sequences in tests/data.py. The node values are powers of two, so the - individual values identify the nodes that contributed to them. + tree sequences of tests.data. """ def test_binary_tree(self, individual_genetic_value): - # Individual 0 is nodes 4 and 5, individual 1 is nodes 0 and 1, and - # individual 2 is nodes 2 and 3. Node 6 has no individual. + # Individual 0 is nodes 4 and 5, individual 1 nodes 0 and 1, and + # individual 2 nodes 2 and 3. ts = binary_tree() - np.testing.assert_array_equal(ts.nodes_individual, [1, 1, 2, 2, 0, 0, -1]) np.testing.assert_array_equal( individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64]), [48, 3, 12] ) def test_diff_ind_tree(self, individual_genetic_value): - # The same tree, with the leaves paired up the other way around. ts = diff_ind_tree() - np.testing.assert_array_equal(ts.nodes_individual, [1, 2, 1, 2, 0, 0, -1]) np.testing.assert_array_equal( individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64]), [48, 5, 10] ) def test_triploid_tree(self, individual_genetic_value): - # Two triploids: nodes 0, 2 and 4, and nodes 1, 3 and 5. + # Individual 0 is nodes 0, 2 and 4, and individual 1 is nodes 1, 3 + # and 5. ts = triploid_tree() - np.testing.assert_array_equal(ts.nodes_individual, [0, 1, 0, 1, 0, 1, -1, -1]) np.testing.assert_array_equal( individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64, 128]), [21, 42] ) @@ -376,7 +491,7 @@ def test_no_individuals(self, individual_genetic_value): ts = tskit.Tree.generate_balanced(4).tree_sequence assert ts.num_individuals == 0 np.testing.assert_array_equal( - individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64]), [] + individual_genetic_value(ts, np.ones(ts.num_nodes)), [] ) @@ -398,17 +513,19 @@ class TestNodeAndIndividualValues: def test_internal_node(self, node_genetic_value, individual_genetic_value): ts = binary_tree() - value = node_genetic_value(ts.first(), [4]) + value = node_genetic_value(one_site(ts.first(), [(4, "T")]), causal_allele="T") np.testing.assert_array_equal(value, [1, 1, 0, 0, 1, 0, 0]) # Individual 0 has one copy through node 4, individual 1 has two # copies through nodes 0 and 1, and individual 2 has none. np.testing.assert_array_equal(individual_genetic_value(ts, value), [1, 2, 0]) - def test_virtual_root(self, node_genetic_value, individual_genetic_value): + def test_ancestral_state_is_causal( + self, node_genetic_value, individual_genetic_value + ): # The causal allele is the ancestral state, so every node carries it # and every diploid has two copies. ts = binary_tree() - value = node_genetic_value(ts.first(), [7]) + value = node_genetic_value(one_site(ts.first(), [])) np.testing.assert_array_equal(value, [1, 1, 1, 1, 1, 1, 1]) np.testing.assert_array_equal(individual_genetic_value(ts, value), [2, 2, 2]) @@ -423,7 +540,7 @@ def test_triploid(self, node_genetic_value, individual_genetic_value): # 0 1 2 3 4 5 # ts = triploid_tree() - value = node_genetic_value(ts.first(), [6]) + value = node_genetic_value(one_site(ts.first(), [(6, "T")]), causal_allele="T") np.testing.assert_array_equal(value, [0, 0, 0, 1, 1, 1, 1, 0]) # Individual 0 is nodes 0, 2 and 4, and individual 1 is nodes 1, 3 # and 5. diff --git a/tstrait/genetic_value.py b/tstrait/genetic_value.py index 382eb61..eddc6fa 100644 --- a/tstrait/genetic_value.py +++ b/tstrait/genetic_value.py @@ -1,24 +1,117 @@ import numpy as np import pandas as pd import tskit +import tskit.jit.numba as tskit_numba from . import jit from .base import _check_dataframe, _check_instance, _check_non_decreasing # noreorder -def _accumulate_edge_values(nodes_genetic_value, nodes_edge, num_nodes, num_edges): +def _causal_mutations(ts, trait_df): """ - Accumulate the edge genetic values by summing their node contributions. + Expand the trait dataframe to one entry per (row, mutation at that row's + site) pair, and return the row, the mutation and the change in causal + allele state for the mutations that change the state. Mutations that leave + the state unchanged contribute nothing and are dropped. """ - nodes_edge = nodes_edge[:num_nodes] - nodes_genetic_value = nodes_genetic_value[:num_nodes] - has_edge = nodes_edge != tskit.NULL - return np.bincount( - nodes_edge[has_edge], - weights=nodes_genetic_value[has_edge], - minlength=num_edges, + site_id = trait_df["site_id"].to_numpy() + # Mutations are sorted by site, so each site owns a contiguous run of IDs. + offset = np.searchsorted(ts.mutations_site, np.arange(ts.num_sites + 1)) + start = offset[site_id] + count = offset[site_id + 1] - start + row = np.repeat(np.arange(len(trait_df)), count) + mutation = np.arange(count.sum()) + np.repeat( + start - (np.cumsum(count) - count), count ) + derived_state = ts.mutations_derived_state + causal_allele = np.asarray( + trait_df["causal_allele"].to_numpy(), dtype=derived_state.dtype + )[row] + had_causal_allele = ts.mutations_inherited_state[mutation] == causal_allele + has_causal_allele = derived_state[mutation] == causal_allele + state_change = has_causal_allele.astype(np.int8) - had_causal_allele.astype(np.int8) + + changed = state_change != 0 + return row[changed], mutation[changed], state_change[changed] + + +def _root_runs(ts): + """ + Return the roots of the tree sequence as ``(node, left, right)`` arrays, + where each root spans a maximal interval over which it is a root. + + The roots are taken as the children of the virtual root, so that a node + counts as being in a tree exactly when the tree based implementation would + have reached it from the virtual root. Isolated samples are roots. + """ + node = [] + left = [] + right = [] + tree = tskit.Tree(ts) + left_child_array = tree.left_child_array + right_sib_array = tree.right_sib_array + virtual_root = tree.virtual_root + tree.first() + while True: + interval_left, interval_right = tree.interval + u = left_child_array[virtual_root] + while u != tskit.NULL: + if len(node) > 0 and node[-1] == u and right[-1] == interval_left: + right[-1] = interval_right + else: + node.append(u) + left.append(interval_left) + right.append(interval_right) + u = right_sib_array[u] + if not tree.next(): + break + return ( + np.array(node, dtype=np.int32), + np.array(left, dtype=float), + np.array(right, dtype=float), + ) + + +def _group_by(key, num_key, weight, site): + """ + Group the causal mutation weights by ``key``, returning the offsets of each + key's group, the causal site index of each entry sorted within its group, + and the cumulative sum of the weights with a leading zero. + """ + order = np.argsort(key, kind="stable") + offset = np.searchsorted(key[order], np.arange(num_key + 1)).astype(np.int32) + weight_sum = np.zeros(len(order) + 1) + np.cumsum(weight[order], out=weight_sum[1:]) + return offset, site[order].astype(np.int32), weight_sum + + +def _root_mutation_groups( + node, site, weight, roots, roots_site_start, roots_site_stop, num_causal_site +): + """ + Group the causal mutations that sit above a root by the root run they + belong to. + + A mutation with no edge is above a root, so its effect applies to that root + and to everything below it. Root runs for a given node are disjoint, so the + run a mutation belongs to is the one for its node containing its site. + """ + run = np.full(len(node), tskit.NULL, dtype=np.int32) + if len(node) > 0: + # Runs keyed by node and start, so that one search finds the last run + # of the mutation's node that starts at or before the mutation's site. + scale = num_causal_site + 1 + order = np.lexsort((roots_site_start, roots)) + key = roots[order] * scale + roots_site_start[order] + found = np.searchsorted(key, node * scale + site, side="right") - 1 + valid = found >= 0 + candidate = order[np.where(valid, found, 0)] + valid &= roots[candidate] == node + valid &= site < roots_site_stop[candidate] + run[valid] = candidate[valid] + return _group_by(run, len(roots), weight, site) + def _check_trait_df(ts, trait_df): """ @@ -47,6 +140,15 @@ class _GeneticValue: """ GeneticValue class to compute genetic values of individuals, nodes, or edges. + The genetic values of every causal site are accumulated in a single descent + of the ARG. Each causal mutation changes the causal allele state by + ``[derived == causal] - [inherited == causal]``, and that change applies to + every node below it. These changes telescope down a path, so a node's value + is the sum of the changes on the mutations above it, plus the effect size + of every causal site whose ancestral state is itself the causal allele. + That makes the value additive along a root to node path, which is what the + descent accumulates. Nested and back mutations need no special handling. + Parameters ---------- ts : tskit.TreeSequence @@ -59,32 +161,118 @@ class _GeneticValue: def __init__(self, ts, trait_df): self.trait_df = trait_df[["site_id", "effect_size", "trait_id", "causal_allele"]] self.ts = ts + self.child_index = tskit_numba.jitwrap(ts).child_index() + + site_id = self.trait_df["site_id"].to_numpy() + self.effect_size = self.trait_df["effect_size"].to_numpy() + self.trait_id = self.trait_df["trait_id"].to_numpy() + self.num_trait = np.max(self.trait_id) + 1 + + # Everything below indexes positions by their rank among the causal + # sites, so that the descent prunes on integers and never considers a + # position that cannot matter. + causal_site = np.unique(site_id) + causal_position = ts.sites_position[causal_site] + self.num_causal_site = len(causal_site) + self.rows_site = np.searchsorted(causal_site, site_id) + self.edges_site_start = np.searchsorted(causal_position, ts.edges_left).astype( + np.int32 + ) + self.edges_site_stop = np.searchsorted(causal_position, ts.edges_right).astype( + np.int32 + ) + + self.row, self.mutation, state_change = _causal_mutations(ts, self.trait_df) + self.weight = self.effect_size[self.row] * state_change + self.mutations_edge = ts.mutations_edge[self.mutation] + # A mutation above a root has no edge to sit on, so its effect applies + # to the root itself rather than to an edge. + self.on_edge = self.mutations_edge != tskit.NULL + + # An effect size counts towards every node in a tree when the ancestral + # state of its site is the causal allele. + self.ancestral_is_causal = ( + ts.sites_ancestral_state[site_id] + == self.trait_df["causal_allele"].to_numpy() + ) + self.roots, roots_left, roots_right = _root_runs(ts) + self.roots_site_start = np.searchsorted(causal_position, roots_left).astype( + np.int32 + ) + self.roots_site_stop = np.searchsorted(causal_position, roots_right).astype( + np.int32 + ) - def _node_genetic_values(self, tree, site, causal_allele, effect_size): + def _descent_arguments(self, level, trait): """ - Returns a numpy array with node genetic values. + Return the arguments to the descent kernel for one trait, with the + contributions directed at nodes or at edges according to ``level``. """ - has_mutation = np.zeros(self.ts.num_nodes + 1, dtype=bool) - state_transitions = {tree.virtual_root: site.ancestral_state} - for m in site.mutations: - state_transitions[m.node] = m.derived_state - has_mutation[m.node] = True - causal_nodes = np.array( - [ - node - for node, allele in state_transitions.items() - if allele == causal_allele - ], - dtype=np.int32, + ts = self.ts + if level == "edge": + edges_output = np.arange(ts.num_edges, dtype=np.int32) + roots_output = np.full(len(self.roots), tskit.NULL, dtype=np.int32) + output_size = ts.num_edges + else: + edges_output = ts.edges_child + roots_output = self.roots + output_size = ts.num_nodes + + in_trait = self.trait_id[self.row] == trait + on_edge = in_trait & self.on_edge + above_root = in_trait & ~self.on_edge + + mutations_offset, mutations_site, mutations_weight_sum = _group_by( + self.mutations_edge[on_edge], + ts.num_edges, + self.weight[on_edge], + self.rows_site[self.row[on_edge]], + ) + ( + roots_mutations_offset, + roots_mutations_site, + roots_mutations_weight_sum, + ) = _root_mutation_groups( + ts.mutations_node[self.mutation[above_root]], + self.rows_site[self.row[above_root]], + self.weight[above_root], + self.roots, + self.roots_site_start, + self.roots_site_stop, + self.num_causal_site, ) - return jit._compute_nodes_genetic_value( - left_child_array=tree.left_child_array, - right_sib_array=tree.right_sib_array, - causal_nodes=causal_nodes, - has_mutation=has_mutation, - effect_size=effect_size, + + ancestral = self.ancestral_is_causal & (self.trait_id == trait) + ancestral_sum = np.zeros(self.num_causal_site + 1) + np.cumsum( + np.bincount( + self.rows_site[ancestral], + weights=self.effect_size[ancestral], + minlength=self.num_causal_site, + ), + out=ancestral_sum[1:], ) + return { + "child_index": self.child_index, + "edges_child": ts.edges_child, + "edges_output": edges_output, + "edges_site_start": self.edges_site_start, + "edges_site_stop": self.edges_site_stop, + "mutations_offset": mutations_offset, + "mutations_site": mutations_site, + "mutations_weight_sum": mutations_weight_sum, + "ancestral_sum": ancestral_sum, + "roots": self.roots, + "roots_output": roots_output, + "roots_site_start": self.roots_site_start, + "roots_site_stop": self.roots_site_stop, + "roots_mutations_offset": roots_mutations_offset, + "roots_mutations_site": roots_mutations_site, + "roots_mutations_weight_sum": roots_mutations_weight_sum, + "output": np.zeros(output_size), + } + def _run(self, level): """ Computes genetic values of individuals, nodes, or edges @@ -95,48 +283,30 @@ def _run(self, level): pandas.DataFrame Dataframe with trait ID, [individual|node|edge] ID, and genetic value. """ - ts = self.ts - size_map = { + N = { "individual": ts.num_individuals, "node": ts.num_nodes, "edge": ts.num_edges, - } - N = size_map[level] - - num_trait = np.max(self.trait_df.trait_id) + 1 - genetic_value_table = np.zeros((num_trait, N)) - tree = tskit.Tree(self.ts) - - for data in self.trait_df.itertuples(): - site = self.ts.site(data.site_id) - tree.seek(site.position) - genetic_value = self._node_genetic_values( - tree=tree, - site=site, - causal_allele=data.causal_allele, - effect_size=data.effect_size, - ) + }[level] + + genetic_value_table = np.zeros((self.num_trait, N)) + for trait in range(self.num_trait): + output = jit._descend_arg(**self._descent_arguments(level, trait)) if level == "individual": - genetic_value = jit._accumulate_individual_values( - genetic_value, ts.nodes_individual, ts.num_individuals + output = jit._accumulate_individual_values( + output, ts.nodes_individual, ts.num_individuals ) - elif level == "edge": - genetic_value = _accumulate_edge_values( - genetic_value, tree.edge_array, ts.num_nodes, ts.num_edges - ) - genetic_value_table[data.trait_id, :] += genetic_value + genetic_value_table[trait, :] = output - df = pd.DataFrame( + return pd.DataFrame( { - "trait_id": np.repeat(np.arange(num_trait), N), - f"{level}_id": np.tile(np.arange(N), num_trait), + "trait_id": np.repeat(np.arange(self.num_trait), N), + f"{level}_id": np.tile(np.arange(N), self.num_trait), "genetic_value": genetic_value_table.flatten(), } ) - return df - def genetic_value(ts, trait_df, level="individual"): """ @@ -231,34 +401,10 @@ def edge_effect(ts, trait_df): N = ts.num_edges num_trait = np.max(trait_df.trait_id) + 1 - site_id = trait_df["site_id"].to_numpy() effect_size = trait_df["effect_size"].to_numpy() trait_id = trait_df["trait_id"].to_numpy() - # Mutations are sorted by site, so each site owns a contiguous run of IDs. - offset = np.searchsorted(ts.mutations_site, np.arange(ts.num_sites + 1)) - start = offset[site_id] - count = offset[site_id + 1] - start - # Expand to one entry per (trait_df row, mutation at that row's site) pair. - row = np.repeat(np.arange(len(trait_df)), count) - mutation = np.arange(count.sum()) + np.repeat( - start - (np.cumsum(count) - count), count - ) - - derived_state = ts.mutations_derived_state - causal_allele = np.asarray( - trait_df["causal_allele"].to_numpy(), dtype=derived_state.dtype - )[row] - had_causal_allele = ts.mutations_inherited_state[mutation] == causal_allele - has_causal_allele = derived_state[mutation] == causal_allele - state_change = has_causal_allele.astype(np.int8) - had_causal_allele.astype(np.int8) - - # Mutations that do not change the causal allele state contribute nothing, and - # are allowed to sit above a root. - changed = state_change != 0 - row = row[changed] - mutation = mutation[changed] - state_change = state_change[changed] + row, mutation, state_change = _causal_mutations(ts, trait_df) edge = ts.mutations_edge[mutation] if np.any(edge == tskit.NULL): bad_mutation = mutation[edge == tskit.NULL][0] diff --git a/tstrait/jit.py b/tstrait/jit.py index 20929c3..ec7a67b 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -3,9 +3,11 @@ Two conventions apply throughout this module: -1. Node indexed arrays follow the tskit quintuply linked tree encoding, in - which the arrays have ``num_nodes + 1`` entries and the last entry - corresponds to the virtual root. +1. Genetic values are accumulated by descending the ARG rather than by + working through the trees one at a time. Positions are expressed as + indexes into the sorted array of causal sites, so that the descent can + prune whole subtrees with integer comparisons and never looks at a + genomic coordinate. 2. Numba compiles with bounds checking disabled by default, so an out of bounds access here fails silently rather than raising ``IndexError``. @@ -19,45 +21,228 @@ import numpy as np import tskit +# The stacks start at this size and double whenever they fill up. A few +# hundred entries is typical, so this is usually the only allocation. +INITIAL_STACK_SIZE = 1024 + + +@numba.njit +def _search_sorted(values, start, stop, key): + """ + Return the first index in ``[start, stop)`` at which ``values`` is greater + than or equal to ``key``. The values in that range must be sorted. + """ + while start < stop: + mid = (start + stop) // 2 + if values[mid] < key: + start = mid + 1 + else: + stop = mid + return start + + +@numba.njit +def _group_weight(offset, site, weight_sum, group, start, stop): + """ + Return the total weight of the entries in ``group`` whose causal site index + is in ``[start, stop)``. Entries are grouped by ``offset`` and sorted by + site within each group, and ``weight_sum`` is their cumulative sum with a + leading zero. + """ + group_start = offset[group] + group_stop = offset[group + 1] + if group_start == group_stop: + return 0.0 + first = _search_sorted(site, group_start, group_stop, start) + last = _search_sorted(site, group_start, group_stop, stop) + return weight_sum[last] - weight_sum[first] + @numba.njit -def _compute_nodes_genetic_value( - left_child_array, - right_sib_array, - causal_nodes, - has_mutation, - effect_size, +def _descend_arg( + child_index, + edges_child, + edges_output, + edges_site_start, + edges_site_stop, + mutations_offset, + mutations_site, + mutations_weight_sum, + ancestral_sum, + roots, + roots_output, + roots_site_start, + roots_site_stop, + roots_mutations_offset, + roots_mutations_site, + roots_mutations_weight_sum, + output, ): """ - Compute the genetic value of each node by assigning ``effect_size`` to - every node in ``causal_nodes`` and to each of their descendants that is - not separated from them by a node carrying a mutation. - - The ``left_child_array``, ``right_sib_array`` and ``has_mutation`` inputs - are indexed by node ID and have ``num_nodes + 1`` entries. The - ``causal_nodes`` are distinct node IDs, and may include the virtual root, - which happens when the ancestral state is the causal allele. The returned - array has one entry per node and does not include the virtual root. + Accumulate the genetic value of every causal site in one descent of the ARG. + + Each entry on the stack is a node together with the half open range of + causal sites over which its path back to the root is unchanged, and the + genetic value accumulated along that path. Descending an edge intersects + the range with the edge's own range of causal sites, so a subtree spanning + no causal site is never entered, and adds the effects of the causal + mutations the edge carries within the intersected range. + + A node is reached once per such range, and the ranges partition the causal + sites the node spans, so the values accumulate to the total over every + causal site. + + The value carried down is a sum over the sites in the current range, so + when an edge narrows the range it is no longer the value the child should + inherit: the mutations above at the sites that dropped out do not apply to + the child. The path back to the root is therefore kept in ``path_edge``, and + the value is recomputed over the narrowed range whenever that happens. This + is rare, a hundredth of a percent of descents on a large tree sequence, + because an edge usually spans the whole of the range reaching it. + + ``child_index`` is the tskit child index, in which a node that is never a + parent has the range ``(-1, -1)``. ``ancestral_sum`` is the cumulative + effect size of the causal sites whose ancestral state is the causal allele, + which every node in a tree carries. ``edges_output`` and ``roots_output`` + give the index in ``output`` that each contribution is added to, which is + how the same descent serves both node and edge genetic values; a negative + index discards the contribution. """ - num_nodes = len(left_child_array) - 1 - # One extra entry, so that the virtual root can be a causal node. - genetic_value = np.zeros(num_nodes + 1) - # Each node is pushed at most once, since causal nodes other than the - # virtual root all carry a mutation and so are never pushed as a child. - stack = np.empty(num_nodes + 1, dtype=np.int32) - stack_top = len(causal_nodes) - stack[:stack_top] = causal_nodes - while stack_top > 0: - stack_top -= 1 - parent_node_id = stack[stack_top] - genetic_value[parent_node_id] = effect_size - child_node_id = left_child_array[parent_node_id] - while child_node_id != tskit.NULL: - if not has_mutation[child_node_id]: - stack[stack_top] = child_node_id - stack_top += 1 - child_node_id = right_sib_array[child_node_id] - return genetic_value[:num_nodes] + stack_size = INITIAL_STACK_SIZE + stack_node = np.empty(stack_size, dtype=np.int32) + stack_start = np.empty(stack_size, dtype=np.int32) + stack_stop = np.empty(stack_size, dtype=np.int32) + stack_value = np.empty(stack_size, dtype=np.float64) + stack_depth = np.empty(stack_size, dtype=np.int32) + stack_edge = np.empty(stack_size, dtype=np.int32) + path_size = INITIAL_STACK_SIZE + path_edge = np.empty(path_size, dtype=np.int32) + + for j in range(len(roots)): + root_start = roots_site_start[j] + root_stop = roots_site_stop[j] + if root_start >= root_stop: + continue + root = roots[j] + root_value = ( + ancestral_sum[root_stop] + - ancestral_sum[root_start] + + _group_weight( + roots_mutations_offset, + roots_mutations_site, + roots_mutations_weight_sum, + j, + root_start, + root_stop, + ) + ) + if roots_output[j] >= 0: + output[roots_output[j]] += root_value + if child_index[root, 0] < 0: + continue + stack_node[0] = root + stack_start[0] = root_start + stack_stop[0] = root_stop + stack_value[0] = root_value + stack_depth[0] = 0 + stack_edge[0] = tskit.NULL + top = 1 + + while top > 0: + top -= 1 + parent = stack_node[top] + start = stack_start[top] + stop = stack_stop[top] + parent_value = stack_value[top] + depth = stack_depth[top] + # Depth first order guarantees that the entries below this one in + # path_edge are the edges back to the root. + if depth > 0: + path_edge[depth - 1] = stack_edge[top] + + for e in range(child_index[parent, 0], child_index[parent, 1]): + edge_start = edges_site_start[e] + child_start = start if start > edge_start else edge_start + edge_stop = edges_site_stop[e] + child_stop = stop if stop < edge_stop else edge_stop + if child_start >= child_stop: + continue + + if child_start == start and child_stop == stop: + child_value = parent_value + else: + child_value = ( + ancestral_sum[child_stop] + - ancestral_sum[child_start] + + _group_weight( + roots_mutations_offset, + roots_mutations_site, + roots_mutations_weight_sum, + j, + child_start, + child_stop, + ) + ) + for k in range(depth): + child_value += _group_weight( + mutations_offset, + mutations_site, + mutations_weight_sum, + path_edge[k], + child_start, + child_stop, + ) + child_value += _group_weight( + mutations_offset, + mutations_site, + mutations_weight_sum, + e, + child_start, + child_stop, + ) + + if edges_output[e] >= 0: + output[edges_output[e]] += child_value + + # A node that is never a parent has nothing below it, so there + # is no point putting it on the stack. + child = edges_child[e] + if child_index[child, 0] < 0: + continue + if top == stack_size: + stack_size *= 2 + grown_node = np.empty(stack_size, dtype=np.int32) + grown_start = np.empty(stack_size, dtype=np.int32) + grown_stop = np.empty(stack_size, dtype=np.int32) + grown_value = np.empty(stack_size, dtype=np.float64) + grown_depth = np.empty(stack_size, dtype=np.int32) + grown_edge = np.empty(stack_size, dtype=np.int32) + grown_node[: len(stack_node)] = stack_node + grown_start[: len(stack_start)] = stack_start + grown_stop[: len(stack_stop)] = stack_stop + grown_value[: len(stack_value)] = stack_value + grown_depth[: len(stack_depth)] = stack_depth + grown_edge[: len(stack_edge)] = stack_edge + stack_node = grown_node + stack_start = grown_start + stack_stop = grown_stop + stack_value = grown_value + stack_depth = grown_depth + stack_edge = grown_edge + if depth + 1 == path_size: + path_size *= 2 + grown_path = np.empty(path_size, dtype=np.int32) + grown_path[: len(path_edge)] = path_edge + path_edge = grown_path + stack_node[top] = child + stack_start[top] = child_start + stack_stop[top] = child_stop + stack_value[top] = child_value + stack_depth[top] = depth + 1 + stack_edge[top] = e + top += 1 + + return output @numba.njit From ba5b2a63dfff00754d71d4ed6000b88d360102d9 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Wed, 2 Sep 2026 16:12:14 +0100 Subject: [PATCH 03/19] Push causal mutation effects down the ARG in time order Replace the range carrying descent with a single sweep of the nodes from the past to the present. Each node conceptually holds a set of (causal site, effect) tuples, seeded at the node of each causal mutation. A parent is always older than its child, so time order is a topological order and every node is reached after all of its ancestors: for each tuple the outbound edges spanning its causal site are found, and the tuple is credited to the edge and to the node at the far end before being passed on to it. Nodes are credited when a tuple arrives rather than when they are swept, so a node that is never a parent never holds a tuple, which is about half of the work on a large tree sequence. Each tuple carries a single causal site rather than a range, so there is nothing to narrow and none of the path recomputation the descent needed. Tuples live in an arena threaded as a per node linked list, with a free list, since every tuple of a node is dead once that node has been swept and only those in flight need to be held. The causal allele being the ancestral state needs no separate treatment either: there is no mutation to seed from, so the roots are seeded instead and the effect reaches exactly the nodes of the tree. Mutations above a root need nothing at all, since seeding at the node and pushing down is already right. Cost is the number of nodes that carry a causal allele, rather than the size of the tree sequence, so the gain depends entirely on allele frequency. On 100k samples at level="node", against the tree by tree implementation, causal sites drawn from those below a frequency of 0.001 go from 0.436s to 0.060s at 1000 sites, 3.761s to 0.099s at 10,000, and 38.374s to 0.665s at 100,000. Causal sites drawn uniformly over all sites are around 2.8 times slower throughout, because the common variants in the tail of the frequency spectrum touch 7.4% of the nodes on average where the median site touches 0.39%, and a dense pass over every node is sequential where following carriers is not. Retarget the jit tests at the new kernel, dropping the binary search helpers that only existed for the range searches. --- tests/test_jit.py | 71 +-------- tstrait/genetic_value.py | 175 ++++++++------------- tstrait/jit.py | 321 ++++++++++++++------------------------- 3 files changed, 178 insertions(+), 389 deletions(-) diff --git a/tests/test_jit.py b/tests/test_jit.py index 0b517de..b36cce0 100644 --- a/tests/test_jit.py +++ b/tests/test_jit.py @@ -10,7 +10,8 @@ The examples are small trees from the tskit generators, drawn in the class docstrings so that the expected values can be checked against the topology. Each is turned into a tree sequence with a single causal site, since the -descent works from the mutations of the ARG rather than from a tree. +kernel pushes the effects of mutations down the ARG rather than working from +a tree. """ import numpy as np @@ -65,7 +66,7 @@ def node_genetic_value(request): state of "A", and a node's value is ``effect_size`` when the allele it inherits is ``causal_allele``. """ - func = kernel(jit._descend_arg, request.param) + func = kernel(jit._push_down_arg, request.param) def f( tree, @@ -140,72 +141,6 @@ def empty_ts(): return tables.tree_sequence() -class TestSearchSorted: - """ - The binary search the descent uses to find the mutations of an edge that - lie in a range of causal sites. - """ - - @pytest.fixture(params=["jit", "nojit"]) - def search_sorted(self, request): - return kernel(jit._search_sorted, request.param) - - @pytest.mark.parametrize( - ("key", "expected"), [(-1, 0), (0, 0), (1, 1), (3, 1), (4, 1), (5, 3), (9, 4), (10, 5)] - ) - def test_search(self, search_sorted, key, expected): - values = np.array([0, 4, 4, 7, 9], dtype=np.int32) - assert search_sorted(values, 0, len(values), key) == expected - - def test_restricted_range(self, search_sorted): - # Only the values in [1, 4) are searched, so a key below the range - # comes back as the start of it. - values = np.array([0, 4, 4, 7, 9], dtype=np.int32) - assert search_sorted(values, 1, 4, -1) == 1 - assert search_sorted(values, 1, 4, 8) == 4 - - def test_empty_range(self, search_sorted): - values = np.array([0, 4, 4, 7, 9], dtype=np.int32) - assert search_sorted(values, 2, 2, 4) == 2 - - -class TestGroupWeight: - """ - The total weight of the entries of one group that fall in a range of - causal sites. Group 0 holds sites 1 and 3, group 1 is empty and group 2 - holds sites 0, 2 and 2. - """ - - @pytest.fixture(params=["jit", "nojit"]) - def group_weight(self, request): - func = kernel(jit._group_weight, request.param) - offset = np.array([0, 2, 2, 5], dtype=np.int32) - site = np.array([1, 3, 0, 2, 2], dtype=np.int32) - weight_sum = np.array([0.0, 1.0, 3.0, 7.0, 15.0, 31.0]) - - def f(group, start, stop): - return func(offset, site, weight_sum, group, start, stop) - - return f - - def test_whole_group(self, group_weight): - assert group_weight(0, 0, 4) == 3.0 - assert group_weight(2, 0, 4) == 28.0 - - def test_empty_group(self, group_weight): - assert group_weight(1, 0, 4) == 0.0 - - def test_partial_range(self, group_weight): - assert group_weight(0, 0, 2) == 1.0 - assert group_weight(0, 2, 4) == 2.0 - # Both entries at site 2 are inside the range or outside it together. - assert group_weight(2, 2, 3) == 24.0 - assert group_weight(2, 0, 1) == 4.0 - - def test_range_outside_group(self, group_weight): - assert group_weight(0, 4, 8) == 0.0 - - class TestBalancedBinaryTree: """ tskit.Tree.generate_balanced(4):: diff --git a/tstrait/genetic_value.py b/tstrait/genetic_value.py index eddc6fa..1ff2f0b 100644 --- a/tstrait/genetic_value.py +++ b/tstrait/genetic_value.py @@ -73,46 +73,6 @@ def _root_runs(ts): ) -def _group_by(key, num_key, weight, site): - """ - Group the causal mutation weights by ``key``, returning the offsets of each - key's group, the causal site index of each entry sorted within its group, - and the cumulative sum of the weights with a leading zero. - """ - order = np.argsort(key, kind="stable") - offset = np.searchsorted(key[order], np.arange(num_key + 1)).astype(np.int32) - weight_sum = np.zeros(len(order) + 1) - np.cumsum(weight[order], out=weight_sum[1:]) - return offset, site[order].astype(np.int32), weight_sum - - -def _root_mutation_groups( - node, site, weight, roots, roots_site_start, roots_site_stop, num_causal_site -): - """ - Group the causal mutations that sit above a root by the root run they - belong to. - - A mutation with no edge is above a root, so its effect applies to that root - and to everything below it. Root runs for a given node are disjoint, so the - run a mutation belongs to is the one for its node containing its site. - """ - run = np.full(len(node), tskit.NULL, dtype=np.int32) - if len(node) > 0: - # Runs keyed by node and start, so that one search finds the last run - # of the mutation's node that starts at or before the mutation's site. - scale = num_causal_site + 1 - order = np.lexsort((roots_site_start, roots)) - key = roots[order] * scale + roots_site_start[order] - found = np.searchsorted(key, node * scale + site, side="right") - 1 - valid = found >= 0 - candidate = order[np.where(valid, found, 0)] - valid &= roots[candidate] == node - valid &= site < roots_site_stop[candidate] - run[valid] = candidate[valid] - return _group_by(run, len(roots), weight, site) - - def _check_trait_df(ts, trait_df): """ Check the trait dataframe against the tree sequence, returning the required @@ -164,112 +124,99 @@ def __init__(self, ts, trait_df): self.child_index = tskit_numba.jitwrap(ts).child_index() site_id = self.trait_df["site_id"].to_numpy() - self.effect_size = self.trait_df["effect_size"].to_numpy() + effect_size = self.trait_df["effect_size"].to_numpy() self.trait_id = self.trait_df["trait_id"].to_numpy() self.num_trait = np.max(self.trait_id) + 1 - # Everything below indexes positions by their rank among the causal - # sites, so that the descent prunes on integers and never considers a - # position that cannot matter. + # Causal sites are identified by their rank among the causal sites, so + # that matching an edge to one is an integer comparison and a position + # that cannot matter is never considered. causal_site = np.unique(site_id) causal_position = ts.sites_position[causal_site] - self.num_causal_site = len(causal_site) - self.rows_site = np.searchsorted(causal_site, site_id) + rows_site = np.searchsorted(causal_site, site_id) self.edges_site_start = np.searchsorted(causal_position, ts.edges_left).astype( np.int32 ) self.edges_site_stop = np.searchsorted(causal_position, ts.edges_right).astype( np.int32 ) - - self.row, self.mutation, state_change = _causal_mutations(ts, self.trait_df) - self.weight = self.effect_size[self.row] * state_change - self.mutations_edge = ts.mutations_edge[self.mutation] - # A mutation above a root has no edge to sit on, so its effect applies - # to the root itself rather than to an edge. - self.on_edge = self.mutations_edge != tskit.NULL - - # An effect size counts towards every node in a tree when the ancestral - # state of its site is the causal allele. - self.ancestral_is_causal = ( + # tskit does not require the node IDs to be in time order. + self.nodes_by_time = np.argsort(-ts.nodes_time, kind="stable").astype(np.int32) + + row, mutation, state_change = _causal_mutations(ts, self.trait_df) + self.seed_trait = self.trait_id[row] + self.seed_node = ts.mutations_node[mutation].astype(np.int32) + self.seed_site = rows_site[row].astype(np.int32) + self.seed_weight = effect_size[row] * state_change + # A mutation above a root has no edge, and needs no special handling: + # seeding at its node and pushing down is already right, and the + # missing edge contribution matches the tree based implementation. + self.seed_edge = ts.mutations_edge[mutation] + + # When the ancestral state of a site is the causal allele, every node + # in its tree carries it. There is no mutation to seed from, so the + # roots are seeded instead and the effect reaches the same nodes. + ancestral = np.flatnonzero( ts.sites_ancestral_state[site_id] == self.trait_df["causal_allele"].to_numpy() ) - self.roots, roots_left, roots_right = _root_runs(ts) - self.roots_site_start = np.searchsorted(causal_position, roots_left).astype( - np.int32 - ) - self.roots_site_stop = np.searchsorted(causal_position, roots_right).astype( - np.int32 - ) + if len(ancestral) > 0: + roots, roots_left, roots_right = _root_runs(ts) + start = np.searchsorted(causal_position, roots_left) + stop = np.searchsorted(causal_position, roots_right) + # Each root run takes the ancestral rows spanned by its interval. + ancestral_site = rows_site[ancestral] + first = np.searchsorted(ancestral_site, start) + count = np.searchsorted(ancestral_site, stop) - first + run = np.repeat(np.arange(len(roots)), count) + index = np.arange(count.sum()) + np.repeat( + first - (np.cumsum(count) - count), count + ) + self.seed_trait = np.concatenate( + [self.seed_trait, self.trait_id[ancestral[index]]] + ) + self.seed_node = np.concatenate([self.seed_node, roots[run]]) + self.seed_site = np.concatenate( + [self.seed_site, ancestral_site[index].astype(np.int32)] + ) + self.seed_weight = np.concatenate( + [self.seed_weight, effect_size[ancestral[index]]] + ) + # A root has no edge above it, so it contributes nothing at the + # edge level, which is what the tree based implementation does too. + self.seed_edge = np.concatenate( + [self.seed_edge, np.full(len(run), tskit.NULL, dtype=np.int32)] + ) def _descent_arguments(self, level, trait): """ - Return the arguments to the descent kernel for one trait, with the + Return the arguments to the push down kernel for one trait, with the contributions directed at nodes or at edges according to ``level``. """ ts = self.ts + in_trait = self.seed_trait == trait if level == "edge": edges_output = np.arange(ts.num_edges, dtype=np.int32) - roots_output = np.full(len(self.roots), tskit.NULL, dtype=np.int32) + # A seed is credited to the edge above the mutation, which does not + # exist when the mutation is above a root. + seed_output = self.seed_edge[in_trait].astype(np.int32) output_size = ts.num_edges else: edges_output = ts.edges_child - roots_output = self.roots + seed_output = self.seed_node[in_trait] output_size = ts.num_nodes - in_trait = self.trait_id[self.row] == trait - on_edge = in_trait & self.on_edge - above_root = in_trait & ~self.on_edge - - mutations_offset, mutations_site, mutations_weight_sum = _group_by( - self.mutations_edge[on_edge], - ts.num_edges, - self.weight[on_edge], - self.rows_site[self.row[on_edge]], - ) - ( - roots_mutations_offset, - roots_mutations_site, - roots_mutations_weight_sum, - ) = _root_mutation_groups( - ts.mutations_node[self.mutation[above_root]], - self.rows_site[self.row[above_root]], - self.weight[above_root], - self.roots, - self.roots_site_start, - self.roots_site_stop, - self.num_causal_site, - ) - - ancestral = self.ancestral_is_causal & (self.trait_id == trait) - ancestral_sum = np.zeros(self.num_causal_site + 1) - np.cumsum( - np.bincount( - self.rows_site[ancestral], - weights=self.effect_size[ancestral], - minlength=self.num_causal_site, - ), - out=ancestral_sum[1:], - ) - return { "child_index": self.child_index, "edges_child": ts.edges_child, "edges_output": edges_output, "edges_site_start": self.edges_site_start, "edges_site_stop": self.edges_site_stop, - "mutations_offset": mutations_offset, - "mutations_site": mutations_site, - "mutations_weight_sum": mutations_weight_sum, - "ancestral_sum": ancestral_sum, - "roots": self.roots, - "roots_output": roots_output, - "roots_site_start": self.roots_site_start, - "roots_site_stop": self.roots_site_stop, - "roots_mutations_offset": roots_mutations_offset, - "roots_mutations_site": roots_mutations_site, - "roots_mutations_weight_sum": roots_mutations_weight_sum, + "nodes_by_time": self.nodes_by_time, + "seed_node": self.seed_node[in_trait], + "seed_site": self.seed_site[in_trait], + "seed_weight": self.seed_weight[in_trait], + "seed_output": seed_output, "output": np.zeros(output_size), } @@ -292,7 +239,7 @@ def _run(self, level): genetic_value_table = np.zeros((self.num_trait, N)) for trait in range(self.num_trait): - output = jit._descend_arg(**self._descent_arguments(level, trait)) + output = jit._push_down_arg(**self._descent_arguments(level, trait)) if level == "individual": output = jit._accumulate_individual_values( output, ts.nodes_individual, ts.num_individuals diff --git a/tstrait/jit.py b/tstrait/jit.py index ec7a67b..ffdf0c6 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -3,11 +3,11 @@ Two conventions apply throughout this module: -1. Genetic values are accumulated by descending the ARG rather than by - working through the trees one at a time. Positions are expressed as - indexes into the sorted array of causal sites, so that the descent can - prune whole subtrees with integer comparisons and never looks at a - genomic coordinate. +1. Genetic values are accumulated by pushing the effect of each causal + mutation down the ARG, rather than by working through the trees one at a + time. Positions are expressed as indexes into the sorted array of causal + sites, so that an edge is matched to a causal site with an integer + comparison and a genomic coordinate never appears. 2. Numba compiles with bounds checking disabled by default, so an out of bounds access here fails silently rather than raising ``IndexError``. @@ -21,226 +21,133 @@ import numpy as np import tskit -# The stacks start at this size and double whenever they fill up. A few -# hundred entries is typical, so this is usually the only allocation. -INITIAL_STACK_SIZE = 1024 +# The arena starts at this many tuples and doubles whenever it fills up. +INITIAL_ARENA_SIZE = 1024 @numba.njit -def _search_sorted(values, start, stop, key): - """ - Return the first index in ``[start, stop)`` at which ``values`` is greater - than or equal to ``key``. The values in that range must be sorted. - """ - while start < stop: - mid = (start + stop) // 2 - if values[mid] < key: - start = mid + 1 - else: - stop = mid - return start - - -@numba.njit -def _group_weight(offset, site, weight_sum, group, start, stop): - """ - Return the total weight of the entries in ``group`` whose causal site index - is in ``[start, stop)``. Entries are grouped by ``offset`` and sorted by - site within each group, and ``weight_sum`` is their cumulative sum with a - leading zero. - """ - group_start = offset[group] - group_stop = offset[group + 1] - if group_start == group_stop: - return 0.0 - first = _search_sorted(site, group_start, group_stop, start) - last = _search_sorted(site, group_start, group_stop, stop) - return weight_sum[last] - weight_sum[first] - - -@numba.njit -def _descend_arg( +def _push_down_arg( child_index, edges_child, edges_output, edges_site_start, edges_site_stop, - mutations_offset, - mutations_site, - mutations_weight_sum, - ancestral_sum, - roots, - roots_output, - roots_site_start, - roots_site_stop, - roots_mutations_offset, - roots_mutations_site, - roots_mutations_weight_sum, + nodes_by_time, + seed_node, + seed_site, + seed_weight, + seed_output, output, ): """ - Accumulate the genetic value of every causal site in one descent of the ARG. - - Each entry on the stack is a node together with the half open range of - causal sites over which its path back to the root is unchanged, and the - genetic value accumulated along that path. Descending an edge intersects - the range with the edge's own range of causal sites, so a subtree spanning - no causal site is never entered, and adds the effects of the causal - mutations the edge carries within the intersected range. + Accumulate the genetic value of every causal site in one sweep of the nodes. - A node is reached once per such range, and the ranges partition the causal - sites the node spans, so the values accumulate to the total over every - causal site. + Each node conceptually holds a set of ``(causal site, effect)`` tuples, + seeded at the node of each causal mutation. Sweeping the nodes from the + past to the present visits every node after all of its ancestors, because + a parent is always older than its child, so a tuple can be pushed from a + node to its children and is guaranteed to have arrived before that child is + reached. For each tuple the outbound edges spanning its causal site are + found, and the tuple is credited to the edge and to the node at the far end + before being passed on to it. - The value carried down is a sum over the sites in the current range, so - when an edge narrows the range it is no longer the value the child should - inherit: the mutations above at the sites that dropped out do not apply to - the child. The path back to the root is therefore kept in ``path_edge``, and - the value is recomputed over the narrowed range whenever that happens. This - is rare, a hundredth of a percent of descents on a large tree sequence, - because an edge usually spans the whole of the range reaching it. + A node is credited when a tuple arrives rather than when it is swept, so a + node that is never a parent never holds a tuple. On a large tree sequence + that is about half of the work. ``child_index`` is the tskit child index, in which a node that is never a - parent has the range ``(-1, -1)``. ``ancestral_sum`` is the cumulative - effect size of the causal sites whose ancestral state is the causal allele, - which every node in a tree carries. ``edges_output`` and ``roots_output`` - give the index in ``output`` that each contribution is added to, which is - how the same descent serves both node and edge genetic values; a negative - index discards the contribution. + parent has the range ``(-1, -1)``. ``nodes_by_time`` lists the nodes from + the oldest to the youngest. ``edges_output`` and ``seed_output`` give the + index in ``output`` that each contribution is added to, which is how the + same sweep serves both node and edge genetic values; a negative index + discards the contribution. """ - stack_size = INITIAL_STACK_SIZE - stack_node = np.empty(stack_size, dtype=np.int32) - stack_start = np.empty(stack_size, dtype=np.int32) - stack_stop = np.empty(stack_size, dtype=np.int32) - stack_value = np.empty(stack_size, dtype=np.float64) - stack_depth = np.empty(stack_size, dtype=np.int32) - stack_edge = np.empty(stack_size, dtype=np.int32) - path_size = INITIAL_STACK_SIZE - path_edge = np.empty(path_size, dtype=np.int32) - - for j in range(len(roots)): - root_start = roots_site_start[j] - root_stop = roots_site_stop[j] - if root_start >= root_stop: + num_nodes = len(child_index) + head = np.full(num_nodes, tskit.NULL, dtype=np.int32) + arena_size = INITIAL_ARENA_SIZE + tuple_site = np.empty(arena_size, dtype=np.int32) + tuple_weight = np.empty(arena_size, dtype=np.float64) + tuple_next = np.empty(arena_size, dtype=np.int32) + # Slots are taken from the free list first, and from the top of the arena + # only when it is empty. Every tuple of a node is dead once that node has + # been swept, so the arena only ever holds the tuples still in flight. + arena_top = 0 + free = tskit.NULL + + for j in range(len(seed_node)): + u = seed_node[j] + weight = seed_weight[j] + if seed_output[j] >= 0: + output[seed_output[j]] += weight + if child_index[u, 0] < 0: continue - root = roots[j] - root_value = ( - ancestral_sum[root_stop] - - ancestral_sum[root_start] - + _group_weight( - roots_mutations_offset, - roots_mutations_site, - roots_mutations_weight_sum, - j, - root_start, - root_stop, - ) - ) - if roots_output[j] >= 0: - output[roots_output[j]] += root_value - if child_index[root, 0] < 0: + if free != tskit.NULL: + slot = free + free = tuple_next[slot] + else: + if arena_top == arena_size: + arena_size *= 2 + grown_site = np.empty(arena_size, dtype=np.int32) + grown_weight = np.empty(arena_size, dtype=np.float64) + grown_next = np.empty(arena_size, dtype=np.int32) + grown_site[:arena_top] = tuple_site + grown_weight[:arena_top] = tuple_weight + grown_next[:arena_top] = tuple_next + tuple_site = grown_site + tuple_weight = grown_weight + tuple_next = grown_next + slot = arena_top + arena_top += 1 + tuple_site[slot] = seed_site[j] + tuple_weight[slot] = weight + tuple_next[slot] = head[u] + head[u] = slot + + for i in range(len(nodes_by_time)): + parent = nodes_by_time[i] + slot = head[parent] + if slot == tskit.NULL: continue - stack_node[0] = root - stack_start[0] = root_start - stack_stop[0] = root_stop - stack_value[0] = root_value - stack_depth[0] = 0 - stack_edge[0] = tskit.NULL - top = 1 - - while top > 0: - top -= 1 - parent = stack_node[top] - start = stack_start[top] - stop = stack_stop[top] - parent_value = stack_value[top] - depth = stack_depth[top] - # Depth first order guarantees that the entries below this one in - # path_edge are the edges back to the root. - if depth > 0: - path_edge[depth - 1] = stack_edge[top] - - for e in range(child_index[parent, 0], child_index[parent, 1]): - edge_start = edges_site_start[e] - child_start = start if start > edge_start else edge_start - edge_stop = edges_site_stop[e] - child_stop = stop if stop < edge_stop else edge_stop - if child_start >= child_stop: - continue - - if child_start == start and child_stop == stop: - child_value = parent_value - else: - child_value = ( - ancestral_sum[child_stop] - - ancestral_sum[child_start] - + _group_weight( - roots_mutations_offset, - roots_mutations_site, - roots_mutations_weight_sum, - j, - child_start, - child_stop, - ) - ) - for k in range(depth): - child_value += _group_weight( - mutations_offset, - mutations_site, - mutations_weight_sum, - path_edge[k], - child_start, - child_stop, - ) - child_value += _group_weight( - mutations_offset, - mutations_site, - mutations_weight_sum, - e, - child_start, - child_stop, - ) - - if edges_output[e] >= 0: - output[edges_output[e]] += child_value - - # A node that is never a parent has nothing below it, so there - # is no point putting it on the stack. - child = edges_child[e] - if child_index[child, 0] < 0: - continue - if top == stack_size: - stack_size *= 2 - grown_node = np.empty(stack_size, dtype=np.int32) - grown_start = np.empty(stack_size, dtype=np.int32) - grown_stop = np.empty(stack_size, dtype=np.int32) - grown_value = np.empty(stack_size, dtype=np.float64) - grown_depth = np.empty(stack_size, dtype=np.int32) - grown_edge = np.empty(stack_size, dtype=np.int32) - grown_node[: len(stack_node)] = stack_node - grown_start[: len(stack_start)] = stack_start - grown_stop[: len(stack_stop)] = stack_stop - grown_value[: len(stack_value)] = stack_value - grown_depth[: len(stack_depth)] = stack_depth - grown_edge[: len(stack_edge)] = stack_edge - stack_node = grown_node - stack_start = grown_start - stack_stop = grown_stop - stack_value = grown_value - stack_depth = grown_depth - stack_edge = grown_edge - if depth + 1 == path_size: - path_size *= 2 - grown_path = np.empty(path_size, dtype=np.int32) - grown_path[: len(path_edge)] = path_edge - path_edge = grown_path - stack_node[top] = child - stack_start[top] = child_start - stack_stop[top] = child_stop - stack_value[top] = child_value - stack_depth[top] = depth + 1 - stack_edge[top] = e - top += 1 + head[parent] = tskit.NULL + edge_start = child_index[parent, 0] + edge_stop = child_index[parent, 1] + while slot != tskit.NULL: + site = tuple_site[slot] + weight = tuple_weight[slot] + for e in range(edge_start, edge_stop): + if edges_site_start[e] <= site and site < edges_site_stop[e]: + if edges_output[e] >= 0: + output[edges_output[e]] += weight + child = edges_child[e] + if child_index[child, 0] < 0: + # Nothing below, so there is no tuple to hold. + continue + if free != tskit.NULL: + child_slot = free + free = tuple_next[child_slot] + else: + if arena_top == arena_size: + arena_size *= 2 + grown_site = np.empty(arena_size, dtype=np.int32) + grown_weight = np.empty(arena_size, dtype=np.float64) + grown_next = np.empty(arena_size, dtype=np.int32) + grown_site[:arena_top] = tuple_site + grown_weight[:arena_top] = tuple_weight + grown_next[:arena_top] = tuple_next + tuple_site = grown_site + tuple_weight = grown_weight + tuple_next = grown_next + child_slot = arena_top + arena_top += 1 + tuple_site[child_slot] = site + tuple_weight[child_slot] = weight + tuple_next[child_slot] = head[child] + head[child] = child_slot + # This tuple is finished with, so its slot can be reused. + next_slot = tuple_next[slot] + tuple_next[slot] = free + free = slot + slot = next_slot return output From 43f01c6ca1ee9ade8d2b9fbf09eb6fc4459133f6 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Wed, 2 Sep 2026 16:40:11 +0100 Subject: [PATCH 04/19] Hold the pending seeds of a node in a numba typed list Replace the hand rolled array of linked lists with a typed list per node, holding indexes into the seed arrays. The causal site and effect size of a seed are fixed for the whole sweep, so an item is one integer and is looked up rather than copied. A node's list is made when something first reaches it, since most nodes are never reached when the causal alleles are rare, and dropping the reference once the node has been swept hands the storage straight back. This removes the arena, the free list and the doubling block that was written out twice. All figures below are level="node" on a tree sequence of 100,000 samples, 217,091 nodes, 285,979 edges, 22,523 trees and 235,458 sites, taking causal sites either uniformly over all sites or from those with an allele frequency below 0.001. Against the array of linked lists it replaces, per causal site count: 1,000 10,000 100,000 uniform typed 0.656s 5.615s 54.8s manual 1.121s 10.427s 109.7s 1.71x 1.86x 2.00x rare typed 0.076s 0.163s 0.558s manual 0.060s 0.099s 0.665s 0.79x 0.61x 1.19x It is faster wherever the structure is under any pressure, and slower only where the lists are a few items long and the sweep is short enough that making them is a noticeable part of it, a difference of tens of milliseconds. Peak memory at 100,000 uniformly drawn causal sites falls from 4.57GB to 1.22GB, because each node's storage is released as the sweep passes it rather than being held to the high water mark with up to twice as much again in slack. Against the tree by tree implementation this branch started from, which took 0.436s, 3.761s and 38.374s for 1000, 10,000 and 100,000 causal sites, and whose cost comes from passes over every node rather than from the number of carriers: 1,000 10,000 100,000 rare 5.7x 23.1x 68.8x uniform 0.66x 0.67x 0.70x The remaining loss on uniformly drawn causal sites is the common variants in the tail of the frequency spectrum, which touch 7.4% of the nodes on average where the median site touches 0.39%. Following carriers is random access at around 50ns each, while a dense pass over every node is sequential at under a nanosecond, so the two cross over at a carrier fraction of one or two percent. --- CHANGELOG.md | 9 ++++ tstrait/jit.py | 121 ++++++++++++++++++------------------------------- 2 files changed, 52 insertions(+), 78 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index b032bda..0e119b3 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -11,6 +11,15 @@ In development {pr}`189` - Added `edge_effect` to compute introduced effects on edges {pr}`189` +### Performance + +- `genetic_value` pushes the effect of each causal mutation down the ARG + instead of working through the trees one at a time, so its cost is the + number of nodes carrying a causal allele rather than the size of the tree + sequence. On 100,000 samples, a trait with 100,000 rare causal sites is + around 70 times faster; traits whose causal sites are mostly common + variants are around 1.5 times slower. + ### Breaking changes - `genetic_value` and `edge_effect` now raise a `ValueError` if a `site_id` in diff --git a/tstrait/jit.py b/tstrait/jit.py index ffdf0c6..9aaf7a9 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -20,9 +20,13 @@ import numba import numpy as np import tskit +from numba.core import types +from numba.typed import List -# The arena starts at this many tuples and doubles whenever it fills up. -INITIAL_ARENA_SIZE = 1024 +# What a node holds while it waits to be swept: the indexes of the seeds whose +# effect has reached it. The causal site and effect size of a seed are looked +# up in arrays that are fixed for the whole sweep, so an item is one integer. +_SEED_LIST = types.ListType(types.int32) @numba.njit @@ -42,17 +46,17 @@ def _push_down_arg( """ Accumulate the genetic value of every causal site in one sweep of the nodes. - Each node conceptually holds a set of ``(causal site, effect)`` tuples, - seeded at the node of each causal mutation. Sweeping the nodes from the - past to the present visits every node after all of its ancestors, because - a parent is always older than its child, so a tuple can be pushed from a - node to its children and is guaranteed to have arrived before that child is - reached. For each tuple the outbound edges spanning its causal site are - found, and the tuple is credited to the edge and to the node at the far end - before being passed on to it. + Each node holds the seeds whose effect has reached it, seeded at the node + of each causal mutation. Sweeping the nodes from the past to the present + visits every node after all of its ancestors, because a parent is always + older than its child, so a seed can be passed from a node to its children + and is guaranteed to have arrived before that child is reached. For each + seed the outbound edges spanning its causal site are found, and the seed is + credited to the edge and to the node at the far end before being passed on + to it. - A node is credited when a tuple arrives rather than when it is swept, so a - node that is never a parent never holds a tuple. On a large tree sequence + A node is credited when a seed arrives rather than when it is swept, so a + node that is never a parent never holds anything. On a large tree sequence that is about half of the work. ``child_index`` is the tskit child index, in which a node that is never a @@ -63,91 +67,52 @@ def _push_down_arg( discards the contribution. """ num_nodes = len(child_index) - head = np.full(num_nodes, tskit.NULL, dtype=np.int32) - arena_size = INITIAL_ARENA_SIZE - tuple_site = np.empty(arena_size, dtype=np.int32) - tuple_weight = np.empty(arena_size, dtype=np.float64) - tuple_next = np.empty(arena_size, dtype=np.int32) - # Slots are taken from the free list first, and from the top of the arena - # only when it is empty. Every tuple of a node is dead once that node has - # been swept, so the arena only ever holds the tuples still in flight. - arena_top = 0 - free = tskit.NULL + # A node's list is made when something first reaches it. Most nodes are + # never reached when the causal alleles are rare, and making a list for + # every one of them up front then costs more than the sweep does. + empty = List.empty_list(types.int32) + pending = List.empty_list(_SEED_LIST) + for _ in range(num_nodes): + pending.append(empty) + reached = np.zeros(num_nodes, dtype=np.bool_) for j in range(len(seed_node)): u = seed_node[j] - weight = seed_weight[j] if seed_output[j] >= 0: - output[seed_output[j]] += weight + output[seed_output[j]] += seed_weight[j] if child_index[u, 0] < 0: continue - if free != tskit.NULL: - slot = free - free = tuple_next[slot] - else: - if arena_top == arena_size: - arena_size *= 2 - grown_site = np.empty(arena_size, dtype=np.int32) - grown_weight = np.empty(arena_size, dtype=np.float64) - grown_next = np.empty(arena_size, dtype=np.int32) - grown_site[:arena_top] = tuple_site - grown_weight[:arena_top] = tuple_weight - grown_next[:arena_top] = tuple_next - tuple_site = grown_site - tuple_weight = grown_weight - tuple_next = grown_next - slot = arena_top - arena_top += 1 - tuple_site[slot] = seed_site[j] - tuple_weight[slot] = weight - tuple_next[slot] = head[u] - head[u] = slot + if not reached[u]: + pending[u] = List.empty_list(types.int32) + reached[u] = True + pending[u].append(np.int32(j)) for i in range(len(nodes_by_time)): parent = nodes_by_time[i] - slot = head[parent] - if slot == tskit.NULL: + if not reached[parent]: continue - head[parent] = tskit.NULL + items = pending[parent] edge_start = child_index[parent, 0] edge_stop = child_index[parent, 1] - while slot != tskit.NULL: - site = tuple_site[slot] - weight = tuple_weight[slot] + for k in range(len(items)): + item = items[k] + site = seed_site[item] + weight = seed_weight[item] for e in range(edge_start, edge_stop): if edges_site_start[e] <= site and site < edges_site_stop[e]: if edges_output[e] >= 0: output[edges_output[e]] += weight child = edges_child[e] if child_index[child, 0] < 0: - # Nothing below, so there is no tuple to hold. + # Nothing below, so there is nothing to hold. continue - if free != tskit.NULL: - child_slot = free - free = tuple_next[child_slot] - else: - if arena_top == arena_size: - arena_size *= 2 - grown_site = np.empty(arena_size, dtype=np.int32) - grown_weight = np.empty(arena_size, dtype=np.float64) - grown_next = np.empty(arena_size, dtype=np.int32) - grown_site[:arena_top] = tuple_site - grown_weight[:arena_top] = tuple_weight - grown_next[:arena_top] = tuple_next - tuple_site = grown_site - tuple_weight = grown_weight - tuple_next = grown_next - child_slot = arena_top - arena_top += 1 - tuple_site[child_slot] = site - tuple_weight[child_slot] = weight - tuple_next[child_slot] = head[child] - head[child] = child_slot - # This tuple is finished with, so its slot can be reused. - next_slot = tuple_next[slot] - tuple_next[slot] = free - free = slot - slot = next_slot + if not reached[child]: + pending[child] = List.empty_list(types.int32) + reached[child] = True + pending[child].append(np.int32(item)) + # A node is swept once, so dropping the reference here hands its + # storage back for the nodes still to come. + pending[parent] = empty return output From f004ae331fd055265988909465438dbe094134ef Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Thu, 3 Sep 2026 10:44:40 +0100 Subject: [PATCH 05/19] Benchmark genetic_value on a simulation small enough to iterate on The tree sequence the ARG sweep was written against takes nine minutes for a full grid and 48 seconds for its longest single call, which is too slow to work against, and the benchmark measured wall time and nothing else. The peak memory in the last commit message and the carrier fractions in the one before it were both produced out of tree. Add a --preset flag and default it to a simulation of 30,000 samples over 1Mb, which has 63,287 nodes against 217,091 and runs the grid in thirty seconds. What makes it a fair substitute is not the timings but the distribution underneath them, since the sweep costs the number of nodes that carry a causal allele: the mean fraction of the nodes that a causal site reaches is 7.19% against 5.99%, and the median 0.406% against 0.184%, which is as close as 150 sampled sites from a distribution with that tail can show. Both presets then have a flat per call floor, a uniform cost per causal site that falls to a plateau from 1000 sites upwards, a rare cost still falling at the top of the grid, a uniform to rare ratio growing to two orders of magnitude, and parity between the three levels. The plateau is 160us against 483us per causal site, a factor of 3.0 on a node count ratio of 3.4. Sixteen times smaller was tried and is measurably less faithful, and the grid stops at 10,000 causal sites because raising the mutation rate to lift the ceiling would take sites per node from 0.68 to 2.7 against the large preset's 1.08 and inflate the setup floor. Measure four things besides the wall clock. --phases times the check, the setup, the kernel and the dataframe separately, without which a flat setup cost and a kernel that grows with the causal sites cannot be told apart; setup goes from 8ms to 17ms across the whole small grid and the dataframe is 1ms, so at 10,000 causal sites the kernel is 98% of the call. --counters runs a counting only copy of the sweep, since perf cannot see inside the kernel. --structure reports the carrier fractions above. --memory reports peak RSS, resetting VmHWM through clear_refs because it never falls on its own. The counters say three things that were assumed otherwise. The scan of a node's out edges is not where the waste is, because 94% of the trips find an edge spanning the seed's causal site. Uniform selection costs 178 times as many scans as rare at 10,000 causal sites but only 46 times the time, because rare costs 63ns a scan against 28ns, which is the random access against sequential difference measured directly rather than inferred. And the kernel has a floor in the number of nodes: it builds a pending entry for every node and sweeps every node whether or not anything reached it, so one rare causal site, ten scans, still takes 2ms. Add profile_genetic_value.py, which profiles one cell of the grid with cProfile for the setup and sets up a cell to run just the kernel under perf, printing the commands. Two things had to be established for that to be worth anything. Source lines inside the kernel are not available, because this llvmlite has no LLVM PerfJITEventListener and so nothing writes a perf map or a jitdump; what perf does give is the split between the kernel, the numba runtime with its typed list functions named, the interpreter and compilation. And NUMBA_ENABLE_PROFILING=1, the documented way to profile numba, is the wrong thing to use here: it would only help through the listener that is missing, and it defaults NUMBA_DEBUGINFO to 1, which measured 2.62s a call against 1.61s. perf finds the JIT mappings by itself. Record that _GeneticValue takes the _root_runs branch, a Python loop over every tree costing 7ms on the small preset and 38ms on the large one. It is not the edge case it looks like: drawing 10,000 sites uniformly makes it near certain that one of them has the ancestral state as its causal allele, and it fires at 10,000 and 100,000 causal sites on both presets, so on the large one it is over a third of the setup. Check in the small preset grid to diff against. The counts in it are identical between runs and are the part worth treating as a regression test; the timings are specific to the machine they were taken on. --- benchmarks/README.md | 208 ++++++++++- benchmarks/baseline_small.csv | 85 +++++ benchmarks/baseline_small_counters.csv | 9 + benchmarks/benchmark_genetic_value.py | 498 +++++++++++++++++++++++-- benchmarks/profile_genetic_value.py | 226 +++++++++++ 5 files changed, 978 insertions(+), 48 deletions(-) create mode 100644 benchmarks/baseline_small.csv create mode 100644 benchmarks/baseline_small_counters.csv create mode 100644 benchmarks/profile_genetic_value.py diff --git a/benchmarks/README.md b/benchmarks/README.md index 16796f4..7a3b6e0 100644 --- a/benchmarks/README.md +++ b/benchmarks/README.md @@ -6,17 +6,23 @@ that performance work can be repeated and compared. ## `benchmark_genetic_value.py` Measures how `tstrait.genetic_value` scales with the number of causal sites in -a trait, which is the regime where the current implementation is slow: a trait -with a large number of weakly causal sites. +a trait. ``` uv run --group test benchmarks/benchmark_genetic_value.py ``` -By default this simulates 100,000 sample nodes (50,000 diploid individuals) -over 5Mb, and times `genetic_value` for 1, 100, 1000, 10,000 and 100,000 causal -sites at each of the `individual`, `node` and `edge` levels, taking the minimum -of three replicates. +The cost of the ARG sweep is the number of nodes that carry a causal allele, so +the allele frequency of the causal sites matters more than how many there are. +`--selections` therefore draws them two ways: `uniform` over all sites, which +the common variants in the tail of the frequency spectrum dominate, and `rare`, +restricted to sites below `--rare-threshold`. The two differ by two orders of +magnitude and behave differently, so a single number for "the cost of a causal +site" is meaningless without saying which. + +`sim_trait` is timed separately, because it has a per-site Python loop of its +own that we do not want folded into the `genetic_value` numbers. The numba +kernels are compiled by a warm up call that is not timed. The mutation rate defaults to 1e-7, ten times the human rate, so that there are enough sites in a genome short enough to simulate quickly. The mutation rate @@ -24,17 +30,187 @@ does not affect the allele frequency spectrum, so the causal sites are as weakly causal as they would be under a realistic rate; only the number of sites per tree is inflated. -`sim_trait` is timed separately, because it has a per-site Python loop of its -own that we do not want folded into the `genetic_value` numbers. The numba -kernels are compiled by a warm up call that is not timed. +### Presets + +`--preset small` is the default and takes about 30 seconds; `--preset large` is +the tree sequence the ARG sweep was written against and takes about nine +minutes, with its longest single call at 48 seconds. + +| | small | large | +|---|---|---| +| samples | 30,000 | 100,000 | +| sequence length | 1Mb | 5Mb | +| nodes | 63,287 | 217,091 | +| edges | 75,586 | 285,979 | +| trees | 4,108 | 22,523 | +| sites | 42,794 | 235,458 | +| causal sites | 1 … 10,000 | 1 … 100,000 | +| cached tree sequence | 6.5MB | 27MB | + +A preset only fills in `--samples`, `--length` and `--num-causal`, each of which +still overrides it when given explicitly. The tree sequence is cached in +`_output/` under a name built from the simulation parameters, so `small` +simulates itself in 0.2s the first time and is loaded after that. + +`small` is the default because iterating against a nine minute grid is not +practical. It reproduces the patterns the large one shows, at `level="node"`: + +| num_causal | small uniform | small rare | ratio | large uniform | large rare | ratio | +|---|---|---|---|---|---|---| +| 1 | 0.016s | 0.011s | 1.4 | 0.058s | 0.051s | 1.1 | +| 100 | 0.036s | 0.013s | 2.8 | 0.152s | 0.056s | 2.7 | +| 1,000 | 0.189s | 0.015s | 12.3 | 0.593s | 0.065s | 9.1 | +| 10,000 | 1.596s | 0.035s | 45.7 | 5.171s | 0.108s | 48 | +| 100,000 | — | — | | 48.310s | 0.476s | 101 | + +Both show a flat per call floor, a uniform µs/site that falls to a plateau from +1,000 causal sites upwards, a rare µs/site that is still falling at the top of +the grid, a uniform to rare ratio that grows with the number of causal sites, +and parity between the three levels. The uniform plateau is 160µs/site against +483µs/site, a factor of 3.0 on a node count ratio of 3.4. + +What makes the small preset a fair substitute is not the timings but the +distribution underneath them: the fraction of the nodes that a causal site's +effect reaches, which is what the sweep costs. `--structure` measures it. + +| | mean carrier fraction | median | +|---|---|---| +| small | 7.19% | 0.406% | +| large | 5.99% | 0.184% | + +Those are 150 sampled sites from a distribution with a heavy tail, so they +agree about as well as they can. Check this again before trusting a new preset. + +The grid stops at 10,000 causal sites on `small` because the rare pool is only +about 37% of its 42,794 sites. Raising the mutation rate to lift the ceiling +was tried and rejected: it takes `sites/nodes` from 0.68 to about 2.7 against +the large preset's 1.08, and the per-call setup floor grows with the number of +sites, so the low end of the curve stops being comparable. + +### What else it measures + +Wall time on its own does not say why a configuration is slow. Four optional +modes say more. `--phases` is the expensive one, roughly doubling the run +because it times the same work again a piece at a time; the other three add +seconds. + +`--phases` times `_check_trait_df`, `_GeneticValue.__init__`, the +`_push_down_arg` kernel and the output dataframe separately. The end to end +number stays as the headline. This is what tells an algorithmic win from a +setup win: setup barely grows with the causal sites, going from 8ms to 17ms +across the whole `small` grid and sitting near 100ms on `large`, so at one +causal site the public call is measuring almost nothing else, while at 10,000 +the kernel is 98% of it. Most of setup is `tskit.jit.numba.jitwrap`, which runs +three Python-speed `max(map(len, ...))` passes over the site and mutation +tables; `_root_runs` is most of the rest when it fires. The dataframe is about +1ms and is not worth thinking about. + +`--counters` runs a counting-only copy of the sweep and reports the work it +does. perf cannot attribute time to source lines inside a numba kernel here +(see below), so counting what the kernel does and dividing is the way to say +where the time goes. On `small`: + +``` +selection num_causal seeds edge_scans edge_hits appends reached scans/seed hit rate reached +uniform 1 1 35,915 33,818 16,909 16,909 35915.0 94.2% 26.7% +uniform 10000 10,107 56,556,279 53,117,294 26,558,647 32,966 5595.8 93.9% 52.1% +rare 1 1 10 10 5 5 10.0 100.0% 0.0% +rare 10000 10,028 318,619 302,898 151,449 29,775 31.8 95.1% 47.0% +``` + +`edge_scans` is the trip count of the kernel's innermost loop and is what the +run time is proportional to, so `ns/scan` in the summary table is the constant +an optimisation has to move. Three things fall out of the table above: -The simulation parameters, the grid of causal site counts, the levels and the -number of replicates are all command line arguments; see `--help`. +- The scan of a node's out edges is not where the waste is: 94% of trips find + an edge that spans the seed's causal site. +- Uniform selection costs 178 times as many scans as rare at 10,000 causal + sites but only 46 times the time, because rare is more expensive per scan: + 63ns against 28ns for the kernel alone. Following a few carriers is random + access; a dense pass over half the nodes is sequential. Compare the `kernel` + rows rather than `genetic_value`, or the setup floor swamps the rare ones. +- The kernel has an O(num_nodes) floor. It builds a pending entry for every + node before it starts and visits every node whether or not anything reached + it, so at one rare causal site — 10 scans — the kernel still takes 2ms. + +`--structure` reports the shape of the tree sequence and the carrier fraction +distribution described above. + +`--memory` reports the peak resident set size of each call, which is how the +typed list rewrite was justified. VmHWM never falls, so it is reset before each +call by writing to `/proc/self/clear_refs`; on a kernel without that the column +reads `unavailable`. This is the one thing the small preset is a poor substitute +for: 10,000 uniform causal sites peak at 0.02GB over the baseline there, against +the gigabytes the large preset reaches at 100,000. ### Output -The simulated tree sequence is cached in `_output/`, keyed by the simulation -parameters, so that repeated runs do not resimulate it. Results are written to -`_output/genetic_value.csv` in long format, one row per replicate, together -with the dimensions of the tree sequence they were measured on. `_output/` is -gitignored. +Results are written to `_output/genetic_value.csv` in long format, one row per +replicate, together with the dimensions of the tree sequence they were measured +on; `--counters` writes a second file alongside it. `_output/` is gitignored. + +`baseline_small.csv` and `baseline_small_counters.csv` are the `small` preset at +the tip of the ARG sweep work, for diffing against. The timings in the first are +specific to the machine they were taken on; the counts in the second are not, +and are the part worth treating as a regression test. + +## `profile_genetic_value.py` + +Profiles a single cell of the grid, in two modes because the Python setup and +the numba kernel need different tools. + +``` +uv run --group test benchmarks/profile_genetic_value.py --mode python +uv run --group test benchmarks/profile_genetic_value.py --mode kernel +``` + +`--mode python` is cProfile around the public call. It is the only way to see +the setup, and it shows `jitwrap` and its `builtins.max` rows plainly. The +kernel appears in it as one opaque dispatcher frame. + +`--mode kernel` sets one cell up and runs only the kernel, so that a sampling +profile is not swamped by simulating effect sizes for the whole site pool. Run +on its own it prints the perf commands to copy; `--run` is what those commands +invoke. `perf_event_paranoid` is usually high enough to need `sudo`. + +Two things to know before reading a perf profile of this code. + +**Source lines inside the kernel are not available.** This llvmlite has no LLVM +`PerfJITEventListener`, so nothing writes a `/tmp/perf-.map` or a jitdump +and the kernel shows up as raw addresses under `[JIT]`. Use `--counters` for +attribution inside the kernel. What perf does give is the split between the +kernel, the numba runtime, the interpreter and LLVM compilation, with the typed +list functions resolved by name — which is how to size the typed list overhead: + +``` +64.42% [JIT] tid 4722 <- the kernel +12.18% python3.11 <- setup + 6.27% libc.so.6 + 6.02% libllvmlite.so <- compilation, not work + 4.93% _helperlib.cpython-311-...so <- numba_list_append, numba_list_resize +``` + +That is `--repeats 10`; the setup is a fixed few seconds, so raise `--repeats` +until the `[JIT]` share stops moving. + +**Do not set `NUMBA_ENABLE_PROFILING=1`.** It is the documented way to profile +numba and is wrong here: it would only help through the listener llvmlite does +not have, and it defaults `NUMBA_DEBUGINFO` to 1, which changes the generated +code. The same kernel measured 2.62s a call with it and 1.61s without. perf +finds the JIT mappings by itself, so the recipe does not set it. + +## Gotchas + +- `_GeneticValue` takes a `_root_runs` branch whenever a drawn causal allele is + the ancestral state of its site. It is a Python loop over every tree, costing + 7ms on `small` and 38ms on `large`, so on `large` it is over a third of the + setup. It is not an edge case: drawing 10,000 sites uniformly makes it near + certain that one of them qualifies, and it fires at 10,000 and 100,000 causal + sites on both presets. At the low end of the grid it usually does not, so it + appears part way up the curve and looks like a step in the setup cost. The + benchmark prints a line when a cell takes it. +- Never compare replicate 0 against replicates 1 and up. `mutations_inherited_state` + and its neighbours are built lazily and cached on the tree sequence, so the + first call on a given tree sequence pays for all of them. The warm up call + covers numba compilation but not this. +- The summary takes the minimum over replicates, not the mean. diff --git a/benchmarks/baseline_small.csv b/benchmarks/baseline_small.csv new file mode 100644 index 0000000..483fac3 --- /dev/null +++ b/benchmarks/baseline_small.csv @@ -0,0 +1,85 @@ +phase,num_causal,selection,level,replicate,seconds,num_samples,num_individuals,num_nodes,num_edges,num_trees,num_sites +sim_trait,1,uniform,,0,0.0024821249999149586,30000,15000,63287,75586,4108,42794 +sim_trait,1,uniform,,1,0.0022028849998605438,30000,15000,63287,75586,4108,42794 +sim_trait,1,uniform,,2,0.0021113359998707892,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,individual,0,0.017362073000185774,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,individual,1,0.016651905999424343,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,individual,2,0.016448290999505844,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,node,0,0.01680002900047839,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,node,1,0.016700003000551078,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,node,2,0.016296245999910752,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,edge,0,0.01801782100028504,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,edge,1,0.01743433499996172,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,edge,2,0.01691987899994274,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,individual,0,0.012258568000106607,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,individual,1,0.01165933300035249,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,individual,2,0.011369154000021808,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,node,0,0.011771557999963989,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,node,1,0.011883568000484956,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,node,2,0.011892135999914899,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,edge,0,0.012155798999629042,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,edge,1,0.012474035000195727,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,edge,2,0.01207672599957732,30000,15000,63287,75586,4108,42794 +sim_trait,100,uniform,,0,0.007151913000598142,30000,15000,63287,75586,4108,42794 +sim_trait,100,uniform,,1,0.005853265000041574,30000,15000,63287,75586,4108,42794 +sim_trait,100,uniform,,2,0.005971398000838235,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,individual,0,0.03708483099944715,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,individual,1,0.03696376599964424,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,individual,2,0.03942140500021196,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,node,0,0.0380683780003892,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,node,1,0.039289737999752106,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,node,2,0.03719520599952375,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,edge,0,0.0392425390000426,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,edge,1,0.03797104800014495,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,edge,2,0.04035435600053461,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,individual,0,0.01287275499998941,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,individual,1,0.012763894000272558,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,individual,2,0.01330246699944837,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,node,0,0.01320572200074821,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,node,1,0.012976414999684494,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,node,2,0.013185578000047826,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,edge,0,0.013328967999768793,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,edge,1,0.013860037999620545,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,edge,2,0.0130741749999288,30000,15000,63287,75586,4108,42794 +sim_trait,1000,uniform,,0,0.0281717200005005,30000,15000,63287,75586,4108,42794 +sim_trait,1000,uniform,,1,0.028771799999958603,30000,15000,63287,75586,4108,42794 +sim_trait,1000,uniform,,2,0.03157023900075728,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,individual,0,0.18984438400002546,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,individual,1,0.21684775500034448,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,individual,2,0.2687714949997826,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,node,0,0.226016618000358,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,node,1,0.20893606600020576,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,node,2,0.1926977330003865,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,edge,0,0.19492765799986955,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,edge,1,0.2099747740003295,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,edge,2,0.20740931999989698,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,individual,0,0.017229248999683477,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,individual,1,0.018185036999966542,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,individual,2,0.018869548999646213,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,node,0,0.018393275000562426,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,node,1,0.016106918999867048,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,node,2,0.015783592999468965,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,edge,0,0.015947148999657657,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,edge,1,0.016200306999962777,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,edge,2,0.015847360999941884,30000,15000,63287,75586,4108,42794 +sim_trait,10000,uniform,,0,0.24855346700041991,30000,15000,63287,75586,4108,42794 +sim_trait,10000,uniform,,1,0.25112143699971057,30000,15000,63287,75586,4108,42794 +sim_trait,10000,uniform,,2,0.2850705849996302,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,individual,0,1.704894198999682,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,individual,1,1.676236326000435,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,individual,2,1.6262085170001228,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,node,0,1.7224619290000192,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,node,1,1.6860980929996003,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,node,2,1.7367036840005312,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,edge,0,1.675767371000802,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,edge,1,1.613268662000337,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,edge,2,1.7291035789994567,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,individual,0,0.04531490799945459,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,individual,1,0.037717979999797535,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,individual,2,0.03916538799967384,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,node,0,0.044803013000091596,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,node,1,0.05839010900035646,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,node,2,0.05877978399985295,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,edge,0,0.03882118100045773,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,edge,1,0.03812910400029068,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,edge,2,0.03929764599979535,30000,15000,63287,75586,4108,42794 diff --git a/benchmarks/baseline_small_counters.csv b/benchmarks/baseline_small_counters.csv new file mode 100644 index 0000000..d730e93 --- /dev/null +++ b/benchmarks/baseline_small_counters.csv @@ -0,0 +1,9 @@ +num_causal,selection,seeds,edge_scans,edge_hits,appends,reached,num_nodes,num_edges +1,uniform,1,35915,33818,16909,16909,63287,75586 +1,rare,1,10,10,5,5,63287,75586 +100,uniform,101,558839,524756,262378,31784,63287,75586 +100,rare,100,1507,1454,727,706,63287,75586 +1000,uniform,1008,5856520,5501688,2750844,32611,63287,75586 +1000,rare,1002,11559,11228,5614,4671,63287,75586 +10000,uniform,10107,56556279,53117294,26558647,32966,63287,75586 +10000,rare,10028,318619,302898,151449,29775,63287,75586 diff --git a/benchmarks/benchmark_genetic_value.py b/benchmarks/benchmark_genetic_value.py index 237645a..35ba92e 100644 --- a/benchmarks/benchmark_genetic_value.py +++ b/benchmarks/benchmark_genetic_value.py @@ -2,10 +2,11 @@ Benchmark tstrait.genetic_value as the number of causal sites grows. The interesting regime is a trait with a large number of weakly causal sites, -i.e. many causal sites that are each carried by a small number of samples. -Causal sites are chosen uniformly at random over all sites, and because the -site frequency spectrum is dominated by rare variants this puts us in that -regime by default. +i.e. many causal sites that are each carried by a small number of samples. The +cost of the ARG sweep is the number of nodes that carry a causal allele, so the +allele frequency of the causal sites matters more than how many there are, and +``--selections`` draws them either uniformly over all sites or from the rare +ones alone. Run with ``uv run --group test benchmarks/benchmark_genetic_value.py``. """ @@ -14,19 +15,136 @@ import csv import functools import pathlib +import statistics import sys import time import msprime +import numba import numpy as np +import pandas as pd import tskit +from numba.core import types +from numba.typed import List import tstrait +from tstrait import jit +from tstrait.genetic_value import _check_trait_df, _GeneticValue -DEFAULT_NUM_CAUSAL = [1, 100, 1000, 10_000, 100_000] LEVELS = ["individual", "node", "edge"] SELECTIONS = ["uniform", "rare"] +# The small preset is the default because the large one takes about six minutes +# and its longest single call takes almost a minute, which is too slow to +# iterate against. It reproduces the same patterns; see the README for the +# evidence, and --structure to check it again after a change. +PRESETS = { + "small": { + "samples": 30_000, + "length": 1e6, + "num_causal": [1, 100, 1000, 10_000], + }, + "large": { + "samples": 100_000, + "length": 5e6, + "num_causal": [1, 100, 1000, 10_000, 100_000], + }, +} + +# The counting kernel mirrors tstrait.jit._push_down_arg, so it holds what a +# node holds there: the indexes of the seeds whose effect has reached it. +_SEED_LIST = types.ListType(types.int32) + +COUNTERS = ["seeds", "edge_scans", "edge_hits", "appends", "reached"] + + +@numba.njit +def _count_push_down_arg( + child_index, + edges_child, + edges_site_start, + edges_site_stop, + nodes_by_time, + seed_node, + seed_site, +): + """ + Count the work that ``tstrait.jit._push_down_arg`` does, without doing any + of it. + + perf cannot attribute time to source lines inside a numba kernel here, so + the way to say where the time goes is to count the things the kernel does + and divide. This is a copy of the sweep with the output writes and the + weight lookups taken out and a counter put in their place, which is the + only reason it may diverge from the kernel it mirrors: keep them in step. + + Returns the counters named in ``COUNTERS``: + + ``seeds`` the causal mutations, plus the roots seeded when the causal + allele is the ancestral state + ``edge_scans`` trips of the innermost loop, i.e. the sum over every + (swept node, seed held there) pair of the node's out degree + ``edge_hits`` those trips where the edge spans the seed's causal site, so + edge_scans - edge_hits is the scan that was wasted + ``appends`` seeds pushed onto a node's list + ``reached`` nodes that held a list, so reached / num_nodes is the + fraction of the tree sequence the sweep touched, and, since + a list is made for a node the first time anything reaches + it, also the number of typed lists allocated + """ + num_nodes = len(child_index) + empty = List.empty_list(types.int32) + pending = List.empty_list(_SEED_LIST) + for _ in range(num_nodes): + pending.append(empty) + reached = np.zeros(num_nodes, dtype=np.bool_) + + edge_scans = 0 + edge_hits = 0 + appends = 0 + + for j in range(len(seed_node)): + u = seed_node[j] + if child_index[u, 0] < 0: + continue + if not reached[u]: + pending[u] = List.empty_list(types.int32) + reached[u] = True + pending[u].append(np.int32(j)) + appends += 1 + + for i in range(len(nodes_by_time)): + parent = nodes_by_time[i] + if not reached[parent]: + continue + items = pending[parent] + edge_start = child_index[parent, 0] + edge_stop = child_index[parent, 1] + for k in range(len(items)): + item = items[k] + site = seed_site[item] + edge_scans += edge_stop - edge_start + for e in range(edge_start, edge_stop): + if edges_site_start[e] <= site and site < edges_site_stop[e]: + edge_hits += 1 + child = edges_child[e] + if child_index[child, 0] < 0: + continue + if not reached[child]: + pending[child] = List.empty_list(types.int32) + reached[child] = True + pending[child].append(np.int32(item)) + appends += 1 + pending[parent] = empty + + return ( + len(seed_node), + edge_scans, + edge_hits, + appends, + int(np.sum(reached)), + ) + def cached_simulation(args): """ @@ -76,6 +194,11 @@ def time_call(func, replicates): """ Call func the given number of times, returning the result of the last call and the elapsed time of each call. + + The first call on a tree sequence is systematically more expensive than the + ones after it, because tskit builds mutations_inherited_state and its + neighbours lazily and caches them on the tree sequence, so a replicate is + only comparable with another replicate of the same call. """ times = [] for _ in range(replicates): @@ -85,6 +208,31 @@ def time_call(func, replicates): return result, times +def _read_status(key): + for line in open("/proc/self/status"): + if line.startswith(key): + return int(line.split()[1]) * 1024 + return None + + +def peak_memory(func): + """ + Return the result of func and the peak resident set size while it ran. + + VmHWM is a high water mark that never falls, so it has to be reset before + the call rather than subtracted after it; writing 5 to clear_refs does + that. Returns None for the memory where the kernel does not support it. + """ + try: + with open("/proc/self/clear_refs", "w") as f: + f.write("5") + except OSError: + return func(), None + before = _read_status("VmRSS") + result = func() + return result, (_read_status("VmHWM"), before) + + def causal_site_pool(ts, model, seed): """ Simulate effect sizes for every site, so that causal sites can be drawn @@ -103,7 +251,7 @@ def select_causal(pool, selection, num_causal, rare_threshold, rng): Draw num_causal rows from the pool, either uniformly over the sites or restricted to the rare ones. Uniform selection is dominated by the common variants in the tail of the frequency spectrum, which behave quite - differently from the weakly causal sites the ARG descent is aimed at. + differently from the weakly causal sites the ARG sweep is aimed at. """ if selection == "rare": pool = pool[pool["allele_freq"] < rare_threshold] @@ -113,7 +261,7 @@ def select_causal(pool, selection, num_causal, rare_threshold, rng): return pool.iloc[np.sort(keep)] -def warm_up(ts, model, levels): +def warm_up(ts, model, levels, counters): """ Run the smallest possible simulation through each code path so that the numba kernels are compiled before we start timing. @@ -122,19 +270,156 @@ def warm_up(ts, model, levels): trait_df = tstrait.sim_trait(ts, model=model, num_causal=1, random_seed=1) for level in levels: tstrait.genetic_value(ts, trait_df, level=level) + if counters: + # A kernel of its own to compile, so only pay for it when it is used. + count_work(ts, _check_trait_df(ts, trait_df)) print(f"Warm up (includes numba compilation): {time.perf_counter() - before:.1f}s") +def time_phases(ts, trait_df, level, replicates): + """ + Time the phases of a genetic_value call separately, returning a list of + (phase, times) pairs. + + The public call is a single number in which a flat per call setup cost and + a kernel that grows with the causal sites are indistinguishable, and at one + causal site it is almost all setup. Splitting them is what tells an + algorithmic win from a setup win. + """ + phases = [] + _, times = time_call(functools.partial(_check_trait_df, ts, trait_df), replicates) + phases.append(("check", times)) + checked = _check_trait_df(ts, trait_df) + + _, times = time_call(functools.partial(_GeneticValue, ts, checked), replicates) + phases.append(("setup", times)) + genetic = _GeneticValue(ts, checked) + + arguments = genetic._descent_arguments(level, 0) + _, times = time_call( + # A fresh output array each time, since the kernel accumulates into it. + lambda: jit._push_down_arg( + **{**arguments, "output": np.zeros(len(arguments["output"]))} + ), + replicates, + ) + phases.append(("kernel", times)) + + size = len(arguments["output"]) + values = np.zeros((genetic.num_trait, size)) + _, times = time_call( + functools.partial(_build_frame, genetic.num_trait, size, level, values), + replicates, + ) + phases.append(("frame", times)) + return phases + + +def _build_frame(num_trait, size, level, values): + """ + The dataframe _run builds from the kernel output, at genetic_value.py:249. + """ + return pd.DataFrame( + { + "trait_id": np.repeat(np.arange(num_trait), size), + f"{level}_id": np.tile(np.arange(size), num_trait), + "genetic_value": values.flatten(), + } + ) + + +def count_work(ts, trait_df): + """ + Return the work counters for a trait, as a dict keyed by COUNTERS. + + The counts do not depend on the level, because the same sweep serves all + three and only the array the contributions land in differs. + """ + genetic = _GeneticValue(ts, trait_df) + arguments = genetic._descent_arguments("node", 0) + counts = _count_push_down_arg( + arguments["child_index"], + arguments["edges_child"], + arguments["edges_site_start"], + arguments["edges_site_stop"], + arguments["nodes_by_time"], + arguments["seed_node"], + arguments["seed_site"], + ) + return dict(zip(COUNTERS, counts)) + + +def root_runs_fired(ts, trait_df): + """ + Whether this trait takes the _root_runs branch in _GeneticValue. + + That branch is a Python loop over every tree, so it costs O(num_trees) at + Python speed, and it fires only when a drawn causal allele happens to be + the ancestral state of its site. It therefore appears and vanishes with the + seed, which makes for confusing non-monotonic timings unless it is reported. + """ + site_id = trait_df["site_id"].to_numpy() + causal_allele = trait_df["causal_allele"].to_numpy() + return bool(np.any(ts.sites_ancestral_state[site_id] == causal_allele)) + + +def carrier_fractions(ts, pool, sample_size, rng): + """ + Return the fraction of the nodes that each of a sample of causal sites + reaches. + + The cost of the sweep is the number of nodes a causal allele is carried by, + so this distribution, not the size of the tree sequence, is what a smaller + simulation has to reproduce. Giving each sampled site its own trait_id gets + all of them from a single setup, and a node the effect never reached is + exactly a node left at zero. + """ + size = min(sample_size, len(pool)) + keep = np.sort(rng.choice(len(pool), size=size, replace=False)) + trait_df = pool.iloc[keep].copy() + trait_df["trait_id"] = np.arange(len(trait_df)) + values = tstrait.genetic_value(ts, trait_df, level="node")["genetic_value"] + reached = values.to_numpy().reshape(len(trait_df), ts.num_nodes) + return np.count_nonzero(reached, axis=1) / ts.num_nodes + + +def report_structure(ts, pool, args): + """ + Print the shape of the tree sequence and the carrier fraction distribution, + which is the check that a preset still reproduces the large simulation. + """ + rng = np.random.default_rng(args.seed) + fractions = carrier_fractions(ts, pool, args.structure_sites, rng) + print(f"\n{describe(ts)}") + print( + f"sites/nodes {ts.num_sites / ts.num_nodes:.2f} " + f"edges/nodes {ts.num_edges / ts.num_nodes:.2f}" + ) + print(f"Carrier fraction of nodes over {len(fractions)} sampled sites:") + print( + f" mean {fractions.mean() * 100:.2f}% " + f"median {statistics.median(fractions) * 100:.3f}% " + f"max {fractions.max() * 100:.2f}%" + ) + + def run_benchmark(ts, args): """ - Time each cell of the grid, returning the timings and the causal site - counts that we got through before running out of time budget. + Time each cell of the grid, returning the timings, the work counters, the + peak memory and the causal site counts that we got through before running + out of time budget. """ model = tstrait.trait_model(distribution="normal", mean=0, var=1) - warm_up(ts, model, args.levels) + warm_up(ts, model, args.levels, args.counters) pool = causal_site_pool(ts, model, args.seed) + if args.structure: + report_structure(ts, pool, args) + print() + rows = [] + counts = {} + memory = {} completed = [] for num_causal in args.num_causal: _, times = time_call( @@ -159,11 +444,22 @@ def run_benchmark(ts, args): if trait_df is None: print(f" too few {selection} sites for num_causal={num_causal}") continue + if root_runs_fired(ts, trait_df): + print( + f" {selection} num_causal={num_causal} takes the _root_runs " + "branch, a Python loop over every tree" + ) + if args.counters: + counts[(num_causal, selection)] = count_work( + ts, _check_trait_df(ts, trait_df) + ) for level in args.levels: - _, times = time_call( - functools.partial(tstrait.genetic_value, ts, trait_df, level=level), - args.replicates, + call = functools.partial( + tstrait.genetic_value, ts, trait_df, level=level ) + if args.memory: + _, memory[(num_causal, selection, level)] = peak_memory(call) + _, times = time_call(call, args.replicates) for replicate, seconds in enumerate(times): rows.append( ( @@ -177,9 +473,18 @@ def run_benchmark(ts, args): ) report(rows[-1]) slowest = max(slowest, max(times)) + if args.phases: + for phase, phase_times in time_phases( + ts, trait_df, level, args.replicates + ): + for replicate, seconds in enumerate(phase_times): + rows.append( + (phase, num_causal, selection, level, replicate, seconds) + ) + report(rows[-1]) completed.append(num_causal) - # The cost is superlinear in num_causal, so once a single call is over - # budget the next point on the grid is not worth waiting for. + # The cost grows with the number of causal sites, so once a single call + # is over budget the next point on the grid is not worth waiting for. if slowest > args.max_seconds: skipped = args.num_causal[len(completed) :] if len(skipped) > 0: @@ -188,7 +493,7 @@ def run_benchmark(ts, args): f"budget: skipping {skipped}" ) break - return rows, completed + return rows, counts, memory, completed def report(row): @@ -196,10 +501,11 @@ def report(row): print(f" {phase:<14} {num_causal:>7} {selection:<8} {level:<11} {seconds:8.3f}s") -def summarise(rows, completed, ts, args): +def summarise(rows, counts, memory, completed, ts, args): """ Print the minimum time over the replicates of each cell of the grid, along - with the time per causal site. + with the time per causal site and, where we counted the work, the time per + trip of the kernel's innermost loop. """ best = {} for phase, num_causal, selection, level, _, seconds in rows: @@ -208,15 +514,22 @@ def summarise(rows, completed, ts, args): print(f"\n{describe(ts)}") print(f"Minimum of {args.replicates} replicates\n") - header = ( - f"{'phase':<14} {'selection':<10} {'level':<11} {'num_causal':>10} " - f"{'seconds':>10} {'us/site':>10}" - ) - print(header) - print("-" * len(header)) + columns = f"{'phase':<14} {'selection':<10} {'level':<11} {'num_causal':>10} " + columns += f"{'seconds':>10} {'us/site':>10}" + if args.counters: + columns += f" {'ns/scan':>9}" + print(columns) + print("-" * len(columns)) + phases = ["sim_trait", "genetic_value"] + if args.phases: + phases += ["check", "setup", "kernel", "frame"] combinations = [("sim_trait", "uniform", "")] combinations += [ - ("genetic_value", s, x) for s in args.selections for x in args.levels + (p, s, x) + for p in phases + if p != "sim_trait" + for s in args.selections + for x in args.levels ] for phase, selection, level in combinations: for num_causal in completed: @@ -224,10 +537,60 @@ def summarise(rows, completed, ts, args): if key not in best: continue seconds = best[key] - print( + line = ( f"{phase:<14} {selection:<10} {level:<11} {num_causal:>10} " f"{seconds:>10.3f} {seconds / num_causal * 1e6:>10.1f}" ) + if args.counters: + # Only the phases that run the sweep have a per trip cost. + scans = counts.get((num_causal, selection), {}).get("edge_scans") + sweeps = phase in ("genetic_value", "kernel") + line += ( + f" {seconds / scans * 1e9:>9.2f}" + if sweeps and scans is not None + else f" {'':>9}" + ) + print(line) + + if args.counters: + print() + header = f"{'selection':<10} {'num_causal':>10} " + header += " ".join(f"{name:>12}" for name in COUNTERS) + header += f" {'scans/seed':>11} {'hit rate':>9} {'reached':>8}" + print(header) + print("-" * len(header)) + # The sweep visits every node whether or not anything reached it, and + # builds a pending entry for every node before it starts, so a low + # reached fraction is a kernel spending its time on the prologue. + for selection in args.selections: + for num_causal in completed: + got = counts.get((num_causal, selection)) + if got is None: + continue + line = f"{selection:<10} {num_causal:>10} " + line += " ".join(f"{got[name]:>12,}" for name in COUNTERS) + line += f" {got['edge_scans'] / max(got['seeds'], 1):>11.1f}" + line += f" {got['edge_hits'] / max(got['edge_scans'], 1) * 100:>8.1f}%" + line += f" {got['reached'] / ts.num_nodes * 100:>7.1f}%" + print(line) + + if len(memory) > 0: + print() + header = f"{'selection':<10} {'level':<11} {'num_causal':>10} " + header += f"{'peak RSS':>12} {'over base':>12}" + print(header) + print("-" * len(header)) + for (num_causal, selection, level), (peak, base) in memory.items(): + if peak is None: + shown, added = "unavailable", "" + else: + # The peak is mostly the tree sequence and the interpreter, so + # what the call added on top of them is the interesting number. + shown = f"{peak / 1e9:.2f} GB" + added = f"{(peak - base) / 1e9:.2f} GB" + print( + f"{selection:<10} {level:<11} {num_causal:>10} {shown:>12} {added:>12}" + ) def write_csv(rows, ts, args): @@ -265,18 +628,55 @@ def write_csv(rows, ts, args): print(f"\nWrote {args.output}") +def write_counts_csv(counts, ts, args): + path = args.output.with_name(args.output.stem + "_counters.csv") + with open(path, "w", newline="") as f: + writer = csv.writer(f) + writer.writerow(["num_causal", "selection", *COUNTERS, "num_nodes", "num_edges"]) + for (num_causal, selection), got in counts.items(): + writer.writerow( + [ + num_causal, + selection, + *[got[name] for name in COUNTERS], + ts.num_nodes, + ts.num_edges, + ] + ) + print(f"Wrote {path}") + + +def write_memory_csv(memory, args): + path = args.output.with_name(args.output.stem + "_memory.csv") + with open(path, "w", newline="") as f: + writer = csv.writer(f) + writer.writerow( + ["num_causal", "selection", "level", "peak_rss_bytes", "base_rss_bytes"] + ) + for (num_causal, selection, level), (peak, base) in memory.items(): + writer.writerow([num_causal, selection, level, peak, base]) + print(f"Wrote {path}") + + def parse_args(): default_output = pathlib.Path(__file__).parent / "_output" parser = argparse.ArgumentParser(description=__doc__) parser.add_argument( - "--samples", type=int, default=100_000, help="Number of sample nodes" + "--preset", + choices=list(PRESETS), + default="small", + help=( + "Simulation size. Sets --samples, --length and --num-causal, each " + "of which overrides it when given explicitly" + ), ) - parser.add_argument("--length", type=float, default=5e6, help="Sequence length") + parser.add_argument("--samples", type=int, help="Number of sample nodes") + parser.add_argument("--length", type=float, help="Sequence length") parser.add_argument("--recombination-rate", type=float, default=1e-8) parser.add_argument("--mutation-rate", type=float, default=1e-7) parser.add_argument("--population-size", type=float, default=10_000) parser.add_argument("--seed", type=int, default=42) - parser.add_argument("--num-causal", type=int, nargs="+", default=DEFAULT_NUM_CAUSAL) + parser.add_argument("--num-causal", type=int, nargs="+") parser.add_argument("--levels", nargs="+", choices=LEVELS, default=LEVELS) parser.add_argument( "--selections", nargs="+", choices=SELECTIONS, default=SELECTIONS @@ -297,11 +697,41 @@ def parse_args(): "taken longer than this" ), ) + parser.add_argument( + "--phases", + action="store_true", + help="Also time the check, setup, kernel and dataframe phases separately", + ) + parser.add_argument( + "--counters", + action="store_true", + help="Also count the work the kernel does, which perf cannot see inside", + ) + parser.add_argument( + "--structure", + action="store_true", + help="Report the shape of the tree sequence and the carrier fractions", + ) + parser.add_argument( + "--structure-sites", + type=int, + default=150, + help="Causal sites to sample for the carrier fraction distribution", + ) + parser.add_argument( + "--memory", + action="store_true", + help="Also measure the peak resident set size of each call", + ) parser.add_argument("--cache-dir", type=pathlib.Path, default=default_output) parser.add_argument( "--output", type=pathlib.Path, default=default_output / "genetic_value.csv" ) - return parser.parse_args() + args = parser.parse_args() + for name, value in PRESETS[args.preset].items(): + if getattr(args, name) is None: + setattr(args, name, value) + return args def main(): @@ -315,9 +745,13 @@ def main(): f"Only {ts.num_sites} sites in the tree sequence, but {max_causal} " "causal sites were requested. Increase --length or --mutation-rate." ) - rows, completed = run_benchmark(ts, args) - summarise(rows, completed, ts, args) + rows, counts, memory, completed = run_benchmark(ts, args) + summarise(rows, counts, memory, completed, ts, args) write_csv(rows, ts, args) + if args.counters: + write_counts_csv(counts, ts, args) + if args.memory: + write_memory_csv(memory, args) if __name__ == "__main__": diff --git a/benchmarks/profile_genetic_value.py b/benchmarks/profile_genetic_value.py new file mode 100644 index 0000000..9e4bc22 --- /dev/null +++ b/benchmarks/profile_genetic_value.py @@ -0,0 +1,226 @@ +""" +Profile a single cell of the genetic_value benchmark grid. + +The benchmark says how long a call takes; this says where the time goes. Two +modes, because the Python setup and the numba kernel need different tools: + +``--mode python`` + cProfile around the public ``genetic_value`` call. This is the only way to + see the setup, which is a flat per call cost dominated by + ``tskit.jit.numba.jitwrap``, but the kernel appears in it as a single + opaque dispatcher frame. + +``--mode kernel`` + Set up one cell and run only the kernel, so that a sampling profile is not + swamped by the effect size simulation over the whole site pool. Run it + under perf; ``--mode kernel`` on its own prints the commands. + +Two things about perf and numba are worth knowing before reading a profile. + +perf cannot attribute time to source lines inside the kernel here: this +llvmlite has no LLVM PerfJITEventListener, so nothing emits a +/tmp/perf-.map or a jitdump and the kernel shows up as raw addresses +under ``[JIT]``. What perf does give is the split between the kernel, the +numba runtime, the interpreter and LLVM compilation, with the typed list +runtime functions resolved by name. For attribution inside the kernel use +``benchmark_genetic_value.py --counters``. + +NUMBA_ENABLE_PROFILING=1 is the documented way to profile numba and is the +wrong thing to use here. It would only help by way of the listener llvmlite +does not have, and it defaults NUMBA_DEBUGINFO to 1, which changes the code +that is generated: the kernel measured 2.62s a call with it and 1.61s +without. perf finds the JIT mappings on its own, so the recipe below does not +set it, and a profile taken with it is a profile of a different kernel. + +Run with ``uv run --group test benchmarks/profile_genetic_value.py``. +""" + +import argparse +import cProfile +import pathlib +import pstats +import sys +import time + +import numpy as np +from benchmark_genetic_value import ( + LEVELS, + PRESETS, + SELECTIONS, + cached_simulation, + causal_site_pool, + describe, + select_causal, +) + +import tstrait +from tstrait import jit +from tstrait.genetic_value import _check_trait_df, _GeneticValue + +PERF_RECORD = """\ +sudo perf record -F 999 -g -o {data} -- \\ + {python} {script} {arguments} --run""" + +PERF_REPORT = """\ +sudo perf report -i {data} --stdio --no-children -g none --sort dso +sudo perf report -i {data} --stdio --no-children -g none \\ + --dsos={helper} --percent-limit 0.01""" + + +def prepare(args): + """ + Return the tree sequence and the trait dataframe for one cell of the grid. + """ + ts = cached_simulation(args) + print(describe(ts), file=sys.stderr) + model = tstrait.trait_model(distribution="normal", mean=0, var=1) + pool = causal_site_pool(ts, model, args.seed) + rng = np.random.default_rng(args.seed) + trait_df = select_causal( + pool, args.selection, args.num_causal, args.rare_threshold, rng + ) + if trait_df is None: + raise ValueError( + f"Too few {args.selection} sites for num_causal={args.num_causal}" + ) + return ts, trait_df + + +def profile_python(args): + """ + cProfile the public call, which is where the setup path is visible. + """ + ts, trait_df = prepare(args) + # Neither kernel is compiled with cache=True, so without this the profile + # is mostly LLVM. The lazily built tskit mutation state arrays are cached + # on the tree sequence by the same call. + tstrait.genetic_value(ts, trait_df, level=args.level) + + profiler = cProfile.Profile() + profiler.enable() + for _ in range(args.repeats): + tstrait.genetic_value(ts, trait_df, level=args.level) + profiler.disable() + pstats.Stats(profiler).sort_stats("cumulative").print_stats(args.lines) + + +def run_kernel(args): + """ + Run only the kernel, the loop a sampling profiler should be pointed at. + """ + ts, trait_df = prepare(args) + genetic = _GeneticValue(ts, _check_trait_df(ts, trait_df)) + arguments = genetic._descent_arguments(args.level, 0) + size = len(arguments["output"]) + jit._push_down_arg(**{**arguments, "output": np.zeros(size)}) + + before = time.perf_counter() + for _ in range(args.repeats): + jit._push_down_arg(**{**arguments, "output": np.zeros(size)}) + elapsed = time.perf_counter() - before + print( + f"{args.repeats} kernel calls in {elapsed:.3f}s " + f"({elapsed / args.repeats:.3f}s each)", + file=sys.stderr, + ) + + +def print_perf_commands(args): + """ + Print the perf invocation for this cell rather than running it, since it + needs sudo and writing out the command is more useful than hiding it. + """ + script = pathlib.Path(__file__).resolve() + arguments = ( + f"--preset {args.preset} --num-causal {args.num_causal} " + f"--selection {args.selection} --level {args.level} " + f"--repeats {args.repeats} --mode kernel" + ) + print( + f"# {args.repeats} kernel calls, against a setup of a few seconds that " + "is in the\n# profile too. Raise --repeats until the [JIT] share stops " + "moving.\n" + ) + print("# Record:") + print( + PERF_RECORD.format( + data=args.perf_data, + python=sys.executable, + script=script, + arguments=arguments, + ) + ) + print("\n# Report, first the split by shared object, then the numba runtime:") + print( + PERF_REPORT.format( + data=args.perf_data, + helper="_helperlib.cpython-*.so", + ) + ) + print( + "\n# The kernel is the [JIT] rows, and libllvmlite is compilation, not\n" + "# work. numba_list_append and numba_list_resize in the second report\n" + "# are the typed list traffic. Source lines inside the kernel are not\n" + "# available; use --counters on the benchmark instead." + ) + + +def parse_args(): + default_output = pathlib.Path(__file__).parent / "_output" + parser = argparse.ArgumentParser(description=__doc__) + parser.add_argument("--preset", choices=list(PRESETS), default="small") + parser.add_argument("--samples", type=int) + parser.add_argument("--length", type=float) + parser.add_argument("--recombination-rate", type=float, default=1e-8) + parser.add_argument("--mutation-rate", type=float, default=1e-7) + parser.add_argument("--population-size", type=float, default=10_000) + parser.add_argument("--seed", type=int, default=42) + parser.add_argument("--num-causal", type=int, default=10_000) + parser.add_argument("--selection", choices=SELECTIONS, default="uniform") + parser.add_argument("--level", choices=LEVELS, default="node") + parser.add_argument("--rare-threshold", type=float, default=0.001) + parser.add_argument( + "--repeats", + type=int, + help="Calls to profile. Defaults to 3 for --mode python, 10 for kernel", + ) + parser.add_argument( + "--mode", + choices=["python", "kernel"], + default="python", + help="Profile the Python setup path, or run just the kernel for perf", + ) + parser.add_argument( + "--run", + action="store_true", + help="With --mode kernel, run it rather than printing the perf commands", + ) + parser.add_argument( + "--lines", type=int, default=30, help="Rows of cProfile output to print" + ) + parser.add_argument("--cache-dir", type=pathlib.Path, default=default_output) + parser.add_argument( + "--perf-data", type=pathlib.Path, default=default_output / "perf.data" + ) + args = parser.parse_args() + if args.repeats is None: + args.repeats = 3 if args.mode == "python" else 10 + for name, value in PRESETS[args.preset].items(): + if name != "num_causal" and getattr(args, name) is None: + setattr(args, name, value) + return args + + +def main(): + sys.stdout.reconfigure(line_buffering=True) + args = parse_args() + if args.mode == "python": + profile_python(args) + elif args.run: + run_kernel(args) + else: + print_perf_commands(args) + + +if __name__ == "__main__": + main() From 0273e8ca47237a97f9f10a867d3ba96059927b45 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Thu, 3 Sep 2026 11:46:15 +0100 Subject: [PATCH 06/19] Make the individual level an output mapping of the descent The three levels differ only in which slot a contribution is credited to, and the kernel already discards a negative slot, so individuals need no more than a node to individual mapping: nodes_individual is exactly that array, and tskit.NULL for a node belonging to no individual is already negative. That removes the separate O(num_nodes) accumulation pass the individual level ran after the descent, one per trait, and with it _accumulate_individual_values. Accumulate into the caller's output array rather than allocating one per trait, so that a second implementation can write into the same table and combining the two is nothing at all. Split the causal mutation expansion in two. A descent of the trees needs every mutation at a causal site, because a mutation blocks the inheritance of the allele above it whatever it changes the state to, where a sum of state changes needs only those that change the state. _row_mutations returns all of them and _causal_mutations filters, which is what the push down seeds and edge_effect keep using. Keep the row each seed came from for the same reason: a subset of the rows can then be selected without recomputing the seeds. No change in what is computed. Verified against the tree by tree implementation at 31fec9a on 500 causal sites over a 30,000 sample tree sequence, at each level and for one and three traits, agreeing to 3e-14. --- benchmarks/benchmark_genetic_value.py | 6 +- benchmarks/profile_genetic_value.py | 4 +- tests/test_jit.py | 150 ++++++++++++-------------- tstrait/genetic_value.py | 81 +++++++++----- tstrait/jit.py | 17 --- 5 files changed, 129 insertions(+), 129 deletions(-) diff --git a/benchmarks/benchmark_genetic_value.py b/benchmarks/benchmark_genetic_value.py index 35ba92e..da6920b 100644 --- a/benchmarks/benchmark_genetic_value.py +++ b/benchmarks/benchmark_genetic_value.py @@ -295,7 +295,8 @@ def time_phases(ts, trait_df, level, replicates): phases.append(("setup", times)) genetic = _GeneticValue(ts, checked) - arguments = genetic._descent_arguments(level, 0) + size = genetic._output_size(level) + arguments = genetic._descent_arguments(level, 0, np.zeros(size)) _, times = time_call( # A fresh output array each time, since the kernel accumulates into it. lambda: jit._push_down_arg( @@ -305,7 +306,6 @@ def time_phases(ts, trait_df, level, replicates): ) phases.append(("kernel", times)) - size = len(arguments["output"]) values = np.zeros((genetic.num_trait, size)) _, times = time_call( functools.partial(_build_frame, genetic.num_trait, size, level, values), @@ -336,7 +336,7 @@ def count_work(ts, trait_df): three and only the array the contributions land in differs. """ genetic = _GeneticValue(ts, trait_df) - arguments = genetic._descent_arguments("node", 0) + arguments = genetic._descent_arguments("node", 0, np.zeros(ts.num_nodes)) counts = _count_push_down_arg( arguments["child_index"], arguments["edges_child"], diff --git a/benchmarks/profile_genetic_value.py b/benchmarks/profile_genetic_value.py index 9e4bc22..a18f7a0 100644 --- a/benchmarks/profile_genetic_value.py +++ b/benchmarks/profile_genetic_value.py @@ -110,8 +110,8 @@ def run_kernel(args): """ ts, trait_df = prepare(args) genetic = _GeneticValue(ts, _check_trait_df(ts, trait_df)) - arguments = genetic._descent_arguments(args.level, 0) - size = len(arguments["output"]) + size = genetic._output_size(args.level) + arguments = genetic._descent_arguments(args.level, 0, np.zeros(size)) jit._push_down_arg(**{**arguments, "output": np.zeros(size)}) before = time.perf_counter() diff --git a/tests/test_jit.py b/tests/test_jit.py index b36cce0..6a89092 100644 --- a/tests/test_jit.py +++ b/tests/test_jit.py @@ -85,25 +85,8 @@ def f( } ) genetic = _GeneticValue(ts, trait_df) - return func(**genetic._descent_arguments(level, 0)) - - return f - - -@pytest.fixture(params=["jit", "nojit"]) -def individual_genetic_value(request): - """ - Return a function accumulating node genetic values over the individuals of - a tree sequence. - """ - func = kernel(jit._accumulate_individual_values, request.param) - - def f(ts, nodes_genetic_value): - return func( - np.asarray(nodes_genetic_value, dtype=float), - ts.nodes_individual, - ts.num_individuals, - ) + output = np.zeros(genetic._output_size(level)) + return func(**genetic._descent_arguments(level, 0, output)) return f @@ -394,47 +377,11 @@ def test_multiple_roots_one_blocked(self, node_genetic_value): ) -class TestAccumulateIndividualValues: - """ - Sum the node genetic values over the nodes of each individual, using the - tree sequences of tests.data. +class TestIndividualLevel: """ - - def test_binary_tree(self, individual_genetic_value): - # Individual 0 is nodes 4 and 5, individual 1 nodes 0 and 1, and - # individual 2 nodes 2 and 3. - ts = binary_tree() - np.testing.assert_array_equal( - individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64]), [48, 3, 12] - ) - - def test_diff_ind_tree(self, individual_genetic_value): - ts = diff_ind_tree() - np.testing.assert_array_equal( - individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64]), [48, 5, 10] - ) - - def test_triploid_tree(self, individual_genetic_value): - # Individual 0 is nodes 0, 2 and 4, and individual 1 is nodes 1, 3 - # and 5. - ts = triploid_tree() - np.testing.assert_array_equal( - individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64, 128]), [21, 42] - ) - - def test_no_individuals(self, individual_genetic_value): - ts = tskit.Tree.generate_balanced(4).tree_sequence - assert ts.num_individuals == 0 - np.testing.assert_array_equal( - individual_genetic_value(ts, np.ones(ts.num_nodes)), [] - ) - - -class TestNodeAndIndividualValues: - """ - The two kernels composed, as they are used in tstrait.genetic_value, so - that the individual values can be traced back to the tree topology. The - tree of tests.data.binary_tree is:: + The individual level, which is the node level with a different output + mapping rather than a pass of its own, over the tree sequences of + tests.data. The tree of tests.data.binary_tree is:: 6 +-+-+ @@ -442,29 +389,33 @@ class TestNodeAndIndividualValues: +++ +++ 0 1 2 3 - with individual 0 being nodes 4 and 5, individual 1 nodes 0 and 1, and - individual 2 nodes 2 and 3. + A mutation on node 4 is inherited by nodes 0 and 1, which lands on + different individuals in each of these tree sequences and so pins the + mapping rather than only the total. """ - def test_internal_node(self, node_genetic_value, individual_genetic_value): - ts = binary_tree() - value = node_genetic_value(one_site(ts.first(), [(4, "T")]), causal_allele="T") - np.testing.assert_array_equal(value, [1, 1, 0, 0, 1, 0, 0]) - # Individual 0 has one copy through node 4, individual 1 has two - # copies through nodes 0 and 1, and individual 2 has none. - np.testing.assert_array_equal(individual_genetic_value(ts, value), [1, 2, 0]) - - def test_ancestral_state_is_causal( - self, node_genetic_value, individual_genetic_value - ): - # The causal allele is the ancestral state, so every node carries it - # and every diploid has two copies. + def test_binary_tree(self, node_genetic_value): + # Individual 0 is nodes 4 and 5, individual 1 nodes 0 and 1, and + # individual 2 nodes 2 and 3, so individual 1 has two copies. ts = binary_tree() - value = node_genetic_value(one_site(ts.first(), [])) - np.testing.assert_array_equal(value, [1, 1, 1, 1, 1, 1, 1]) - np.testing.assert_array_equal(individual_genetic_value(ts, value), [2, 2, 2]) + site = one_site(ts.first(), [(4, "T")]) + np.testing.assert_array_equal( + node_genetic_value(site, causal_allele="T"), [1, 1, 0, 0, 1, 0, 0] + ) + np.testing.assert_array_equal( + node_genetic_value(site, causal_allele="T", level="individual"), [1, 2, 0] + ) - def test_triploid(self, node_genetic_value, individual_genetic_value): + def test_diff_ind_tree(self, node_genetic_value): + # The same tree and mutation, but individual 1 is nodes 0 and 2 and + # individual 2 is nodes 1 and 3, so the two copies are split. + ts = diff_ind_tree() + value = node_genetic_value( + one_site(ts.first(), [(4, "T")]), causal_allele="T", level="individual" + ) + np.testing.assert_array_equal(value, [1, 1, 1]) + + def test_triploid_tree(self, node_genetic_value): # tests.data.triploid_tree, in which node 6 is the parent of nodes 3, # 4 and 5 and node 7 is the root: # @@ -474,9 +425,44 @@ def test_triploid(self, node_genetic_value, individual_genetic_value): # | | | +-+-+ # 0 1 2 3 4 5 # - ts = triploid_tree() - value = node_genetic_value(one_site(ts.first(), [(6, "T")]), causal_allele="T") - np.testing.assert_array_equal(value, [0, 0, 0, 1, 1, 1, 1, 0]) # Individual 0 is nodes 0, 2 and 4, and individual 1 is nodes 1, 3 # and 5. - np.testing.assert_array_equal(individual_genetic_value(ts, value), [1, 2]) + ts = triploid_tree() + site = one_site(ts.first(), [(6, "T")]) + np.testing.assert_array_equal( + node_genetic_value(site, causal_allele="T"), [0, 0, 0, 1, 1, 1, 1, 0] + ) + np.testing.assert_array_equal( + node_genetic_value(site, causal_allele="T", level="individual"), [1, 2] + ) + + def test_ancestral_state_is_causal(self, node_genetic_value): + # The causal allele is the ancestral state, so every node carries it + # and every diploid has two copies. + ts = binary_tree() + value = node_genetic_value(one_site(ts.first(), []), level="individual") + np.testing.assert_array_equal(value, [2, 2, 2]) + + def test_no_individuals(self, node_genetic_value): + # Nodes belong to no individual, so every contribution is discarded. + tree = tskit.Tree.generate_balanced(4) + assert tree.tree_sequence.num_individuals == 0 + value = node_genetic_value( + one_site(tree, [(4, "T")]), causal_allele="T", level="individual" + ) + np.testing.assert_array_equal(value, []) + + def test_matches_node_level(self, node_genetic_value): + # The individual values are the node values summed over each + # individual's nodes, however the nodes are assigned. + ts = diff_ind_tree() + site = one_site(ts.first(), [(4, "T")]) + nodes = node_genetic_value(site, causal_allele="T") + expected = np.bincount( + ts.nodes_individual[ts.nodes_individual >= 0], + weights=nodes[ts.nodes_individual >= 0], + minlength=ts.num_individuals, + ) + np.testing.assert_array_equal( + node_genetic_value(site, causal_allele="T", level="individual"), expected + ) diff --git a/tstrait/genetic_value.py b/tstrait/genetic_value.py index 1ff2f0b..d5ab1c9 100644 --- a/tstrait/genetic_value.py +++ b/tstrait/genetic_value.py @@ -7,12 +7,17 @@ from .base import _check_dataframe, _check_instance, _check_non_decreasing # noreorder -def _causal_mutations(ts, trait_df): +def _row_mutations(ts, trait_df): """ Expand the trait dataframe to one entry per (row, mutation at that row's - site) pair, and return the row, the mutation and the change in causal - allele state for the mutations that change the state. Mutations that leave - the state unchanged contribute nothing and are dropped. + site) pair, returning the row, the mutation, whether the mutation carries + the causal allele and whether the state it replaced did. + + Every mutation at a causal site is returned, including those that leave the + causal allele state unchanged. A descent of the trees needs all of them, + because a mutation blocks the inheritance of the allele above it whatever + it changes the state to; ``_causal_mutations`` drops the ones that a sum of + state changes does not need. """ site_id = trait_df["site_id"].to_numpy() # Mutations are sorted by site, so each site owns a contiguous run of IDs. @@ -30,8 +35,18 @@ def _causal_mutations(ts, trait_df): )[row] had_causal_allele = ts.mutations_inherited_state[mutation] == causal_allele has_causal_allele = derived_state[mutation] == causal_allele - state_change = has_causal_allele.astype(np.int8) - had_causal_allele.astype(np.int8) + return row, mutation, has_causal_allele, had_causal_allele + +def _causal_mutations(ts, trait_df): + """ + Expand the trait dataframe to one entry per (row, mutation at that row's + site) pair, and return the row, the mutation and the change in causal + allele state for the mutations that change the state. Mutations that leave + the state unchanged contribute nothing and are dropped. + """ + row, mutation, has_causal_allele, had_causal_allele = _row_mutations(ts, trait_df) + state_change = has_causal_allele.astype(np.int8) - had_causal_allele.astype(np.int8) changed = state_change != 0 return row[changed], mutation[changed], state_change[changed] @@ -144,6 +159,9 @@ def __init__(self, ts, trait_df): self.nodes_by_time = np.argsort(-ts.nodes_time, kind="stable").astype(np.int32) row, mutation, state_change = _causal_mutations(ts, self.trait_df) + # The row a seed came from, so that a subset of the rows can be + # selected without recomputing the seeds. + self.seed_row = row self.seed_trait = self.trait_id[row] self.seed_node = ts.mutations_node[mutation].astype(np.int32) self.seed_site = rows_site[row].astype(np.int32) @@ -172,6 +190,7 @@ def __init__(self, ts, trait_df): index = np.arange(count.sum()) + np.repeat( first - (np.cumsum(count) - count), count ) + self.seed_row = np.concatenate([self.seed_row, ancestral[index]]) self.seed_trait = np.concatenate( [self.seed_trait, self.trait_id[ancestral[index]]] ) @@ -188,10 +207,32 @@ def __init__(self, ts, trait_df): [self.seed_edge, np.full(len(run), tskit.NULL, dtype=np.int32)] ) - def _descent_arguments(self, level, trait): + def _output_size(self, level): + return { + "individual": self.ts.num_individuals, + "node": self.ts.num_nodes, + "edge": self.ts.num_edges, + }[level] + + def _node_output(self, level): + """ + Return the output slot that a contribution arriving at each node is + credited to, with a negative slot discarding it. + + Nodes and individuals differ only in this mapping, so the individual + values fall out of the same descent rather than needing a pass of their + own. A node belonging to no individual is tskit.NULL, which is already + negative. + """ + if level == "individual": + return self.ts.nodes_individual + return np.arange(self.ts.num_nodes, dtype=np.int32) + + def _descent_arguments(self, level, trait, output): """ Return the arguments to the push down kernel for one trait, with the - contributions directed at nodes or at edges according to ``level``. + contributions directed at nodes, individuals or edges according to + ``level`` and accumulated into ``output``. """ ts = self.ts in_trait = self.seed_trait == trait @@ -200,11 +241,10 @@ def _descent_arguments(self, level, trait): # A seed is credited to the edge above the mutation, which does not # exist when the mutation is above a root. seed_output = self.seed_edge[in_trait].astype(np.int32) - output_size = ts.num_edges else: - edges_output = ts.edges_child - seed_output = self.seed_node[in_trait] - output_size = ts.num_nodes + node_output = self._node_output(level) + edges_output = node_output[ts.edges_child] + seed_output = node_output[self.seed_node[in_trait]] return { "child_index": self.child_index, @@ -217,7 +257,7 @@ def _descent_arguments(self, level, trait): "seed_site": self.seed_site[in_trait], "seed_weight": self.seed_weight[in_trait], "seed_output": seed_output, - "output": np.zeros(output_size), + "output": output, } def _run(self, level): @@ -230,21 +270,12 @@ def _run(self, level): pandas.DataFrame Dataframe with trait ID, [individual|node|edge] ID, and genetic value. """ - ts = self.ts - N = { - "individual": ts.num_individuals, - "node": ts.num_nodes, - "edge": ts.num_edges, - }[level] - + N = self._output_size(level) genetic_value_table = np.zeros((self.num_trait, N)) for trait in range(self.num_trait): - output = jit._push_down_arg(**self._descent_arguments(level, trait)) - if level == "individual": - output = jit._accumulate_individual_values( - output, ts.nodes_individual, ts.num_individuals - ) - genetic_value_table[trait, :] = output + jit._push_down_arg( + **self._descent_arguments(level, trait, genetic_value_table[trait]) + ) return pd.DataFrame( { diff --git a/tstrait/jit.py b/tstrait/jit.py index 9aaf7a9..3506191 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -19,7 +19,6 @@ import numba import numpy as np -import tskit from numba.core import types from numba.typed import List @@ -115,19 +114,3 @@ def _push_down_arg( pending[parent] = empty return output - - -@numba.njit -def _accumulate_individual_values( - nodes_genetic_value, nodes_individual, num_individuals -): - """ - Accumulate the individual genetic values by summing their node - contributions. - """ - individuals_genetic_value = np.zeros(num_individuals) - for u in range(len(nodes_individual)): - ind = nodes_individual[u] - if ind != tskit.NULL: - individuals_genetic_value[ind] += nodes_genetic_value[u] - return individuals_genetic_value From ecef33ff465cc626eacbc6b645b420e7cd7ce6f2 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Thu, 3 Sep 2026 11:52:14 +0100 Subject: [PATCH 07/19] Build the trees in numba, wired to nothing yet A descent of a tree needs the tree, so add the incremental construction that a second implementation of genetic values will run over: apply the edges leaving and then entering at each tree of tskit's TreeIndex, maintaining the quintuply linked encoding along with the edge above each node. Attaching and detaching a child are both constant time, so a full pass costs one insertion and one removal per edge: 2.25ms over the 75,586 edges of a 30,000 sample tree sequence and 12.35ms over the 285,979 edges of a 100,000 sample one, which is less than the setup that the push down already pays on the same data. There is no virtual root. The descent starts at the mutations of a causal site and needs the roots only when the ancestral state is the causal allele, which is 0.04% of the sites of the smaller tree sequence and 0.07% of the larger. Maintaining the virtual root's children costs something on every edge, and knowing which nodes belong in that list means tracking the samples below every node; walking up from each sample and taking the top of the path finds the same roots on demand, and marking the nodes already walked through keeps the whole thing to one visit per node however many samples there are. That leaves one difference from tskit: it threads the roots together as children of the virtual root, so they are siblings of each other there and have none here. Nothing else differs, and the descent never walks the siblings of a root, so the test asserts the sibling arrays match at every node that has a parent and are null at the roots, along with the parent, left child and edge arrays and the root set itself, tree by tree over every fixture in tests.data, all_trees_ts(2..5), comb trees, multiple roots, isolated samples and a tree sequence with no nodes at all. --- tests/test_jit.py | 132 ++++++++++++++++++++++++++++++++++++++++++++++ tstrait/jit.py | 111 ++++++++++++++++++++++++++++++++++++++ 2 files changed, 243 insertions(+) diff --git a/tests/test_jit.py b/tests/test_jit.py index 6a89092..08af5be 100644 --- a/tests/test_jit.py +++ b/tests/test_jit.py @@ -14,17 +14,23 @@ a tree. """ +import numba import numpy as np import pandas as pd import pytest import tskit +import tskit.jit.numba as tskit_numba from tstrait import jit from tstrait.genetic_value import _GeneticValue from .data import ( + all_trees_ts, binary_tree, + binary_tree_seq, diff_ind_tree, + non_binary_tree, + simple_tree_seq, triploid_tree, ) # noreorder @@ -466,3 +472,129 @@ def test_matches_node_level(self, node_genetic_value): np.testing.assert_array_equal( node_genetic_value(site, causal_allele="T", level="individual"), expected ) + + +@numba.njit +def _walk_trees( + numba_ts, + edges_parent, + edges_child, + samples, + parent, + left_child, + right_sib, + node_edge, + roots, + num_roots, +): + """ + Build every tree in turn, recording the state of each one. + + This is the loop that drives the tree state kernels, kept here rather than + imported because the production driver does its per site work inside it and + has nowhere to record a tree. It is three lines; everything it calls is the + code under test. + """ + tree = jit.tree_state(numba_ts.num_nodes) + marked = np.full(numba_ts.num_nodes, -1, dtype=np.int64) + scratch = np.empty(numba_ts.num_nodes, dtype=np.int32) + tree_index = numba_ts.tree_index() + i = 0 + while tree_index.next(): + jit._apply_edge_diffs(tree_index, edges_parent, edges_child, tree) + parent[i, :] = tree.parent + left_child[i, :] = tree.left_child + right_sib[i, :] = tree.right_sib + node_edge[i, :] = tree.node_edge + n = jit._tree_roots(tree, samples, marked, i, scratch) + num_roots[i] = n + roots[i, :n] = scratch[:n] + i += 1 + return i + + +def walk_trees(ts): + """ + Return the per tree state arrays that _walk_trees records. + """ + numba_ts = tskit_numba.jitwrap(ts) + shape = (max(ts.num_trees, 1), ts.num_nodes) + got = { + name: np.zeros(shape, dtype=np.int32) + for name in ("parent", "left_child", "right_sib", "node_edge", "roots") + } + num_roots = np.zeros(shape[0], dtype=np.int32) + count = _walk_trees( + numba_ts, + ts.edges_parent, + ts.edges_child, + ts.samples().astype(np.int32), + got["parent"], + got["left_child"], + got["right_sib"], + got["node_edge"], + got["roots"], + num_roots, + ) + assert count == ts.num_trees + return got, num_roots + + +def unbalanced_ts(n): + """ + Return a comb tree sequence, the worst case for anything that walks from a + node to its root. + """ + return tskit.Tree.generate_comb(n).tree_sequence + + +class TestTreeState: + """ + The incrementally maintained tree must match what tskit builds, tree for + tree, over topologies that exercise multiple roots, isolated samples and + nodes that are in no tree at all. + """ + + @pytest.mark.parametrize( + "ts", + [ + binary_tree(), + diff_ind_tree(), + non_binary_tree(), + triploid_tree(), + binary_tree_seq(), + simple_tree_seq(), + all_trees_ts(2), + all_trees_ts(3), + all_trees_ts(4), + all_trees_ts(5), + unbalanced_ts(10), + unbalanced_ts(50), + tskit.Tree.generate_balanced(8).tree_sequence, + multiroot_tree().tree_sequence, + isolated_samples_ts(4), + empty_ts(), + ], + ) + def test_matches_tskit(self, ts): + got, num_roots = walk_trees(ts) + n = ts.num_nodes + for i, tree in enumerate(ts.trees()): + # tskit's arrays carry the virtual root in a final entry, which the + # kernel has no use for and does not keep. + np.testing.assert_array_equal(got["parent"][i, :n], tree.parent_array[:n]) + np.testing.assert_array_equal( + got["left_child"][i, :n], tree.left_child_array[:n] + ) + np.testing.assert_array_equal(got["node_edge"][i, :n], tree.edge_array[:n]) + # tskit threads the roots together as children of the virtual root, + # so they are siblings of each other there and have no sibling + # here. The descent never walks the siblings of a root, since it + # starts at the roots that _tree_roots gives it, so the two agree + # everywhere the descent looks. + child = tree.parent_array[:n] != tskit.NULL + np.testing.assert_array_equal( + got["right_sib"][i, :n][child], tree.right_sib_array[:n][child] + ) + assert np.all(got["right_sib"][i, :n][~child] == tskit.NULL) + assert sorted(got["roots"][i, : num_roots[i]]) == sorted(tree.roots) diff --git a/tstrait/jit.py b/tstrait/jit.py index 3506191..645c7a9 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -17,8 +17,11 @@ coverage measurement for free. """ +from collections import namedtuple + import numba import numpy as np +import tskit from numba.core import types from numba.typed import List @@ -114,3 +117,111 @@ def _push_down_arg( pending[parent] = empty return output + + +# The state of one tree, in the tskit quintuply linked encoding. A tree is +# built by applying the edge differences between it and the tree before it, so +# a full pass over the trees costs one insertion and one removal per edge. +# There is no virtual root here: the descent starts at the mutations of a +# causal site, and the roots are only needed when the ancestral state is the +# causal allele, which _tree_roots finds on demand. +_TreeState = namedtuple( + "_TreeState", + ["parent", "left_child", "right_child", "left_sib", "right_sib", "node_edge"], +) + + +@numba.njit +def tree_state(num_nodes): + """ + Return the arrays holding a tree, in the state of a tree with no edges. + """ + return _TreeState( + parent=np.full(num_nodes, tskit.NULL, dtype=np.int32), + left_child=np.full(num_nodes, tskit.NULL, dtype=np.int32), + right_child=np.full(num_nodes, tskit.NULL, dtype=np.int32), + left_sib=np.full(num_nodes, tskit.NULL, dtype=np.int32), + right_sib=np.full(num_nodes, tskit.NULL, dtype=np.int32), + node_edge=np.full(num_nodes, tskit.NULL, dtype=np.int32), + ) + + +@numba.njit +def _remove_edge(tree, parent_node, child_node): + """ + Detach ``child_node`` from ``parent_node``, unlinking it from its siblings. + """ + left_sib = tree.left_sib[child_node] + right_sib = tree.right_sib[child_node] + if left_sib == tskit.NULL: + tree.left_child[parent_node] = right_sib + else: + tree.right_sib[left_sib] = right_sib + if right_sib == tskit.NULL: + tree.right_child[parent_node] = left_sib + else: + tree.left_sib[right_sib] = left_sib + tree.parent[child_node] = tskit.NULL + tree.left_sib[child_node] = tskit.NULL + tree.right_sib[child_node] = tskit.NULL + tree.node_edge[child_node] = tskit.NULL + + +@numba.njit +def _insert_edge(tree, edge, parent_node, child_node): + """ + Attach ``child_node`` to ``parent_node`` as its rightmost child, through + ``edge``. + """ + right_child = tree.right_child[parent_node] + if right_child == tskit.NULL: + tree.left_child[parent_node] = child_node + tree.left_sib[child_node] = tskit.NULL + else: + tree.right_sib[right_child] = child_node + tree.left_sib[child_node] = right_child + tree.right_child[parent_node] = child_node + tree.right_sib[child_node] = tskit.NULL + tree.parent[child_node] = parent_node + tree.node_edge[child_node] = edge + + +@numba.njit +def _apply_edge_diffs(tree_index, edges_parent, edges_child, tree): + """ + Advance ``tree`` to the tree that ``tree_index`` has moved to, by applying + the edges leaving and then the edges entering. + """ + for j in range(tree_index.out_range.start, tree_index.out_range.stop): + edge = tree_index.out_range.order[j] + _remove_edge(tree, edges_parent[edge], edges_child[edge]) + for j in range(tree_index.in_range.start, tree_index.in_range.stop): + edge = tree_index.in_range.order[j] + _insert_edge(tree, edge, edges_parent[edge], edges_child[edge]) + + +@numba.njit +def _tree_roots(tree, samples, marked, mark, roots): + """ + Fill ``roots`` with the roots of the current tree and return how many there + are. + + A root is a node with no parent that has a sample at or below it, which is + what tskit calls a root, so walking up from each sample and taking the top + of the path finds all of them and nothing else. ``marked`` records the + nodes already walked through under the value ``mark``, so that the paths of + all the samples together cost one visit per node rather than one per + sample. + """ + num_roots = 0 + for i in range(len(samples)): + u = samples[i] + while marked[u] != mark: + marked[u] = mark + parent = tree.parent[u] + if parent == tskit.NULL: + roots[num_roots] = u + num_roots += 1 + break + u = parent + return num_roots From b54570a8c6c4fda3a39edbd4afa3271c39bdfd32 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Thu, 3 Sep 2026 12:00:04 +0100 Subject: [PATCH 08/19] Send common causal alleles down a descent of the trees The push down of the ARG costs the nodes that carry a causal allele, at around 28ns each because following carriers is random access. A descent of the trees costs the same nodes at around 20ns, since a child is the next link of a list rather than a search of the out edges of a node, and it needs no typed list per node to hold what is in flight. Against that it has to build the trees, which the push down never does. So neither wins everywhere, and which one a causal site should take follows from its allele frequency: the tree pass is a fixed cost that common sites amortise and rare ones do not. Add the descent as a second implementation and route each row to one or the other on allele_freq, which sim_trait already returns and _check_trait_df now keeps when it is there. A trait dataframe assembled by hand has no such column and takes the push down for everything, which is what it did before. Measured on 30,000 samples over 1Mb at level="node", end to end: 1,000 10,000 uniform before 195ms 1624ms after 129ms 1080ms 1.51x 1.50x rare before 16.3ms 31.0ms after 16.6ms 31.1ms Rare sites are below the threshold and take the same path as before, so the gain is on the common ones, which is where the time was: sites above a frequency of 0.03 are under a third of the sites and 94 to 97% of the time. The descent walks the rows in step with the trees, which works because _check_trait_df already requires the rows to be sorted by site. A row's mutations are marked with the row's own index rather than into an array that has to be cleared, so a descent costs the nodes it reaches and nothing per node of the tree sequence. All of the traits share the pass, where the push down runs once per trait. Two things the descent has to do that the push down does not. It takes every mutation at a causal site, not only the state changing ones, since a mutation blocks the inheritance of the allele above it whatever it changes the state to. And where the causal allele is the ancestral state it seeds the roots, skipping any root that carries a mutation: the push down seeds every root and lets the root's own mutation cancel it, and the descent has no such cancellation, so seeding an ancestral root that carries the allele would count it twice. Removing that skip fails test_mutation_on_root and two of the reference comparisons. Run every comparison against the tree by tree reference through both implementations and a mix of the two, and add tests for the cases where they are most likely to part company: a mutation on a root under both causal alleles, a causal site carrying no mutations at all, a site exactly on a breakpoint, and that a middling threshold really does divide the rows rather than quietly sending them all the same way. --- tests/test_genetic_value.py | 149 +++++++++++++++++++++++++++++++++++- tstrait/genetic_value.py | 107 ++++++++++++++++++++++---- tstrait/jit.py | 104 +++++++++++++++++++++++++ 3 files changed, 341 insertions(+), 19 deletions(-) diff --git a/tests/test_genetic_value.py b/tests/test_genetic_value.py index 05075ea..4a85157 100644 --- a/tests/test_genetic_value.py +++ b/tests/test_genetic_value.py @@ -6,6 +6,7 @@ import tstrait from tstrait.base import _check_numeric_array +from tstrait.genetic_value import _check_trait_df, _GeneticValue from .data import ( all_trees_ts, @@ -116,6 +117,11 @@ def random_trait_df(ts, num_trait, seed, site_id=None): The causal allele of a row is drawn from the alleles that occur at its site, so that the effects are not trivially zero. + + The allele_freq column alternates between two values rather than being the + real frequency. It is only there to decide which implementation a row goes + through, and alternating splits the rows whatever the topology is, which is + what the tests want to exercise. """ rng = np.random.default_rng(seed) if site_id is None: @@ -134,6 +140,9 @@ def random_trait_df(ts, num_trait, seed, site_id=None): "effect_size": rng.normal(size=num_site * num_trait), "trait_id": np.tile(np.arange(num_trait), num_site), "causal_allele": causal_allele, + "allele_freq": np.tile([0.25, 0.75], num_site * num_trait)[ + : num_site * num_trait + ], } ) @@ -1146,14 +1155,23 @@ class TestGeneticValueReference: roots with, so it is covered here rather than separately. """ + # All rows through the descent of the trees, split between the two, and + # all rows through the push down of the ARG. A row whose causal allele is + # the ancestral state always goes through the descent, so the last of these + # is not quite everything. + THRESHOLDS = [0.0, 0.5, np.inf] + def verify(self, ts, num_trait, seed=1, levels=("node", "edge")): trait_df = random_trait_df(ts, num_trait, seed) for level in levels: if level == "individual" and ts.num_individuals == 0: continue expected = naive_genetic_value(ts, trait_df, level) - result = tstrait.genetic_value(ts=ts, trait_df=trait_df, level=level) - pd.testing.assert_frame_equal(result, expected, check_dtype=False) + for threshold in self.THRESHOLDS: + result = tstrait.genetic_value( + ts=ts, trait_df=trait_df, level=level, _threshold=threshold + ) + pd.testing.assert_frame_equal(result, expected, check_dtype=False) @pytest.mark.parametrize( "ts_func", @@ -1257,3 +1275,130 @@ def test_mutation_above_a_root(self): expected = naive_genetic_value(ts, trait_df, level) result = tstrait.genetic_value(ts=ts, trait_df=trait_df, level=level) pd.testing.assert_frame_equal(result, expected, check_dtype=False) + + +class TestDescentAndPushDown: + """ + The two implementations against each other and against the reference, over + the cases where they are most likely to disagree. + + A row goes through the descent of the trees when its allele frequency is at + or above the threshold, so a threshold of zero sends everything there and + an infinite one sends everything to the push down, apart from the rows + whose causal allele is the ancestral state, which always take the descent. + """ + + def simulated_ts(self, seed=3): + ts = msprime.sim_ancestry( + 20, sequence_length=10_000, recombination_rate=1e-4, random_seed=seed + ) + return multi_allelic_mutations(ts, rate=1e-3, seed=seed) + + @pytest.mark.parametrize("level", ["individual", "node", "edge"]) + @pytest.mark.parametrize("num_trait", [1, 3]) + def test_paths_agree(self, level, num_trait): + # naive_genetic_value is too slow to run on anything this size, so the + # two implementations are checked against each other instead. + ts = self.simulated_ts() + trait_df = random_trait_df(ts, num_trait, seed=5) + descent = tstrait.genetic_value(ts, trait_df, level=level, _threshold=0.0) + push_down = tstrait.genetic_value(ts, trait_df, level=level, _threshold=np.inf) + pd.testing.assert_frame_equal(descent, push_down, check_dtype=False) + + def test_split_is_a_split(self): + # The reference tests would pass just as well if every row went the + # same way, so check that a middling threshold really does divide them. + ts = self.simulated_ts() + trait_df = _check_trait_df(ts, random_trait_df(ts, 2, seed=5)) + genetic = _GeneticValue(ts, trait_df, threshold=0.5) + assert np.any(genetic.descent_rows) + assert np.any(~genetic.descent_rows) + assert not np.all(genetic.descent_rows[genetic.seed_row]) + + def test_no_allele_freq_is_all_push_down(self): + # A trait dataframe assembled by hand has no allele_freq, and then + # there is nothing to route on. + ts = self.simulated_ts() + trait_df = _check_trait_df( + ts, random_trait_df(ts, 1, seed=5).drop(columns=["allele_freq"]) + ) + genetic = _GeneticValue(ts, trait_df, threshold=0.0) + assert not np.any(genetic.descent_rows) + + @pytest.mark.parametrize("derived_state", ["A", "T"]) + def test_mutation_on_root(self, derived_state): + # The ancestral state is the causal allele, so the descent seeds the + # roots, and node 6 is the root. It also carries a mutation, which + # replaced the ancestral state there, so seeding it as a root as well + # would count it twice when the mutation carries the causal allele and + # would count it at all when the mutation does not. + ts = mutated_binary_tree([(0, "A", [(6, derived_state)])]) + trait_df = pd.DataFrame( + { + "site_id": [0], + "effect_size": [1.0], + "trait_id": [0], + "causal_allele": ["A"], + "allele_freq": [0.5], + } + ) + expected = 1.0 if derived_state == "A" else 0.0 + for threshold in (0.0, np.inf): + result = tstrait.genetic_value( + ts, trait_df, level="node", _threshold=threshold + ) + np.testing.assert_array_equal( + result["genetic_value"], np.full(ts.num_nodes, expected) + ) + + def test_site_with_no_mutations(self): + # A causal site need not carry any mutation, and when its ancestral + # state is the causal allele every node in the tree carries it. + ts = mutated_binary_tree([(0, "A", [])]) + assert ts.num_mutations == 0 + trait_df = pd.DataFrame( + { + "site_id": [0], + "effect_size": [1.0], + "trait_id": [0], + "causal_allele": ["A"], + "allele_freq": [0.5], + } + ) + for threshold in (0.0, np.inf): + result = tstrait.genetic_value( + ts, trait_df, level="node", _threshold=threshold + ) + np.testing.assert_array_equal(result["genetic_value"], np.ones(ts.num_nodes)) + + def test_site_on_a_breakpoint(self): + # A site at a tree boundary belongs to the tree that starts there, + # which is where the descent has to find it. + ts = msprime.sim_ancestry( + 5, sequence_length=10, recombination_rate=0.5, random_seed=9 + ) + assert ts.num_trees > 1 + breakpoint_position = ts.breakpoints(as_array=True)[1] + tables = ts.dump_tables() + tables.sites.add_row(position=0, ancestral_state="A") + tables.sites.add_row(position=breakpoint_position, ancestral_state="A") + tables.sort() + for site in range(2): + tables.mutations.add_row(site=site, node=0, derived_state="T") + tables.sort() + ts = tables.tree_sequence() + trait_df = pd.DataFrame( + { + "site_id": [0, 1], + "effect_size": [1.0, 2.0], + "trait_id": [0, 0], + "causal_allele": ["T", "T"], + "allele_freq": [0.5, 0.5], + } + ) + expected = naive_genetic_value(ts, trait_df, "node") + for threshold in (0.0, np.inf): + result = tstrait.genetic_value( + ts, trait_df, level="node", _threshold=threshold + ) + pd.testing.assert_frame_equal(result, expected, check_dtype=False) diff --git a/tstrait/genetic_value.py b/tstrait/genetic_value.py index d5ab1c9..f349a4e 100644 --- a/tstrait/genetic_value.py +++ b/tstrait/genetic_value.py @@ -6,6 +6,18 @@ from . import jit from .base import _check_dataframe, _check_instance, _check_non_decreasing # noreorder +# The optional trait dataframe column that routes a causal site to one +# implementation or the other. +ALLELE_FREQ = "allele_freq" + +# Causal sites at or above this allele frequency go through the descent of the +# trees, and the rest through the push down of the ARG. The descent is faster +# almost everywhere, and measurably so from about a thousandth: below that the +# tree pass it starts with is not amortised and the push down, which never +# builds a tree, wins. Private while the two implementations are being +# compared, since the intention is to end with one of them. +_COMMON_THRESHOLD = 0.001 + def _row_mutations(ts, trait_df): """ @@ -92,10 +104,16 @@ def _check_trait_df(ts, trait_df): """ Check the trait dataframe against the tree sequence, returning the required columns with trait_id cast to int. + + ``allele_freq`` is kept when it is there. It is not required, and a trait + dataframe assembled by hand will not have it, but ``sim_trait`` returns it + and it is what decides which of the two implementations a causal site goes + through. """ - trait_df = _check_dataframe( - trait_df, ["site_id", "effect_size", "trait_id", "causal_allele"], "trait_df" - ) + columns = ["site_id", "effect_size", "trait_id", "causal_allele"] + if ALLELE_FREQ in getattr(trait_df, "columns", []): + columns.append(ALLELE_FREQ) + trait_df = _check_dataframe(trait_df, columns, "trait_df") if len(trait_df) == 0: raise ValueError("trait_df must contain at least one row") _check_non_decreasing(trait_df["site_id"], "site_id") @@ -133,10 +151,14 @@ class _GeneticValue: size, and trait ID. """ - def __init__(self, ts, trait_df): - self.trait_df = trait_df[["site_id", "effect_size", "trait_id", "causal_allele"]] + def __init__(self, ts, trait_df, threshold=_COMMON_THRESHOLD): + columns = ["site_id", "effect_size", "trait_id", "causal_allele"] + if ALLELE_FREQ in trait_df.columns: + columns.append(ALLELE_FREQ) + self.trait_df = trait_df[columns] self.ts = ts - self.child_index = tskit_numba.jitwrap(ts).child_index() + self.numba_ts = tskit_numba.jitwrap(ts) + self.child_index = self.numba_ts.child_index() site_id = self.trait_df["site_id"].to_numpy() effect_size = self.trait_df["effect_size"].to_numpy() @@ -158,7 +180,37 @@ def __init__(self, ts, trait_df): # tskit does not require the node IDs to be in time order. self.nodes_by_time = np.argsort(-ts.nodes_time, kind="stable").astype(np.int32) - row, mutation, state_change = _causal_mutations(ts, self.trait_df) + # Every mutation at a causal site, which is what the descent needs, + # and the state changing ones, which is what the push down seeds from. + pair_row, pair_mutation, has, had = _row_mutations(ts, self.trait_df) + self.row_site = site_id.astype(np.int32) + self.row_trait = self.trait_id.astype(np.int32) + self.row_effect = effect_size.astype(float) + self.row_ancestral = ( + ts.sites_ancestral_state[site_id] + == self.trait_df["causal_allele"].to_numpy() + ) + self.pair_offset = np.searchsorted( + pair_row, np.arange(len(self.trait_df) + 1) + ).astype(np.int64) + self.pair_node = ts.mutations_node[pair_mutation].astype(np.int32) + self.pair_carries = has + + # Rows go to the descent of the trees when their causal allele is + # common enough for the tree pass to pay for itself, and when it is the + # ancestral state, where the descent has the roots to hand and the push + # down would need them found for it. + if ALLELE_FREQ in self.trait_df.columns: + self.descent_rows = self.trait_df[ALLELE_FREQ].to_numpy() >= threshold + self.descent_rows |= self.row_ancestral + else: + self.descent_rows = np.zeros(len(self.trait_df), dtype=bool) + + state_change = has.astype(np.int8) - had.astype(np.int8) + changed = state_change != 0 + row = pair_row[changed] + mutation = pair_mutation[changed] + state_change = state_change[changed] # The row a seed came from, so that a subset of the rows can be # selected without recomputing the seeds. self.seed_row = row @@ -174,10 +226,9 @@ def __init__(self, ts, trait_df): # When the ancestral state of a site is the causal allele, every node # in its tree carries it. There is no mutation to seed from, so the # roots are seeded instead and the effect reaches the same nodes. - ancestral = np.flatnonzero( - ts.sites_ancestral_state[site_id] - == self.trait_df["causal_allele"].to_numpy() - ) + # Only for the rows the push down is taking: the descent seeds the + # roots of the tree it is already standing in. + ancestral = np.flatnonzero(self.row_ancestral & ~self.descent_rows) if len(ancestral) > 0: roots, roots_left, roots_right = _root_runs(ts) start = np.searchsorted(causal_position, roots_left) @@ -235,7 +286,7 @@ def _descent_arguments(self, level, trait, output): ``level`` and accumulated into ``output``. """ ts = self.ts - in_trait = self.seed_trait == trait + in_trait = (self.seed_trait == trait) & ~self.descent_rows[self.seed_row] if level == "edge": edges_output = np.arange(ts.num_edges, dtype=np.int32) # A seed is credited to the edge above the mutation, which does not @@ -270,12 +321,34 @@ def _run(self, level): pandas.DataFrame Dataframe with trait ID, [individual|node|edge] ID, and genetic value. """ + ts = self.ts N = self._output_size(level) genetic_value_table = np.zeros((self.num_trait, N)) - for trait in range(self.num_trait): - jit._push_down_arg( - **self._descent_arguments(level, trait, genetic_value_table[trait]) + + if np.any(self.descent_rows): + jit._descend_trees( + self.numba_ts, + ts.edges_parent, + ts.edges_child, + self.row_site, + self.row_trait, + self.row_effect, + self.row_ancestral, + self.descent_rows, + self.pair_offset, + self.pair_node, + self.pair_carries, + ts.samples().astype(np.int32), + self._node_output(level), + level == "edge", + genetic_value_table, ) + # Compiling the push down is not worth it when nothing is left for it. + if not np.all(self.descent_rows[self.seed_row]): + for trait in range(self.num_trait): + jit._push_down_arg( + **self._descent_arguments(level, trait, genetic_value_table[trait]) + ) return pd.DataFrame( { @@ -286,7 +359,7 @@ def _run(self, level): ) -def genetic_value(ts, trait_df, level="individual"): +def genetic_value(ts, trait_df, level="individual", *, _threshold=_COMMON_THRESHOLD): """ Compute genetic values for a tree sequence given a trait dataframe. @@ -331,7 +404,7 @@ def genetic_value(ts, trait_df, level="individual"): raise ValueError("No individuals in the provided tree sequence dataset") trait_df = _check_trait_df(ts, trait_df) - genetic = _GeneticValue(ts=ts, trait_df=trait_df) + genetic = _GeneticValue(ts=ts, trait_df=trait_df, threshold=_threshold) genetic_result = genetic._run(level) diff --git a/tstrait/jit.py b/tstrait/jit.py index 645c7a9..a46b831 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -225,3 +225,107 @@ def _tree_roots(tree, samples, marked, mark, roots): break u = parent return num_roots + + +@numba.njit +def _descend_trees( + numba_ts, + edges_parent, + edges_child, + row_site, + row_trait, + row_effect, + row_ancestral, + row_selected, + pair_offset, + pair_node, + pair_carries, + samples, + node_output, + edge_level, + output, +): + """ + Accumulate the genetic value of the selected causal sites in one pass over + the trees, descending from each causal mutation to the nodes that inherit + the causal allele from it. + + The rows are sorted by site and each site belongs to one tree, so a single + pointer walks the rows in step with the trees. A row's mutations are marked + with the row's own index rather than into a cleared array, so a descent + costs the nodes it reaches and nothing per node of the tree sequence. + + ``node_output`` gives the index in ``output`` that a node's contribution is + added to, with a negative index discarding it; at the edge level the index + is instead the edge above the node in the current tree, which is where the + contribution arrived from. All of the traits share the pass, since a row + knows the trait it belongs to. + """ + num_nodes = numba_ts.num_nodes + tree = tree_state(num_nodes) + # Marked with the row rather than cleared, so that nothing here costs a + # pass over the nodes. + stamp = np.full(num_nodes, -1, dtype=np.int64) + marked = np.full(num_nodes, -1, dtype=np.int64) + carries = np.zeros(num_nodes, dtype=np.bool_) + roots = np.empty(num_nodes, dtype=np.int32) + # A node is reached by at most one seed of a row, because the descent from + # a seed stops at the mutations of the row, so the tree bounds the stack. + stack = np.empty(num_nodes, dtype=np.int32) + + row = 0 + num_rows = len(row_site) + tree_index = numba_ts.tree_index() + while tree_index.next(): + _apply_edge_diffs(tree_index, edges_parent, edges_child, tree) + site_stop = tree_index.site_range[1] + while row < num_rows and row_site[row] < site_stop: + if not row_selected[row]: + row += 1 + continue + start = pair_offset[row] + stop = pair_offset[row + 1] + # Every mutation at the site blocks the allele above it, whatever + # it changes the state to. A node carrying more than one of them + # takes the last, which is the youngest since tskit orders a + # mutation after its parent. + for k in range(start, stop): + stamp[pair_node[k]] = row + carries[pair_node[k]] = pair_carries[k] + + top = 0 + for k in range(start, stop): + node = pair_node[k] + if carries[node]: + # Cleared so that a node carrying several mutations at this + # site is seeded once. + carries[node] = False + stack[top] = node + top += 1 + if row_ancestral[row]: + # The causal allele is the ancestral state, so it reaches every + # node the roots reach. A root carrying a mutation is not one + # of them: the mutation replaced the ancestral state there, and + # it has already been seeded above if it carries the allele. + num_roots = _tree_roots(tree, samples, marked, row, roots) + for i in range(num_roots): + if stamp[roots[i]] != row: + stack[top] = roots[i] + top += 1 + + trait = row_trait[row] + weight = row_effect[row] + while top > 0: + top -= 1 + node = stack[top] + slot = tree.node_edge[node] if edge_level else node_output[node] + if slot >= 0: + output[trait, slot] += weight + child = tree.left_child[node] + while child != tskit.NULL: + if stamp[child] != row: + stack[top] = child + top += 1 + child = tree.right_sib[child] + row += 1 + return output From ec3e541baf45b8d6377266977629e761679cddeb Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Thu, 3 Sep 2026 12:26:28 +0100 Subject: [PATCH 09/19] Default to the descent, which measured faster everywhere Compare the two implementations over the whole benchmark grid, on both presets, at all three levels, for one and three traits, and for causal sites drawn both uniformly and from the rare ones. A threshold of zero, sending everything to the descent of the trees, was fastest or tied in every cell: uniform rare 1,000 10,000 1,000 10,000 small descent 122ms 1058ms 15ms 23ms push down 196ms 1619ms 16ms 35ms large descent 427ms 3709ms 71ms 90ms push down 592ms 4917ms 68ms 108ms The push down led only where the causal sites are few and rare, so the tree pass is not amortised, and there it led by a millisecond or two on a call taking a few. With three traits the gap widens the other way, to 2.27x on rare causal sites, since the push down sweeps once per trait where the descent takes all of them in one pass. So set the threshold to zero. Every allele frequency is at or above zero, so there is nothing to compare against and nothing to look up, which is what lets the default apply to a trait dataframe assembled by hand as well as to one from sim_trait. The push down is kept, and --threshold inf still runs it, only so that the two can go on being compared. Against the tree by tree implementation this branch started from, at level="node" on 100,000 samples, the whole of it is now ahead: 1,000 10,000 100,000 tree by tree 0.436s 3.761s 38.374s rare 0.072s 0.091s 0.216s 6.1x 41x 178x uniform 0.443s 3.848s 32.589s 0.98x 0.98x 1.18x which closes the regression on uniformly drawn causal sites that the push down left behind at 0.70x. Report the work the descent does rather than the work the push down does, since that is the code that runs. The descent kernel returns the number of nodes it visited, so there is no second copy of the loop to keep in step with the first, and the count is the thing its run time is proportional to: seconds over visits is around 20ns once there are enough causal sites to amortise the tree pass. Visits per row as a fraction of the nodes comes to 8.4% for uniformly drawn causal sites and 0.02% for rare ones, against the 7.2% and 0.4% that --structure measures from the other end. The _root_runs branch is only taken by rows that go to the push down, so at the default threshold it never runs; say so rather than reporting that any ancestral causal allele triggers it. --- benchmarks/README.md | 146 ++++++++++------ benchmarks/baseline_small.csv | 170 +++++++++---------- benchmarks/baseline_small_counters.csv | 18 +- benchmarks/benchmark_genetic_value.py | 221 +++++++++---------------- tests/test_genetic_value.py | 12 +- tests/test_jit.py | 4 +- tstrait/genetic_value.py | 29 +++- tstrait/jit.py | 8 +- 8 files changed, 307 insertions(+), 301 deletions(-) diff --git a/benchmarks/README.md b/benchmarks/README.md index 7a3b6e0..973e48a 100644 --- a/benchmarks/README.md +++ b/benchmarks/README.md @@ -32,9 +32,9 @@ per tree is inflated. ### Presets -`--preset small` is the default and takes about 30 seconds; `--preset large` is -the tree sequence the ARG sweep was written against and takes about nine -minutes, with its longest single call at 48 seconds. +`--preset small` is the default and takes about 25 seconds; `--preset large` is +the tree sequence this work was written against and takes about six minutes, +with its longest single call at 33 seconds. | | small | large | |---|---|---| @@ -52,22 +52,22 @@ still overrides it when given explicitly. The tree sequence is cached in `_output/` under a name built from the simulation parameters, so `small` simulates itself in 0.2s the first time and is loaded after that. -`small` is the default because iterating against a nine minute grid is not +`small` is the default because iterating against a six minute grid is not practical. It reproduces the patterns the large one shows, at `level="node"`: | num_causal | small uniform | small rare | ratio | large uniform | large rare | ratio | |---|---|---|---|---|---|---| -| 1 | 0.016s | 0.011s | 1.4 | 0.058s | 0.051s | 1.1 | -| 100 | 0.036s | 0.013s | 2.8 | 0.152s | 0.056s | 2.7 | -| 1,000 | 0.189s | 0.015s | 12.3 | 0.593s | 0.065s | 9.1 | -| 10,000 | 1.596s | 0.035s | 45.7 | 5.171s | 0.108s | 48 | -| 100,000 | — | — | | 48.310s | 0.476s | 101 | +| 1 | 0.013s | 0.013s | 1.0 | 0.062s | 0.061s | 1.0 | +| 100 | 0.024s | 0.014s | 1.8 | 0.118s | 0.067s | 1.8 | +| 1,000 | 0.125s | 0.017s | 7.4 | 0.443s | 0.072s | 6.1 | +| 10,000 | 1.088s | 0.023s | 47 | 3.848s | 0.091s | 42 | +| 100,000 | — | — | | 32.589s | 0.216s | 151 | Both show a flat per call floor, a uniform µs/site that falls to a plateau from 1,000 causal sites upwards, a rare µs/site that is still falling at the top of the grid, a uniform to rare ratio that grows with the number of causal sites, -and parity between the three levels. The uniform plateau is 160µs/site against -483µs/site, a factor of 3.0 on a node count ratio of 3.4. +and parity between the three levels. The uniform plateau is 109µs/site against +326µs/site, a factor of 3.0 on a node count ratio of 3.4. What makes the small preset a fair substitute is not the timings but the distribution underneath them: the fraction of the nodes that a causal site's @@ -87,6 +87,44 @@ was tried and rejected: it takes `sites/nodes` from 0.68 to about 2.7 against the large preset's 1.08, and the per-call setup floor grows with the number of sites, so the low end of the curve stops being comparable. +### The two implementations + +There are two, and `--threshold` chooses between them per causal site. Below +the threshold a site goes through the push down of the ARG, which sweeps the +nodes once from the past to the present carrying the effect of each causal +mutation to its carriers. At or above it the site goes through a descent of the +trees, which builds each tree in turn and walks down from each causal mutation. +`0` sends everything to the descent and `inf` everything to the push down. + +Both cost the nodes that carry a causal allele. The descent is cheaper per +node, around 20ns against 28ns, because a child is the next link of a list +rather than a search of the out edges of a node, and because it holds nothing +per node while it works. Against that it has to build the trees, one insertion +and one removal per edge, which the push down never does. + +That fixed cost is the whole of the argument for a threshold, and the +measurement says it is not enough of one. Comparing thresholds over the whole +grid, on both presets, at all three levels, for one and three traits, and for +causal sites drawn both uniformly and from the rare ones, a threshold of zero +was fastest or tied in every cell: + +| | uniform 1,000 | uniform 10,000 | rare 1,000 | rare 10,000 | +|---|---|---|---|---| +| small, all descent | 122ms | 1058ms | 15ms | 23ms | +| small, all push down | 196ms | 1619ms | 16ms | 35ms | +| large, all descent | 427ms | 3709ms | 71ms | 90ms | +| large, all push down | 592ms | 4917ms | 68ms | 108ms | + +The push down leads only where there are few causal sites and they are rare, so +the tree pass is not amortised, and there it leads by a millisecond or two on a +call that takes a few. With three traits the gap widens the other way, to 2.27x +on rare causal sites, because the push down sweeps once per trait where the +descent takes all of them in one pass. + +So the default threshold is zero and the push down is not used. It is kept, and +`--threshold inf` still runs it, only so that the two can go on being compared; +the intention is to end with the descent alone. + ### What else it measures Wall time on its own does not say why a configuration is slow. Four optional @@ -94,54 +132,54 @@ modes say more. `--phases` is the expensive one, roughly doubling the run because it times the same work again a piece at a time; the other three add seconds. -`--phases` times `_check_trait_df`, `_GeneticValue.__init__`, the -`_push_down_arg` kernel and the output dataframe separately. The end to end +`--phases` times `_check_trait_df`, `_GeneticValue.__init__`, the kernel and +the output dataframe separately. The end to end number stays as the headline. This is what tells an algorithmic win from a setup win: setup barely grows with the causal sites, going from 8ms to 17ms across the whole `small` grid and sitting near 100ms on `large`, so at one causal site the public call is measuring almost nothing else, while at 10,000 the kernel is 98% of it. Most of setup is `tskit.jit.numba.jitwrap`, which runs three Python-speed `max(map(len, ...))` passes over the site and mutation -tables; `_root_runs` is most of the rest when it fires. The dataframe is about -1ms and is not worth thinking about. +tables. The dataframe is about 1ms and is not worth thinking about. -`--counters` runs a counting-only copy of the sweep and reports the work it -does. perf cannot attribute time to source lines inside a numba kernel here -(see below), so counting what the kernel does and dividing is the way to say -where the time goes. On `small`: +`--counters` reports the work the descent does. perf cannot attribute time to +source lines inside a numba kernel here (see below), so counting what the +kernel does and dividing is the way to say where the time goes. The kernel +returns the count itself rather than there being a second copy of the loop to +keep in step with the first. On `small`: ``` -selection num_causal seeds edge_scans edge_hits appends reached scans/seed hit rate reached -uniform 1 1 35,915 33,818 16,909 16,909 35915.0 94.2% 26.7% -uniform 10000 10,107 56,556,279 53,117,294 26,558,647 32,966 5595.8 93.9% 52.1% -rare 1 1 10 10 5 5 10.0 100.0% 0.0% -rare 10000 10,028 318,619 302,898 151,449 29,775 31.8 95.1% 47.0% +selection num_causal rows visits visits/row of num_nodes +uniform 1 1 33,819 33819.0 53.44% +uniform 10000 10,000 52,900,059 5290.0 8.36% +rare 1 1 11 11.0 0.02% +rare 10000 10,000 131,744 13.2 0.02% ``` -`edge_scans` is the trip count of the kernel's innermost loop and is what the -run time is proportional to, so `ns/scan` in the summary table is the constant -an optimisation has to move. Three things fall out of the table above: - -- The scan of a node's out edges is not where the waste is: 94% of trips find - an edge that spans the seed's causal site. -- Uniform selection costs 178 times as many scans as rare at 10,000 causal - sites but only 46 times the time, because rare is more expensive per scan: - 63ns against 28ns for the kernel alone. Following a few carriers is random - access; a dense pass over half the nodes is sequential. Compare the `kernel` - rows rather than `genetic_value`, or the setup floor swamps the rare ones. -- The kernel has an O(num_nodes) floor. It builds a pending entry for every - node before it starts and visits every node whether or not anything reached - it, so at one rare causal site — 10 scans — the kernel still takes 2ms. +`visits` is the number of nodes the descent reached, which is what the run time +is proportional to, so `ns/visit` in the summary table is the constant an +optimisation has to move. Two things fall out of the table: + +- `visits/row` as a fraction of the nodes is the carrier fraction that + `--structure` measures independently, and the two agree: 8.4% for uniformly + drawn causal sites against a measured mean of 7.2%, and 0.02% for rare ones + against a measured median of 0.4%. A single uniformly drawn site reaching + 53% is one draw from a distribution with a long tail. +- `ns/visit` is around 20ns once there are enough causal sites to amortise the + tree pass, and hundreds of nanoseconds below that. The pass is 2.2ms on + `small` and 12.4ms on `large`, so a rare trait of a hundred causal sites is + paying for a tree sequence it barely touches. That is the one regime the + push down still wins, by a millisecond or two. `--structure` reports the shape of the tree sequence and the carrier fraction distribution described above. -`--memory` reports the peak resident set size of each call, which is how the -typed list rewrite was justified. VmHWM never falls, so it is reset before each -call by writing to `/proc/self/clear_refs`; on a kernel without that the column -reads `unavailable`. This is the one thing the small preset is a poor substitute -for: 10,000 uniform causal sites peak at 0.02GB over the baseline there, against -the gigabytes the large preset reaches at 100,000. +`--memory` reports the peak resident set size of each call. VmHWM never falls, +so it is reset before each call by writing to `/proc/self/clear_refs`; on a +kernel without that the column reads `unavailable`. It mattered more when the +push down was the default, whose working set grew with the causal sites and +reached gigabytes; the descent holds a fixed handful of arrays the length of +the nodes however many causal sites there are. ### Output @@ -168,6 +206,10 @@ uv run --group test benchmarks/profile_genetic_value.py --mode kernel the setup, and it shows `jitwrap` and its `builtins.max` rows plainly. The kernel appears in it as one opaque dispatcher frame. +Note that `--mode kernel` runs the push down, not the descent, so its numbers +describe the implementation that `--threshold inf` selects rather than the +default one. The perf figures below were taken that way. + `--mode kernel` sets one cell up and runs only the kernel, so that a sampling profile is not swamped by simulating effect sizes for the whole site pool. Run on its own it prints the perf commands to copy; `--run` is what those commands @@ -201,14 +243,14 @@ finds the JIT mappings by itself, so the recipe does not set it. ## Gotchas -- `_GeneticValue` takes a `_root_runs` branch whenever a drawn causal allele is - the ancestral state of its site. It is a Python loop over every tree, costing - 7ms on `small` and 38ms on `large`, so on `large` it is over a third of the - setup. It is not an edge case: drawing 10,000 sites uniformly makes it near - certain that one of them qualifies, and it fires at 10,000 and 100,000 causal - sites on both presets. At the low end of the grid it usually does not, so it - appears part way up the curve and looks like a step in the setup cost. The - benchmark prints a line when a cell takes it. +- `_GeneticValue` takes a `_root_runs` branch when a drawn causal allele is the + ancestral state of its site *and* that row goes to the push down. It is a + Python loop over every tree, costing 7ms on `small` and 38ms on `large`, so on + `large` it was over a third of the setup. It is not an edge case: drawing + 10,000 sites uniformly makes it near certain that one of them qualifies. At + the default threshold those rows go to the descent, which has the roots to + hand, and the branch never runs; `--threshold inf` brings it back. The + benchmark prints a line when a cell actually takes it. - Never compare replicate 0 against replicates 1 and up. `mutations_inherited_state` and its neighbours are built lazily and cached on the tree sequence, so the first call on a given tree sequence pays for all of them. The warm up call diff --git a/benchmarks/baseline_small.csv b/benchmarks/baseline_small.csv index 483fac3..2c36907 100644 --- a/benchmarks/baseline_small.csv +++ b/benchmarks/baseline_small.csv @@ -1,85 +1,85 @@ -phase,num_causal,selection,level,replicate,seconds,num_samples,num_individuals,num_nodes,num_edges,num_trees,num_sites -sim_trait,1,uniform,,0,0.0024821249999149586,30000,15000,63287,75586,4108,42794 -sim_trait,1,uniform,,1,0.0022028849998605438,30000,15000,63287,75586,4108,42794 -sim_trait,1,uniform,,2,0.0021113359998707892,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,individual,0,0.017362073000185774,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,individual,1,0.016651905999424343,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,individual,2,0.016448290999505844,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,node,0,0.01680002900047839,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,node,1,0.016700003000551078,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,node,2,0.016296245999910752,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,edge,0,0.01801782100028504,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,edge,1,0.01743433499996172,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,edge,2,0.01691987899994274,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,individual,0,0.012258568000106607,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,individual,1,0.01165933300035249,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,individual,2,0.011369154000021808,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,node,0,0.011771557999963989,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,node,1,0.011883568000484956,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,node,2,0.011892135999914899,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,edge,0,0.012155798999629042,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,edge,1,0.012474035000195727,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,edge,2,0.01207672599957732,30000,15000,63287,75586,4108,42794 -sim_trait,100,uniform,,0,0.007151913000598142,30000,15000,63287,75586,4108,42794 -sim_trait,100,uniform,,1,0.005853265000041574,30000,15000,63287,75586,4108,42794 -sim_trait,100,uniform,,2,0.005971398000838235,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,individual,0,0.03708483099944715,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,individual,1,0.03696376599964424,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,individual,2,0.03942140500021196,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,node,0,0.0380683780003892,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,node,1,0.039289737999752106,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,node,2,0.03719520599952375,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,edge,0,0.0392425390000426,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,edge,1,0.03797104800014495,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,edge,2,0.04035435600053461,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,individual,0,0.01287275499998941,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,individual,1,0.012763894000272558,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,individual,2,0.01330246699944837,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,node,0,0.01320572200074821,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,node,1,0.012976414999684494,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,node,2,0.013185578000047826,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,edge,0,0.013328967999768793,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,edge,1,0.013860037999620545,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,edge,2,0.0130741749999288,30000,15000,63287,75586,4108,42794 -sim_trait,1000,uniform,,0,0.0281717200005005,30000,15000,63287,75586,4108,42794 -sim_trait,1000,uniform,,1,0.028771799999958603,30000,15000,63287,75586,4108,42794 -sim_trait,1000,uniform,,2,0.03157023900075728,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,individual,0,0.18984438400002546,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,individual,1,0.21684775500034448,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,individual,2,0.2687714949997826,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,node,0,0.226016618000358,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,node,1,0.20893606600020576,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,node,2,0.1926977330003865,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,edge,0,0.19492765799986955,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,edge,1,0.2099747740003295,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,edge,2,0.20740931999989698,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,individual,0,0.017229248999683477,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,individual,1,0.018185036999966542,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,individual,2,0.018869548999646213,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,node,0,0.018393275000562426,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,node,1,0.016106918999867048,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,node,2,0.015783592999468965,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,edge,0,0.015947148999657657,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,edge,1,0.016200306999962777,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,edge,2,0.015847360999941884,30000,15000,63287,75586,4108,42794 -sim_trait,10000,uniform,,0,0.24855346700041991,30000,15000,63287,75586,4108,42794 -sim_trait,10000,uniform,,1,0.25112143699971057,30000,15000,63287,75586,4108,42794 -sim_trait,10000,uniform,,2,0.2850705849996302,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,individual,0,1.704894198999682,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,individual,1,1.676236326000435,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,individual,2,1.6262085170001228,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,node,0,1.7224619290000192,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,node,1,1.6860980929996003,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,node,2,1.7367036840005312,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,edge,0,1.675767371000802,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,edge,1,1.613268662000337,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,edge,2,1.7291035789994567,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,individual,0,0.04531490799945459,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,individual,1,0.037717979999797535,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,individual,2,0.03916538799967384,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,node,0,0.044803013000091596,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,node,1,0.05839010900035646,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,node,2,0.05877978399985295,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,edge,0,0.03882118100045773,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,edge,1,0.03812910400029068,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,edge,2,0.03929764599979535,30000,15000,63287,75586,4108,42794 +phase,num_causal,selection,level,replicate,seconds,num_samples,num_individuals,num_nodes,num_edges,num_trees,num_sites,threshold +sim_trait,1,uniform,,0,0.0024810039994918043,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,1,uniform,,1,0.0021365499997045845,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,1,uniform,,2,0.002071789000183344,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,uniform,individual,0,0.013280454999403446,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,uniform,individual,1,0.012969833000170183,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,uniform,individual,2,0.01323338799920748,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,uniform,node,0,0.013460804999340326,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,uniform,node,1,0.01379560399982438,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,uniform,node,2,0.01315673300086928,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,uniform,edge,0,0.013742764998823986,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,uniform,edge,1,0.013944136999270995,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,uniform,edge,2,0.016557847999138176,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,rare,individual,0,0.014388511999641196,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,rare,individual,1,0.012964843999725417,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,rare,individual,2,0.012766195999574848,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,rare,node,0,0.013027275999775156,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,rare,node,1,0.012664592999499291,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,rare,node,2,0.013753575000009732,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,rare,edge,0,0.014797145999182248,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,rare,edge,1,0.01879820199974347,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1,rare,edge,2,0.014308412999525899,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,100,uniform,,0,0.006644762999712839,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,100,uniform,,1,0.00603657600004226,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,100,uniform,,2,0.0061505550002038945,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,uniform,individual,0,0.023435578999851714,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,uniform,individual,1,0.02330094599892618,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,uniform,individual,2,0.023336322999966796,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,uniform,node,0,0.02396749099898443,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,uniform,node,1,0.024603556001238758,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,uniform,node,2,0.02379652099989471,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,uniform,edge,0,0.024302648000229965,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,uniform,edge,1,0.02512614199986274,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,uniform,edge,2,0.024602867000794504,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,rare,individual,0,0.014113458999418071,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,rare,individual,1,0.01350924999860581,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,rare,individual,2,0.014417320000575273,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,rare,node,0,0.014358804000949021,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,rare,node,1,0.013511810000636615,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,rare,node,2,0.013791530000162311,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,rare,edge,0,0.016899227999601862,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,rare,edge,1,0.01515005100009148,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,100,rare,edge,2,0.013994134000313352,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,1000,uniform,,0,0.028378383998642676,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,1000,uniform,,1,0.02964526600044337,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,1000,uniform,,2,0.03017084299972339,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,uniform,individual,0,0.1165051690004475,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,uniform,individual,1,0.12006181399920024,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,uniform,individual,2,0.12192240799959109,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,uniform,node,0,0.1254568719996314,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,uniform,node,1,0.12637870099933934,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,uniform,node,2,0.12458551100098703,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,uniform,edge,0,0.12738228200032609,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,uniform,edge,1,0.12425055700077792,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,uniform,edge,2,0.12265691699940362,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,rare,individual,0,0.015345580999564845,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,rare,individual,1,0.015087319001395372,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,rare,individual,2,0.014869897000608034,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,rare,node,0,0.018529604998548166,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,rare,node,1,0.016907263001485262,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,rare,node,2,0.017068429000573815,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,rare,edge,0,0.015638313998351805,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,rare,edge,1,0.016809493999971892,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,1000,rare,edge,2,0.015862450998611166,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,10000,uniform,,0,0.24336689300071157,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,10000,uniform,,1,0.24226529100087646,30000,15000,63287,75586,4108,42794,0.0 +sim_trait,10000,uniform,,2,0.2436596239986102,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,uniform,individual,0,1.0077231790000951,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,uniform,individual,1,1.0158602640003664,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,uniform,individual,2,1.011626936000539,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,uniform,node,0,1.0877643929998158,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,uniform,node,1,1.0880109169993375,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,uniform,node,2,1.0914093689989386,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,uniform,edge,0,1.0646407929998531,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,uniform,edge,1,1.0721369829989271,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,uniform,edge,2,1.0633098569996946,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,rare,individual,0,0.024103609001031145,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,rare,individual,1,0.02428249100012181,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,rare,individual,2,0.024011842999243527,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,rare,node,0,0.02510500999960641,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,rare,node,1,0.023118164001061814,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,rare,node,2,0.023174789999757195,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,rare,edge,0,0.023430856001141365,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,rare,edge,1,0.025804556000366574,30000,15000,63287,75586,4108,42794,0.0 +genetic_value,10000,rare,edge,2,0.02633049299947743,30000,15000,63287,75586,4108,42794,0.0 diff --git a/benchmarks/baseline_small_counters.csv b/benchmarks/baseline_small_counters.csv index d730e93..2bdefaf 100644 --- a/benchmarks/baseline_small_counters.csv +++ b/benchmarks/baseline_small_counters.csv @@ -1,9 +1,9 @@ -num_causal,selection,seeds,edge_scans,edge_hits,appends,reached,num_nodes,num_edges -1,uniform,1,35915,33818,16909,16909,63287,75586 -1,rare,1,10,10,5,5,63287,75586 -100,uniform,101,558839,524756,262378,31784,63287,75586 -100,rare,100,1507,1454,727,706,63287,75586 -1000,uniform,1008,5856520,5501688,2750844,32611,63287,75586 -1000,rare,1002,11559,11228,5614,4671,63287,75586 -10000,uniform,10107,56556279,53117294,26558647,32966,63287,75586 -10000,rare,10028,318619,302898,151449,29775,63287,75586 +num_causal,selection,rows,visits,num_nodes,num_edges +1,uniform,1,33819,63287,75586 +1,rare,1,11,63287,75586 +100,uniform,100,514067,63287,75586 +100,rare,100,1554,63287,75586 +1000,uniform,1000,5501454,63287,75586 +1000,rare,1000,12230,63287,75586 +10000,uniform,10000,52900059,63287,75586 +10000,rare,10000,131744,63287,75586 diff --git a/benchmarks/benchmark_genetic_value.py b/benchmarks/benchmark_genetic_value.py index da6920b..cfbbf28 100644 --- a/benchmarks/benchmark_genetic_value.py +++ b/benchmarks/benchmark_genetic_value.py @@ -20,16 +20,13 @@ import time import msprime -import numba import numpy as np import pandas as pd import tskit -from numba.core import types -from numba.typed import List import tstrait from tstrait import jit -from tstrait.genetic_value import _check_trait_df, _GeneticValue +from tstrait.genetic_value import _COMMON_THRESHOLD, _check_trait_df, _GeneticValue LEVELS = ["individual", "node", "edge"] SELECTIONS = ["uniform", "rare"] @@ -51,100 +48,6 @@ }, } -# The counting kernel mirrors tstrait.jit._push_down_arg, so it holds what a -# node holds there: the indexes of the seeds whose effect has reached it. -_SEED_LIST = types.ListType(types.int32) - -COUNTERS = ["seeds", "edge_scans", "edge_hits", "appends", "reached"] - - -@numba.njit -def _count_push_down_arg( - child_index, - edges_child, - edges_site_start, - edges_site_stop, - nodes_by_time, - seed_node, - seed_site, -): - """ - Count the work that ``tstrait.jit._push_down_arg`` does, without doing any - of it. - - perf cannot attribute time to source lines inside a numba kernel here, so - the way to say where the time goes is to count the things the kernel does - and divide. This is a copy of the sweep with the output writes and the - weight lookups taken out and a counter put in their place, which is the - only reason it may diverge from the kernel it mirrors: keep them in step. - - Returns the counters named in ``COUNTERS``: - - ``seeds`` the causal mutations, plus the roots seeded when the causal - allele is the ancestral state - ``edge_scans`` trips of the innermost loop, i.e. the sum over every - (swept node, seed held there) pair of the node's out degree - ``edge_hits`` those trips where the edge spans the seed's causal site, so - edge_scans - edge_hits is the scan that was wasted - ``appends`` seeds pushed onto a node's list - ``reached`` nodes that held a list, so reached / num_nodes is the - fraction of the tree sequence the sweep touched, and, since - a list is made for a node the first time anything reaches - it, also the number of typed lists allocated - """ - num_nodes = len(child_index) - empty = List.empty_list(types.int32) - pending = List.empty_list(_SEED_LIST) - for _ in range(num_nodes): - pending.append(empty) - reached = np.zeros(num_nodes, dtype=np.bool_) - - edge_scans = 0 - edge_hits = 0 - appends = 0 - - for j in range(len(seed_node)): - u = seed_node[j] - if child_index[u, 0] < 0: - continue - if not reached[u]: - pending[u] = List.empty_list(types.int32) - reached[u] = True - pending[u].append(np.int32(j)) - appends += 1 - - for i in range(len(nodes_by_time)): - parent = nodes_by_time[i] - if not reached[parent]: - continue - items = pending[parent] - edge_start = child_index[parent, 0] - edge_stop = child_index[parent, 1] - for k in range(len(items)): - item = items[k] - site = seed_site[item] - edge_scans += edge_stop - edge_start - for e in range(edge_start, edge_stop): - if edges_site_start[e] <= site and site < edges_site_stop[e]: - edge_hits += 1 - child = edges_child[e] - if child_index[child, 0] < 0: - continue - if not reached[child]: - pending[child] = List.empty_list(types.int32) - reached[child] = True - pending[child].append(np.int32(item)) - appends += 1 - pending[parent] = empty - - return ( - len(seed_node), - edge_scans, - edge_hits, - appends, - int(np.sum(reached)), - ) - def cached_simulation(args): """ @@ -269,10 +172,12 @@ def warm_up(ts, model, levels, counters): before = time.perf_counter() trait_df = tstrait.sim_trait(ts, model=model, num_causal=1, random_seed=1) for level in levels: - tstrait.genetic_value(ts, trait_df, level=level) + # Both implementations, since the grid may use either. + for threshold in (0.0, np.inf): + tstrait.genetic_value(ts, trait_df, level=level, _threshold=threshold) if counters: # A kernel of its own to compile, so only pay for it when it is used. - count_work(ts, _check_trait_df(ts, trait_df)) + count_work(ts, _check_trait_df(ts, trait_df), 0.0) print(f"Warm up (includes numba compilation): {time.perf_counter() - before:.1f}s") @@ -328,39 +233,59 @@ def _build_frame(num_trait, size, level, values): ) -def count_work(ts, trait_df): +COUNTERS = ["rows", "visits"] + + +def count_work(ts, trait_df, threshold): """ - Return the work counters for a trait, as a dict keyed by COUNTERS. + Return the work the descent does, as a dict keyed by COUNTERS. + + ``visits`` is the number of nodes the descent reached, which is what its + run time is proportional to, so seconds divided by it is the per node + constant that an optimisation has to move. The kernel returns it, so there + is no second copy of the loop to keep in step with the first. - The counts do not depend on the level, because the same sweep serves all + The counts do not depend on the level, because the same descent serves all three and only the array the contributions land in differs. """ - genetic = _GeneticValue(ts, trait_df) - arguments = genetic._descent_arguments("node", 0, np.zeros(ts.num_nodes)) - counts = _count_push_down_arg( - arguments["child_index"], - arguments["edges_child"], - arguments["edges_site_start"], - arguments["edges_site_stop"], - arguments["nodes_by_time"], - arguments["seed_node"], - arguments["seed_site"], + genetic = _GeneticValue(ts, trait_df, threshold=threshold) + if not np.any(genetic.descent_rows): + return {"rows": 0, "visits": 0} + visits = jit._descend_trees( + genetic.numba_ts, + ts.edges_parent, + ts.edges_child, + genetic.row_site, + genetic.row_trait, + genetic.row_effect, + genetic.row_ancestral, + genetic.descent_rows, + genetic.pair_offset, + genetic.pair_node, + genetic.pair_carries, + ts.samples().astype(np.int32), + genetic._node_output("node"), + False, + np.zeros((genetic.num_trait, ts.num_nodes)), ) - return dict(zip(COUNTERS, counts)) + return {"rows": int(np.sum(genetic.descent_rows)), "visits": int(visits)} -def root_runs_fired(ts, trait_df): +def root_runs_fired(ts, trait_df, descent_rows): """ Whether this trait takes the _root_runs branch in _GeneticValue. That branch is a Python loop over every tree, so it costs O(num_trees) at - Python speed, and it fires only when a drawn causal allele happens to be - the ancestral state of its site. It therefore appears and vanishes with the - seed, which makes for confusing non-monotonic timings unless it is reported. + Python speed. It fires when a drawn causal allele happens to be the + ancestral state of its site and that row goes to the push down, which needs + the roots found for it; the descent has them to hand. It therefore appears + and vanishes with the seed and the threshold, which makes for confusing + non-monotonic timings unless it is reported. """ site_id = trait_df["site_id"].to_numpy() causal_allele = trait_df["causal_allele"].to_numpy() - return bool(np.any(ts.sites_ancestral_state[site_id] == causal_allele)) + ancestral = ts.sites_ancestral_state[site_id] == causal_allele + return bool(np.any(ancestral & ~descent_rows)) def carrier_fractions(ts, pool, sample_size, rng): @@ -444,18 +369,26 @@ def run_benchmark(ts, args): if trait_df is None: print(f" too few {selection} sites for num_causal={num_causal}") continue - if root_runs_fired(ts, trait_df): + checked = _check_trait_df(ts, trait_df) + descent = _GeneticValue(ts, checked, threshold=args.threshold).descent_rows + print( + f" {selection} num_causal={num_causal}: " + f"{descent.sum()} of {len(descent)} rows take the descent" + ) + if root_runs_fired(ts, checked, descent): print( f" {selection} num_causal={num_causal} takes the _root_runs " "branch, a Python loop over every tree" ) if args.counters: - counts[(num_causal, selection)] = count_work( - ts, _check_trait_df(ts, trait_df) - ) + counts[(num_causal, selection)] = count_work(ts, checked, args.threshold) for level in args.levels: call = functools.partial( - tstrait.genetic_value, ts, trait_df, level=level + tstrait.genetic_value, + ts, + trait_df, + level=level, + _threshold=args.threshold, ) if args.memory: _, memory[(num_causal, selection, level)] = peak_memory(call) @@ -513,11 +446,11 @@ def summarise(rows, counts, memory, completed, ts, args): best[key] = min(best.get(key, seconds), seconds) print(f"\n{describe(ts)}") - print(f"Minimum of {args.replicates} replicates\n") + print(f"Minimum of {args.replicates} replicates, threshold {args.threshold:g}\n") columns = f"{'phase':<14} {'selection':<10} {'level':<11} {'num_causal':>10} " columns += f"{'seconds':>10} {'us/site':>10}" if args.counters: - columns += f" {'ns/scan':>9}" + columns += f" {'ns/visit':>9}" print(columns) print("-" * len(columns)) phases = ["sim_trait", "genetic_value"] @@ -542,12 +475,12 @@ def summarise(rows, counts, memory, completed, ts, args): f"{seconds:>10.3f} {seconds / num_causal * 1e6:>10.1f}" ) if args.counters: - # Only the phases that run the sweep have a per trip cost. - scans = counts.get((num_causal, selection), {}).get("edge_scans") - sweeps = phase in ("genetic_value", "kernel") + # Only the phases that run the descent have a per node cost. + visits = counts.get((num_causal, selection), {}).get("visits") + descends = phase in ("genetic_value", "kernel") line += ( - f" {seconds / scans * 1e9:>9.2f}" - if sweeps and scans is not None + f" {seconds / visits * 1e9:>9.2f}" + if descends and visits else f" {'':>9}" ) print(line) @@ -555,23 +488,19 @@ def summarise(rows, counts, memory, completed, ts, args): if args.counters: print() header = f"{'selection':<10} {'num_causal':>10} " - header += " ".join(f"{name:>12}" for name in COUNTERS) - header += f" {'scans/seed':>11} {'hit rate':>9} {'reached':>8}" + header += " ".join(f"{name:>14}" for name in COUNTERS) + header += f" {'visits/row':>11} {'of num_nodes':>13}" print(header) print("-" * len(header)) - # The sweep visits every node whether or not anything reached it, and - # builds a pending entry for every node before it starts, so a low - # reached fraction is a kernel spending its time on the prologue. for selection in args.selections: for num_causal in completed: got = counts.get((num_causal, selection)) if got is None: continue line = f"{selection:<10} {num_causal:>10} " - line += " ".join(f"{got[name]:>12,}" for name in COUNTERS) - line += f" {got['edge_scans'] / max(got['seeds'], 1):>11.1f}" - line += f" {got['edge_hits'] / max(got['edge_scans'], 1) * 100:>8.1f}%" - line += f" {got['reached'] / ts.num_nodes * 100:>7.1f}%" + line += " ".join(f"{got[name]:>14,}" for name in COUNTERS) + line += f" {got['visits'] / max(got['rows'], 1):>11.1f}" + line += f" {got['visits'] / max(got['rows'], 1) / ts.num_nodes:>12.2%}" print(line) if len(memory) > 0: @@ -611,6 +540,7 @@ def write_csv(rows, ts, args): "num_edges", "num_trees", "num_sites", + "threshold", ] ) for row in rows: @@ -623,6 +553,7 @@ def write_csv(rows, ts, args): ts.num_edges, ts.num_trees, ts.num_sites, + args.threshold, ] ) print(f"\nWrote {args.output}") @@ -697,6 +628,16 @@ def parse_args(): "taken longer than this" ), ) + parser.add_argument( + "--threshold", + type=float, + default=_COMMON_THRESHOLD, + help=( + "Allele frequency at or above which a causal site goes through the " + "descent of the trees. 0 sends everything there and inf sends " + "everything to the push down of the ARG" + ), + ) parser.add_argument( "--phases", action="store_true", diff --git a/tests/test_genetic_value.py b/tests/test_genetic_value.py index 4a85157..5df1274 100644 --- a/tests/test_genetic_value.py +++ b/tests/test_genetic_value.py @@ -1315,15 +1315,17 @@ def test_split_is_a_split(self): assert np.any(~genetic.descent_rows) assert not np.all(genetic.descent_rows[genetic.seed_row]) - def test_no_allele_freq_is_all_push_down(self): - # A trait dataframe assembled by hand has no allele_freq, and then - # there is nothing to route on. + def test_no_allele_freq_needs_no_threshold(self): + # A trait dataframe assembled by hand has no allele_freq, so there is + # nothing for a positive threshold to compare against and every row + # takes the push down. The default threshold of zero needs no + # comparison, so those rows take the descent like any other. ts = self.simulated_ts() trait_df = _check_trait_df( ts, random_trait_df(ts, 1, seed=5).drop(columns=["allele_freq"]) ) - genetic = _GeneticValue(ts, trait_df, threshold=0.0) - assert not np.any(genetic.descent_rows) + assert not np.any(_GeneticValue(ts, trait_df, threshold=0.5).descent_rows) + assert np.all(_GeneticValue(ts, trait_df, threshold=0.0).descent_rows) @pytest.mark.parametrize("derived_state", ["A", "T"]) def test_mutation_on_root(self, derived_state): diff --git a/tests/test_jit.py b/tests/test_jit.py index 08af5be..7d82282 100644 --- a/tests/test_jit.py +++ b/tests/test_jit.py @@ -90,7 +90,9 @@ def f( "causal_allele": [causal_allele], } ) - genetic = _GeneticValue(ts, trait_df) + # These test the push down kernel, so keep every row on it rather + # than letting the default threshold send them to the descent. + genetic = _GeneticValue(ts, trait_df, threshold=np.inf) output = np.zeros(genetic._output_size(level)) return func(**genetic._descent_arguments(level, 0, output)) diff --git a/tstrait/genetic_value.py b/tstrait/genetic_value.py index f349a4e..040af0e 100644 --- a/tstrait/genetic_value.py +++ b/tstrait/genetic_value.py @@ -11,12 +11,20 @@ ALLELE_FREQ = "allele_freq" # Causal sites at or above this allele frequency go through the descent of the -# trees, and the rest through the push down of the ARG. The descent is faster -# almost everywhere, and measurably so from about a thousandth: below that the -# tree pass it starts with is not amortised and the push down, which never -# builds a tree, wins. Private while the two implementations are being -# compared, since the intention is to end with one of them. -_COMMON_THRESHOLD = 0.001 +# trees, and the rest through the push down of the ARG. Zero, because the +# descent measured faster than the push down at every point of the benchmark +# grid on both tree sequences, at every level, for one and three traits, and +# for causal sites drawn both uniformly and from the rare ones: 1.33 to 1.70 +# times on uniformly drawn sites and up to 2.27 times on rare ones with three +# traits, where the push down runs its sweep once per trait and the descent +# takes all of them in one pass. The few cells the push down led were a +# fraction of a millisecond apart. +# +# A threshold of zero needs no allele frequency to compare against, so it +# applies to a trait dataframe assembled by hand as well. Private, and kept +# only so that the two can still be compared, since the intention is to end +# with the descent alone. +_COMMON_THRESHOLD = 0.0 def _row_mutations(ts, trait_df): @@ -199,8 +207,13 @@ def __init__(self, ts, trait_df, threshold=_COMMON_THRESHOLD): # Rows go to the descent of the trees when their causal allele is # common enough for the tree pass to pay for itself, and when it is the # ancestral state, where the descent has the roots to hand and the push - # down would need them found for it. - if ALLELE_FREQ in self.trait_df.columns: + # down would need them found for it. Every frequency is at or above a + # threshold of zero, so there is nothing to look up in that case, which + # is what lets the default apply to a trait dataframe that has no + # allele frequency in it. + if threshold <= 0: + self.descent_rows = np.ones(len(self.trait_df), dtype=bool) + elif ALLELE_FREQ in self.trait_df.columns: self.descent_rows = self.trait_df[ALLELE_FREQ].to_numpy() >= threshold self.descent_rows |= self.row_ancestral else: diff --git a/tstrait/jit.py b/tstrait/jit.py index a46b831..c491a3e 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -260,6 +260,10 @@ def _descend_trees( is instead the edge above the node in the current tree, which is where the contribution arrived from. All of the traits share the pass, since a row knows the trait it belongs to. + + ``output`` is accumulated into, and the count of nodes visited is returned. + That count is what the run time is proportional to, so the benchmark + divides by it rather than keeping a copy of this loop that only counts. """ num_nodes = numba_ts.num_nodes tree = tree_state(num_nodes) @@ -273,6 +277,7 @@ def _descend_trees( # a seed stops at the mutations of the row, so the tree bounds the stack. stack = np.empty(num_nodes, dtype=np.int32) + visits = 0 row = 0 num_rows = len(row_site) tree_index = numba_ts.tree_index() @@ -318,6 +323,7 @@ def _descend_trees( while top > 0: top -= 1 node = stack[top] + visits += 1 slot = tree.node_edge[node] if edge_level else node_output[node] if slot >= 0: output[trait, slot] += weight @@ -328,4 +334,4 @@ def _descend_trees( top += 1 child = tree.right_sib[child] row += 1 - return output + return visits From 8ffd22f98defc9eaf6c2cfe610c3fac7a3b94f30 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Thu, 3 Sep 2026 12:30:36 +0100 Subject: [PATCH 10/19] Record what the descent does to the working set The push down holds the seeds still in flight, which grows with the number of causal sites; the descent holds a fixed handful of arrays the length of the nodes whatever the trait looks like. At 100,000 uniformly drawn causal sites on 100,000 samples, that is 0.01GB over the baseline against 0.89GB, and a peak of 0.42GB against 1.30GB. --- benchmarks/README.md | 15 +++++++++++---- 1 file changed, 11 insertions(+), 4 deletions(-) diff --git a/benchmarks/README.md b/benchmarks/README.md index 973e48a..b073f61 100644 --- a/benchmarks/README.md +++ b/benchmarks/README.md @@ -176,10 +176,17 @@ distribution described above. `--memory` reports the peak resident set size of each call. VmHWM never falls, so it is reset before each call by writing to `/proc/self/clear_refs`; on a -kernel without that the column reads `unavailable`. It mattered more when the -push down was the default, whose working set grew with the causal sites and -reached gigabytes; the descent holds a fixed handful of arrays the length of -the nodes however many causal sites there are. +kernel without that the column reads `unavailable`. + +It is the clearest difference between the two implementations. The push down +holds the seeds still in flight, which grows with the causal sites; the descent +holds a fixed handful of arrays the length of the nodes however many causal +sites there are. At 100,000 uniformly drawn causal sites on `large`: + +| | peak RSS | over baseline | +|---|---|---| +| descent | 0.42 GB | 0.01 GB | +| push down | 1.30 GB | 0.89 GB | ### Output From 9b123468fccc9b4db7ad397df7af55f5dbed7ed7 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Thu, 3 Sep 2026 12:53:38 +0100 Subject: [PATCH 11/19] Delete the push down, which lost to the descent everywhere Two implementations of the genetic value computation have been carried side by side so that they could be compared. The comparison is done, so delete the one that lost, along with the threshold that chose between them and everything built only to feed it. Measured over the whole benchmark grid, on both tree sequences, at each level, for one and three traits, and for causal sites drawn both uniformly and from the rare ones, the push down was slower everywhere the cost mattered, by 1.33 to 1.70 times on uniformly drawn causal sites and up to 2.27 times on rare ones with three traits, where it swept once per trait against a single pass for all of them: uniform rare 1,000 10,000 1,000 10,000 small descent 122ms 1058ms 15ms 23ms push down 196ms 1619ms 16ms 35ms large descent 427ms 3709ms 71ms 90ms push down 592ms 4917ms 68ms 108ms It was ahead in ten of the thirty six cells, every one of them rare causal sites at a thousand or fewer, where the tree building the descent starts with has too little to amortise it against. In none of them was it ahead by more than a few milliseconds on a call taking a few. The note left in jit.py says as much, since a reader can check it against the same benchmark. Deleting it takes the setup with it, which was built on every call whether the push down ran or not: an argsort over the nodes, two searchsorted over the edge table, the tskit child index, and the seed arrays. _GeneticValue setup on 100,000 samples goes from 98ms to 42ms, almost all of what is left being tskit's jitwrap, and every cell of the small preset grid gets between 1.01 and 1.41 times faster with no change to the kernel at all. _root_runs goes too. It was a Python loop over every tree, finding the roots for causal sites whose ancestral state is the causal allele, and only the push down needed it: the descent is already standing in the tree and asks _tree_roots. The threshold argument, the optional allele_freq column that it read, and the row selection the descent kernel took all go with it, so _check_trait_df is back to the four required columns and sim_trait's allele_freq is once again only an output. The thirty five topology tests carry over unchanged by repointing the one fixture they share; the jit and nojit parametrisation survives, though nojit now interprets only the descent's own loop, since the tree building kernels it calls are compiled. Those are covered instead by TestTreeState against tskit.Tree. The tests that existed to compare the two implementations go, and the reference comparison against the tree by tree oracle stops looping over thresholds. Verified against that same tree by tree implementation, recovered from 31fec9a, on both tree sequences at all three levels for one and three traits. --- CHANGELOG.md | 14 +- benchmarks/README.md | 105 +++------- benchmarks/baseline_small.csv | 170 ++++++++-------- benchmarks/benchmark_genetic_value.py | 97 ++------- benchmarks/profile_genetic_value.py | 29 ++- tests/test_genetic_value.py | 102 ++-------- tests/test_jit.py | 16 +- tstrait/genetic_value.py | 283 +++++--------------------- tstrait/jit.py | 126 ++---------- 9 files changed, 245 insertions(+), 697 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 0e119b3..dd215bd 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,12 +13,14 @@ In development ### Performance -- `genetic_value` pushes the effect of each causal mutation down the ARG - instead of working through the trees one at a time, so its cost is the - number of nodes carrying a causal allele rather than the size of the tree - sequence. On 100,000 samples, a trait with 100,000 rare causal sites is - around 70 times faster; traits whose causal sites are mostly common - variants are around 1.5 times slower. +- `genetic_value` descends from the mutations of each causal site instead of + making a pass over every node for each of them, so its cost is the number of + nodes carrying a causal allele rather than the number of causal sites times + the size of the tree sequence. On 100,000 samples, a trait with 100,000 rare + causal sites is around 180 times faster, and one whose causal sites are drawn + uniformly, and so are mostly common variants, around 1.2 times faster. Every + trait of a tree sequence is computed in one pass over the trees rather than + one pass each. ### Breaking changes diff --git a/benchmarks/README.md b/benchmarks/README.md index b073f61..28105bf 100644 --- a/benchmarks/README.md +++ b/benchmarks/README.md @@ -12,8 +12,10 @@ a trait. uv run --group test benchmarks/benchmark_genetic_value.py ``` -The cost of the ARG sweep is the number of nodes that carry a causal allele, so -the allele frequency of the causal sites matters more than how many there are. +`genetic_value` builds each tree in turn and descends from the mutations of the +causal sites it carries, at around 20ns a node once the tree building is +amortised, so its cost is the number of nodes that carry a causal allele and the +allele frequency of the causal sites matters more than how many there are. `--selections` therefore draws them two ways: `uniform` over all sites, which the common variants in the tail of the frequency spectrum dominate, and `rare`, restricted to sites below `--rare-threshold`. The two differ by two orders of @@ -22,7 +24,7 @@ site" is meaningless without saying which. `sim_trait` is timed separately, because it has a per-site Python loop of its own that we do not want folded into the `genetic_value` numbers. The numba -kernels are compiled by a warm up call that is not timed. +kernel is compiled by a warm up call that is not timed. The mutation rate defaults to 1e-7, ten times the human rate, so that there are enough sites in a genome short enough to simulate quickly. The mutation rate @@ -71,7 +73,7 @@ and parity between the three levels. The uniform plateau is 109µs/site against What makes the small preset a fair substitute is not the timings but the distribution underneath them: the fraction of the nodes that a causal site's -effect reaches, which is what the sweep costs. `--structure` measures it. +effect reaches, which is what the descent costs. `--structure` measures it. | | mean carrier fraction | median | |---|---|---| @@ -87,44 +89,6 @@ was tried and rejected: it takes `sites/nodes` from 0.68 to about 2.7 against the large preset's 1.08, and the per-call setup floor grows with the number of sites, so the low end of the curve stops being comparable. -### The two implementations - -There are two, and `--threshold` chooses between them per causal site. Below -the threshold a site goes through the push down of the ARG, which sweeps the -nodes once from the past to the present carrying the effect of each causal -mutation to its carriers. At or above it the site goes through a descent of the -trees, which builds each tree in turn and walks down from each causal mutation. -`0` sends everything to the descent and `inf` everything to the push down. - -Both cost the nodes that carry a causal allele. The descent is cheaper per -node, around 20ns against 28ns, because a child is the next link of a list -rather than a search of the out edges of a node, and because it holds nothing -per node while it works. Against that it has to build the trees, one insertion -and one removal per edge, which the push down never does. - -That fixed cost is the whole of the argument for a threshold, and the -measurement says it is not enough of one. Comparing thresholds over the whole -grid, on both presets, at all three levels, for one and three traits, and for -causal sites drawn both uniformly and from the rare ones, a threshold of zero -was fastest or tied in every cell: - -| | uniform 1,000 | uniform 10,000 | rare 1,000 | rare 10,000 | -|---|---|---|---|---| -| small, all descent | 122ms | 1058ms | 15ms | 23ms | -| small, all push down | 196ms | 1619ms | 16ms | 35ms | -| large, all descent | 427ms | 3709ms | 71ms | 90ms | -| large, all push down | 592ms | 4917ms | 68ms | 108ms | - -The push down leads only where there are few causal sites and they are rare, so -the tree pass is not amortised, and there it leads by a millisecond or two on a -call that takes a few. With three traits the gap widens the other way, to 2.27x -on rare causal sites, because the push down sweeps once per trait where the -descent takes all of them in one pass. - -So the default threshold is zero and the push down is not used. It is kept, and -`--threshold inf` still runs it, only so that the two can go on being compared; -the intention is to end with the descent alone. - ### What else it measures Wall time on its own does not say why a configuration is slow. Four optional @@ -168,8 +132,7 @@ optimisation has to move. Two things fall out of the table: - `ns/visit` is around 20ns once there are enough causal sites to amortise the tree pass, and hundreds of nanoseconds below that. The pass is 2.2ms on `small` and 12.4ms on `large`, so a rare trait of a hundred causal sites is - paying for a tree sequence it barely touches. That is the one regime the - push down still wins, by a millisecond or two. + paying to build a tree sequence it barely touches. `--structure` reports the shape of the tree sequence and the carrier fraction distribution described above. @@ -178,15 +141,10 @@ distribution described above. so it is reset before each call by writing to `/proc/self/clear_refs`; on a kernel without that the column reads `unavailable`. -It is the clearest difference between the two implementations. The push down -holds the seeds still in flight, which grows with the causal sites; the descent -holds a fixed handful of arrays the length of the nodes however many causal -sites there are. At 100,000 uniformly drawn causal sites on `large`: - -| | peak RSS | over baseline | -|---|---|---| -| descent | 0.42 GB | 0.01 GB | -| push down | 1.30 GB | 0.89 GB | +The working set is a fixed handful of arrays the length of the nodes however +many causal sites there are: at 100,000 uniformly drawn causal sites on `large` +the peak is 0.42GB, of which 0.01GB is what the call added over the tree +sequence and the interpreter. ### Output @@ -194,8 +152,8 @@ Results are written to `_output/genetic_value.csv` in long format, one row per replicate, together with the dimensions of the tree sequence they were measured on; `--counters` writes a second file alongside it. `_output/` is gitignored. -`baseline_small.csv` and `baseline_small_counters.csv` are the `small` preset at -the tip of the ARG sweep work, for diffing against. The timings in the first are +`baseline_small.csv` and `baseline_small_counters.csv` are the `small` preset as +it stands, for diffing against. The timings in the first are specific to the machine they were taken on; the counts in the second are not, and are the part worth treating as a regression test. @@ -210,12 +168,9 @@ uv run --group test benchmarks/profile_genetic_value.py --mode kernel ``` `--mode python` is cProfile around the public call. It is the only way to see -the setup, and it shows `jitwrap` and its `builtins.max` rows plainly. The -kernel appears in it as one opaque dispatcher frame. - -Note that `--mode kernel` runs the push down, not the descent, so its numbers -describe the implementation that `--threshold inf` selects rather than the -default one. The perf figures below were taken that way. +the setup, and it shows `jitwrap` and its `builtins.max` rows plainly, which is +now nearly all of what setup costs. The kernel appears in it as one opaque +dispatcher frame. `--mode kernel` sets one cell up and runs only the kernel, so that a sampling profile is not swamped by simulating effect sizes for the whole site pool. Run @@ -228,36 +183,30 @@ Two things to know before reading a perf profile of this code. `PerfJITEventListener`, so nothing writes a `/tmp/perf-.map` or a jitdump and the kernel shows up as raw addresses under `[JIT]`. Use `--counters` for attribution inside the kernel. What perf does give is the split between the -kernel, the numba runtime, the interpreter and LLVM compilation, with the typed -list functions resolved by name — which is how to size the typed list overhead: +kernel, the interpreter, LLVM compilation and the numba runtime: ``` -64.42% [JIT] tid 4722 <- the kernel -12.18% python3.11 <- setup - 6.27% libc.so.6 - 6.02% libllvmlite.so <- compilation, not work - 4.93% _helperlib.cpython-311-...so <- numba_list_append, numba_list_resize +59.19% [JIT] tid 23771 <- the kernel +19.24% python3.11 <- setup +11.78% libllvmlite.so <- compilation, not work + 4.14% [kernel.kallsyms] + 1.75% libc.so.6 ``` +The numba runtime does not appear at all: the descent allocates nothing per +node, so there is no `_helperlib` row. + That is `--repeats 10`; the setup is a fixed few seconds, so raise `--repeats` until the `[JIT]` share stops moving. **Do not set `NUMBA_ENABLE_PROFILING=1`.** It is the documented way to profile numba and is wrong here: it would only help through the listener llvmlite does not have, and it defaults `NUMBA_DEBUGINFO` to 1, which changes the generated -code. The same kernel measured 2.62s a call with it and 1.61s without. perf -finds the JIT mappings by itself, so the recipe does not set it. +code and measured the kernel over 60% slower. perf finds the JIT mappings by +itself, so the recipe does not set it. ## Gotchas -- `_GeneticValue` takes a `_root_runs` branch when a drawn causal allele is the - ancestral state of its site *and* that row goes to the push down. It is a - Python loop over every tree, costing 7ms on `small` and 38ms on `large`, so on - `large` it was over a third of the setup. It is not an edge case: drawing - 10,000 sites uniformly makes it near certain that one of them qualifies. At - the default threshold those rows go to the descent, which has the roots to - hand, and the branch never runs; `--threshold inf` brings it back. The - benchmark prints a line when a cell actually takes it. - Never compare replicate 0 against replicates 1 and up. `mutations_inherited_state` and its neighbours are built lazily and cached on the tree sequence, so the first call on a given tree sequence pays for all of them. The warm up call diff --git a/benchmarks/baseline_small.csv b/benchmarks/baseline_small.csv index 2c36907..cdbba44 100644 --- a/benchmarks/baseline_small.csv +++ b/benchmarks/baseline_small.csv @@ -1,85 +1,85 @@ -phase,num_causal,selection,level,replicate,seconds,num_samples,num_individuals,num_nodes,num_edges,num_trees,num_sites,threshold -sim_trait,1,uniform,,0,0.0024810039994918043,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,1,uniform,,1,0.0021365499997045845,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,1,uniform,,2,0.002071789000183344,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,uniform,individual,0,0.013280454999403446,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,uniform,individual,1,0.012969833000170183,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,uniform,individual,2,0.01323338799920748,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,uniform,node,0,0.013460804999340326,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,uniform,node,1,0.01379560399982438,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,uniform,node,2,0.01315673300086928,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,uniform,edge,0,0.013742764998823986,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,uniform,edge,1,0.013944136999270995,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,uniform,edge,2,0.016557847999138176,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,rare,individual,0,0.014388511999641196,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,rare,individual,1,0.012964843999725417,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,rare,individual,2,0.012766195999574848,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,rare,node,0,0.013027275999775156,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,rare,node,1,0.012664592999499291,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,rare,node,2,0.013753575000009732,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,rare,edge,0,0.014797145999182248,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,rare,edge,1,0.01879820199974347,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1,rare,edge,2,0.014308412999525899,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,100,uniform,,0,0.006644762999712839,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,100,uniform,,1,0.00603657600004226,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,100,uniform,,2,0.0061505550002038945,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,uniform,individual,0,0.023435578999851714,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,uniform,individual,1,0.02330094599892618,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,uniform,individual,2,0.023336322999966796,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,uniform,node,0,0.02396749099898443,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,uniform,node,1,0.024603556001238758,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,uniform,node,2,0.02379652099989471,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,uniform,edge,0,0.024302648000229965,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,uniform,edge,1,0.02512614199986274,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,uniform,edge,2,0.024602867000794504,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,rare,individual,0,0.014113458999418071,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,rare,individual,1,0.01350924999860581,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,rare,individual,2,0.014417320000575273,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,rare,node,0,0.014358804000949021,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,rare,node,1,0.013511810000636615,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,rare,node,2,0.013791530000162311,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,rare,edge,0,0.016899227999601862,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,rare,edge,1,0.01515005100009148,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,100,rare,edge,2,0.013994134000313352,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,1000,uniform,,0,0.028378383998642676,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,1000,uniform,,1,0.02964526600044337,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,1000,uniform,,2,0.03017084299972339,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,uniform,individual,0,0.1165051690004475,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,uniform,individual,1,0.12006181399920024,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,uniform,individual,2,0.12192240799959109,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,uniform,node,0,0.1254568719996314,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,uniform,node,1,0.12637870099933934,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,uniform,node,2,0.12458551100098703,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,uniform,edge,0,0.12738228200032609,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,uniform,edge,1,0.12425055700077792,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,uniform,edge,2,0.12265691699940362,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,rare,individual,0,0.015345580999564845,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,rare,individual,1,0.015087319001395372,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,rare,individual,2,0.014869897000608034,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,rare,node,0,0.018529604998548166,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,rare,node,1,0.016907263001485262,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,rare,node,2,0.017068429000573815,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,rare,edge,0,0.015638313998351805,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,rare,edge,1,0.016809493999971892,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,1000,rare,edge,2,0.015862450998611166,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,10000,uniform,,0,0.24336689300071157,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,10000,uniform,,1,0.24226529100087646,30000,15000,63287,75586,4108,42794,0.0 -sim_trait,10000,uniform,,2,0.2436596239986102,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,uniform,individual,0,1.0077231790000951,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,uniform,individual,1,1.0158602640003664,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,uniform,individual,2,1.011626936000539,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,uniform,node,0,1.0877643929998158,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,uniform,node,1,1.0880109169993375,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,uniform,node,2,1.0914093689989386,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,uniform,edge,0,1.0646407929998531,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,uniform,edge,1,1.0721369829989271,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,uniform,edge,2,1.0633098569996946,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,rare,individual,0,0.024103609001031145,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,rare,individual,1,0.02428249100012181,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,rare,individual,2,0.024011842999243527,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,rare,node,0,0.02510500999960641,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,rare,node,1,0.023118164001061814,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,rare,node,2,0.023174789999757195,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,rare,edge,0,0.023430856001141365,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,rare,edge,1,0.025804556000366574,30000,15000,63287,75586,4108,42794,0.0 -genetic_value,10000,rare,edge,2,0.02633049299947743,30000,15000,63287,75586,4108,42794,0.0 +phase,num_causal,selection,level,replicate,seconds,num_samples,num_individuals,num_nodes,num_edges,num_trees,num_sites +sim_trait,1,uniform,,0,0.0024541959992347984,30000,15000,63287,75586,4108,42794 +sim_trait,1,uniform,,1,0.0020481539995671483,30000,15000,63287,75586,4108,42794 +sim_trait,1,uniform,,2,0.0020018260001961607,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,individual,0,0.012482684000133304,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,individual,1,0.012150070000643609,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,individual,2,0.012130552000598982,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,node,0,0.012228628000229946,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,node,1,0.012219962000017404,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,node,2,0.01234818900047685,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,edge,0,0.012532699000075809,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,edge,1,0.012748904999170918,30000,15000,63287,75586,4108,42794 +genetic_value,1,uniform,edge,2,0.013155032000213396,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,individual,0,0.01225102599892125,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,individual,1,0.011626394998529577,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,individual,2,0.011599434999880032,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,node,0,0.011604778999753762,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,node,1,0.012147759000072256,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,node,2,0.013451222999719903,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,edge,0,0.012008032999801799,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,edge,1,0.012059919001330854,30000,15000,63287,75586,4108,42794 +genetic_value,1,rare,edge,2,0.01325277399882907,30000,15000,63287,75586,4108,42794 +sim_trait,100,uniform,,0,0.006410296000467497,30000,15000,63287,75586,4108,42794 +sim_trait,100,uniform,,1,0.005835321000631666,30000,15000,63287,75586,4108,42794 +sim_trait,100,uniform,,2,0.005629381999824545,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,individual,0,0.021628216998578864,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,individual,1,0.02123565699912433,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,individual,2,0.021466680000230554,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,node,0,0.021798133999254787,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,node,1,0.022633644999586977,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,node,2,0.021683422999558388,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,edge,0,0.021675133999451646,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,edge,1,0.021045367000624537,30000,15000,63287,75586,4108,42794 +genetic_value,100,uniform,edge,2,0.0206327659998351,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,individual,0,0.011630530001639272,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,individual,1,0.011546760999408434,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,individual,2,0.011690916000588913,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,node,0,0.01185587200052396,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,node,1,0.01156106799862755,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,node,2,0.011715804999766988,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,edge,0,0.0116736559994024,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,edge,1,0.014592969999284833,30000,15000,63287,75586,4108,42794 +genetic_value,100,rare,edge,2,0.012491683999542147,30000,15000,63287,75586,4108,42794 +sim_trait,1000,uniform,,0,0.02809734800030128,30000,15000,63287,75586,4108,42794 +sim_trait,1000,uniform,,1,0.026648627999747987,30000,15000,63287,75586,4108,42794 +sim_trait,1000,uniform,,2,0.026458119999006158,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,individual,0,0.11483565899834502,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,individual,1,0.11446014399916749,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,individual,2,0.11398487399856094,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,node,0,0.1097547970002779,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,node,1,0.10957310199955828,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,node,2,0.11040277299980517,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,edge,0,0.10547441699964111,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,edge,1,0.10673052600031951,30000,15000,63287,75586,4108,42794 +genetic_value,1000,uniform,edge,2,0.10701839099965582,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,individual,0,0.013144032000127481,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,individual,1,0.012871730999904685,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,individual,2,0.012051140000039595,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,node,0,0.012279790000320645,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,node,1,0.012358201000097324,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,node,2,0.012161783000919968,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,edge,0,0.012438690000635688,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,edge,1,0.012494279000748065,30000,15000,63287,75586,4108,42794 +genetic_value,1000,rare,edge,2,0.01249857100083318,30000,15000,63287,75586,4108,42794 +sim_trait,10000,uniform,,0,0.22923536500093178,30000,15000,63287,75586,4108,42794 +sim_trait,10000,uniform,,1,0.23691307000080997,30000,15000,63287,75586,4108,42794 +sim_trait,10000,uniform,,2,0.23219180299929576,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,individual,0,0.9973290090001683,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,individual,1,0.9977533809997112,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,individual,2,0.9951367390003725,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,node,0,0.9438101870000537,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,node,1,0.9550599880003574,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,node,2,0.9355915460000688,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,edge,0,0.9174456219989224,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,edge,1,0.9207621040004597,30000,15000,63287,75586,4108,42794 +genetic_value,10000,uniform,edge,2,0.9195124369998666,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,individual,0,0.017908996998812654,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,individual,1,0.017265692000364652,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,individual,2,0.017073662998882355,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,node,0,0.016921062000619713,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,node,1,0.01794522300042445,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,node,2,0.017002496000714018,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,edge,0,0.018678357999306172,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,edge,1,0.017268129999138182,30000,15000,63287,75586,4108,42794 +genetic_value,10000,rare,edge,2,0.017067848000806407,30000,15000,63287,75586,4108,42794 diff --git a/benchmarks/benchmark_genetic_value.py b/benchmarks/benchmark_genetic_value.py index cfbbf28..670e77e 100644 --- a/benchmarks/benchmark_genetic_value.py +++ b/benchmarks/benchmark_genetic_value.py @@ -3,7 +3,7 @@ The interesting regime is a trait with a large number of weakly causal sites, i.e. many causal sites that are each carried by a small number of samples. The -cost of the ARG sweep is the number of nodes that carry a causal allele, so the +cost of the descent is the number of nodes that carry a causal allele, so the allele frequency of the causal sites matters more than how many there are, and ``--selections`` draws them either uniformly over all sites or from the rare ones alone. @@ -26,7 +26,7 @@ import tstrait from tstrait import jit -from tstrait.genetic_value import _COMMON_THRESHOLD, _check_trait_df, _GeneticValue +from tstrait.genetic_value import _check_trait_df, _GeneticValue LEVELS = ["individual", "node", "edge"] SELECTIONS = ["uniform", "rare"] @@ -154,7 +154,7 @@ def select_causal(pool, selection, num_causal, rare_threshold, rng): Draw num_causal rows from the pool, either uniformly over the sites or restricted to the rare ones. Uniform selection is dominated by the common variants in the tail of the frequency spectrum, which behave quite - differently from the weakly causal sites the ARG sweep is aimed at. + differently from the weakly causal sites the descent is aimed at. """ if selection == "rare": pool = pool[pool["allele_freq"] < rare_threshold] @@ -172,12 +172,10 @@ def warm_up(ts, model, levels, counters): before = time.perf_counter() trait_df = tstrait.sim_trait(ts, model=model, num_causal=1, random_seed=1) for level in levels: - # Both implementations, since the grid may use either. - for threshold in (0.0, np.inf): - tstrait.genetic_value(ts, trait_df, level=level, _threshold=threshold) + tstrait.genetic_value(ts, trait_df, level=level) if counters: # A kernel of its own to compile, so only pay for it when it is used. - count_work(ts, _check_trait_df(ts, trait_df), 0.0) + count_work(ts, _check_trait_df(ts, trait_df)) print(f"Warm up (includes numba compilation): {time.perf_counter() - before:.1f}s") @@ -201,12 +199,10 @@ def time_phases(ts, trait_df, level, replicates): genetic = _GeneticValue(ts, checked) size = genetic._output_size(level) - arguments = genetic._descent_arguments(level, 0, np.zeros(size)) + shape = (genetic.num_trait, size) _, times = time_call( # A fresh output array each time, since the kernel accumulates into it. - lambda: jit._push_down_arg( - **{**arguments, "output": np.zeros(len(arguments["output"]))} - ), + lambda: jit._descend_trees(**genetic._descend_arguments(level, np.zeros(shape))), replicates, ) phases.append(("kernel", times)) @@ -236,7 +232,7 @@ def _build_frame(num_trait, size, level, values): COUNTERS = ["rows", "visits"] -def count_work(ts, trait_df, threshold): +def count_work(ts, trait_df): """ Return the work the descent does, as a dict keyed by COUNTERS. @@ -248,44 +244,10 @@ def count_work(ts, trait_df, threshold): The counts do not depend on the level, because the same descent serves all three and only the array the contributions land in differs. """ - genetic = _GeneticValue(ts, trait_df, threshold=threshold) - if not np.any(genetic.descent_rows): - return {"rows": 0, "visits": 0} - visits = jit._descend_trees( - genetic.numba_ts, - ts.edges_parent, - ts.edges_child, - genetic.row_site, - genetic.row_trait, - genetic.row_effect, - genetic.row_ancestral, - genetic.descent_rows, - genetic.pair_offset, - genetic.pair_node, - genetic.pair_carries, - ts.samples().astype(np.int32), - genetic._node_output("node"), - False, - np.zeros((genetic.num_trait, ts.num_nodes)), - ) - return {"rows": int(np.sum(genetic.descent_rows)), "visits": int(visits)} - - -def root_runs_fired(ts, trait_df, descent_rows): - """ - Whether this trait takes the _root_runs branch in _GeneticValue. - - That branch is a Python loop over every tree, so it costs O(num_trees) at - Python speed. It fires when a drawn causal allele happens to be the - ancestral state of its site and that row goes to the push down, which needs - the roots found for it; the descent has them to hand. It therefore appears - and vanishes with the seed and the threshold, which makes for confusing - non-monotonic timings unless it is reported. - """ - site_id = trait_df["site_id"].to_numpy() - causal_allele = trait_df["causal_allele"].to_numpy() - ancestral = ts.sites_ancestral_state[site_id] == causal_allele - return bool(np.any(ancestral & ~descent_rows)) + genetic = _GeneticValue(ts, trait_df) + output = np.zeros((genetic.num_trait, ts.num_nodes)) + visits = jit._descend_trees(**genetic._descend_arguments("node", output)) + return {"rows": len(trait_df), "visits": int(visits)} def carrier_fractions(ts, pool, sample_size, rng): @@ -369,26 +331,13 @@ def run_benchmark(ts, args): if trait_df is None: print(f" too few {selection} sites for num_causal={num_causal}") continue - checked = _check_trait_df(ts, trait_df) - descent = _GeneticValue(ts, checked, threshold=args.threshold).descent_rows - print( - f" {selection} num_causal={num_causal}: " - f"{descent.sum()} of {len(descent)} rows take the descent" - ) - if root_runs_fired(ts, checked, descent): - print( - f" {selection} num_causal={num_causal} takes the _root_runs " - "branch, a Python loop over every tree" - ) if args.counters: - counts[(num_causal, selection)] = count_work(ts, checked, args.threshold) + counts[(num_causal, selection)] = count_work( + ts, _check_trait_df(ts, trait_df) + ) for level in args.levels: call = functools.partial( - tstrait.genetic_value, - ts, - trait_df, - level=level, - _threshold=args.threshold, + tstrait.genetic_value, ts, trait_df, level=level ) if args.memory: _, memory[(num_causal, selection, level)] = peak_memory(call) @@ -446,7 +395,7 @@ def summarise(rows, counts, memory, completed, ts, args): best[key] = min(best.get(key, seconds), seconds) print(f"\n{describe(ts)}") - print(f"Minimum of {args.replicates} replicates, threshold {args.threshold:g}\n") + print(f"Minimum of {args.replicates} replicates\n") columns = f"{'phase':<14} {'selection':<10} {'level':<11} {'num_causal':>10} " columns += f"{'seconds':>10} {'us/site':>10}" if args.counters: @@ -540,7 +489,6 @@ def write_csv(rows, ts, args): "num_edges", "num_trees", "num_sites", - "threshold", ] ) for row in rows: @@ -553,7 +501,6 @@ def write_csv(rows, ts, args): ts.num_edges, ts.num_trees, ts.num_sites, - args.threshold, ] ) print(f"\nWrote {args.output}") @@ -628,16 +575,6 @@ def parse_args(): "taken longer than this" ), ) - parser.add_argument( - "--threshold", - type=float, - default=_COMMON_THRESHOLD, - help=( - "Allele frequency at or above which a causal site goes through the " - "descent of the trees. 0 sends everything there and inf sends " - "everything to the push down of the ARG" - ), - ) parser.add_argument( "--phases", action="store_true", diff --git a/benchmarks/profile_genetic_value.py b/benchmarks/profile_genetic_value.py index a18f7a0..971ab5e 100644 --- a/benchmarks/profile_genetic_value.py +++ b/benchmarks/profile_genetic_value.py @@ -11,9 +11,9 @@ opaque dispatcher frame. ``--mode kernel`` - Set up one cell and run only the kernel, so that a sampling profile is not - swamped by the effect size simulation over the whole site pool. Run it - under perf; ``--mode kernel`` on its own prints the commands. + Set up one cell and run only the descent kernel, so that a sampling profile + is not swamped by the effect size simulation over the whole site pool. Run + it under perf; ``--mode kernel`` on its own prints the commands. Two things about perf and numba are worth knowing before reading a profile. @@ -21,16 +21,15 @@ llvmlite has no LLVM PerfJITEventListener, so nothing emits a /tmp/perf-.map or a jitdump and the kernel shows up as raw addresses under ``[JIT]``. What perf does give is the split between the kernel, the -numba runtime, the interpreter and LLVM compilation, with the typed list -runtime functions resolved by name. For attribution inside the kernel use -``benchmark_genetic_value.py --counters``. +numba runtime, the interpreter and LLVM compilation. For attribution inside +the kernel use ``benchmark_genetic_value.py --counters``. NUMBA_ENABLE_PROFILING=1 is the documented way to profile numba and is the wrong thing to use here. It would only help by way of the listener llvmlite does not have, and it defaults NUMBA_DEBUGINFO to 1, which changes the code -that is generated: the kernel measured 2.62s a call with it and 1.61s -without. perf finds the JIT mappings on its own, so the recipe below does not -set it, and a profile taken with it is a profile of a different kernel. +that is generated and measured the kernel over 60% slower. perf finds the JIT +mappings on its own, so the recipe below does not set it, and a profile taken +with it is a profile of a different kernel. Run with ``uv run --group test benchmarks/profile_genetic_value.py``. """ @@ -110,13 +109,12 @@ def run_kernel(args): """ ts, trait_df = prepare(args) genetic = _GeneticValue(ts, _check_trait_df(ts, trait_df)) - size = genetic._output_size(args.level) - arguments = genetic._descent_arguments(args.level, 0, np.zeros(size)) - jit._push_down_arg(**{**arguments, "output": np.zeros(size)}) + shape = (genetic.num_trait, genetic._output_size(args.level)) + jit._descend_trees(**genetic._descend_arguments(args.level, np.zeros(shape))) before = time.perf_counter() for _ in range(args.repeats): - jit._push_down_arg(**{**arguments, "output": np.zeros(size)}) + jit._descend_trees(**genetic._descend_arguments(args.level, np.zeros(shape))) elapsed = time.perf_counter() - before print( f"{args.repeats} kernel calls in {elapsed:.3f}s " @@ -159,9 +157,8 @@ def print_perf_commands(args): ) print( "\n# The kernel is the [JIT] rows, and libllvmlite is compilation, not\n" - "# work. numba_list_append and numba_list_resize in the second report\n" - "# are the typed list traffic. Source lines inside the kernel are not\n" - "# available; use --counters on the benchmark instead." + "# work. Source lines inside the kernel are not available; use\n" + "# --counters on the benchmark instead." ) diff --git a/tests/test_genetic_value.py b/tests/test_genetic_value.py index 5df1274..20c6d1e 100644 --- a/tests/test_genetic_value.py +++ b/tests/test_genetic_value.py @@ -6,7 +6,6 @@ import tstrait from tstrait.base import _check_numeric_array -from tstrait.genetic_value import _check_trait_df, _GeneticValue from .data import ( all_trees_ts, @@ -37,7 +36,6 @@ def sample_df(): "causal_allele": ["A", "A"], "effect_size": [0.1, 0.1], "trait_id": [0, 0], - "allele_freq": [0.2, 0.3], } ) @@ -117,11 +115,6 @@ def random_trait_df(ts, num_trait, seed, site_id=None): The causal allele of a row is drawn from the alleles that occur at its site, so that the effects are not trivially zero. - - The allele_freq column alternates between two values rather than being the - real frequency. It is only there to decide which implementation a row goes - through, and alternating splits the rows whatever the topology is, which is - what the tests want to exercise. """ rng = np.random.default_rng(seed) if site_id is None: @@ -140,9 +133,6 @@ def random_trait_df(ts, num_trait, seed, site_id=None): "effect_size": rng.normal(size=num_site * num_trait), "trait_id": np.tile(np.arange(num_trait), num_site), "causal_allele": causal_allele, - "allele_freq": np.tile([0.25, 0.75], num_site * num_trait)[ - : num_site * num_trait - ], } ) @@ -1155,23 +1145,14 @@ class TestGeneticValueReference: roots with, so it is covered here rather than separately. """ - # All rows through the descent of the trees, split between the two, and - # all rows through the push down of the ARG. A row whose causal allele is - # the ancestral state always goes through the descent, so the last of these - # is not quite everything. - THRESHOLDS = [0.0, 0.5, np.inf] - def verify(self, ts, num_trait, seed=1, levels=("node", "edge")): trait_df = random_trait_df(ts, num_trait, seed) for level in levels: if level == "individual" and ts.num_individuals == 0: continue expected = naive_genetic_value(ts, trait_df, level) - for threshold in self.THRESHOLDS: - result = tstrait.genetic_value( - ts=ts, trait_df=trait_df, level=level, _threshold=threshold - ) - pd.testing.assert_frame_equal(result, expected, check_dtype=False) + result = tstrait.genetic_value(ts=ts, trait_df=trait_df, level=level) + pd.testing.assert_frame_equal(result, expected, check_dtype=False) @pytest.mark.parametrize( "ts_func", @@ -1277,56 +1258,13 @@ def test_mutation_above_a_root(self): pd.testing.assert_frame_equal(result, expected, check_dtype=False) -class TestDescentAndPushDown: +class TestDescent: """ - The two implementations against each other and against the reference, over - the cases where they are most likely to disagree. - - A row goes through the descent of the trees when its allele frequency is at - or above the threshold, so a threshold of zero sends everything there and - an infinite one sends everything to the push down, apart from the rows - whose causal allele is the ancestral state, which always take the descent. + The cases where the descent of the trees is most likely to go wrong: a + mutation sitting on a root, a causal site carrying no mutations at all, and + a site falling exactly on a tree boundary. """ - def simulated_ts(self, seed=3): - ts = msprime.sim_ancestry( - 20, sequence_length=10_000, recombination_rate=1e-4, random_seed=seed - ) - return multi_allelic_mutations(ts, rate=1e-3, seed=seed) - - @pytest.mark.parametrize("level", ["individual", "node", "edge"]) - @pytest.mark.parametrize("num_trait", [1, 3]) - def test_paths_agree(self, level, num_trait): - # naive_genetic_value is too slow to run on anything this size, so the - # two implementations are checked against each other instead. - ts = self.simulated_ts() - trait_df = random_trait_df(ts, num_trait, seed=5) - descent = tstrait.genetic_value(ts, trait_df, level=level, _threshold=0.0) - push_down = tstrait.genetic_value(ts, trait_df, level=level, _threshold=np.inf) - pd.testing.assert_frame_equal(descent, push_down, check_dtype=False) - - def test_split_is_a_split(self): - # The reference tests would pass just as well if every row went the - # same way, so check that a middling threshold really does divide them. - ts = self.simulated_ts() - trait_df = _check_trait_df(ts, random_trait_df(ts, 2, seed=5)) - genetic = _GeneticValue(ts, trait_df, threshold=0.5) - assert np.any(genetic.descent_rows) - assert np.any(~genetic.descent_rows) - assert not np.all(genetic.descent_rows[genetic.seed_row]) - - def test_no_allele_freq_needs_no_threshold(self): - # A trait dataframe assembled by hand has no allele_freq, so there is - # nothing for a positive threshold to compare against and every row - # takes the push down. The default threshold of zero needs no - # comparison, so those rows take the descent like any other. - ts = self.simulated_ts() - trait_df = _check_trait_df( - ts, random_trait_df(ts, 1, seed=5).drop(columns=["allele_freq"]) - ) - assert not np.any(_GeneticValue(ts, trait_df, threshold=0.5).descent_rows) - assert np.all(_GeneticValue(ts, trait_df, threshold=0.0).descent_rows) - @pytest.mark.parametrize("derived_state", ["A", "T"]) def test_mutation_on_root(self, derived_state): # The ancestral state is the causal allele, so the descent seeds the @@ -1341,17 +1279,13 @@ def test_mutation_on_root(self, derived_state): "effect_size": [1.0], "trait_id": [0], "causal_allele": ["A"], - "allele_freq": [0.5], } ) expected = 1.0 if derived_state == "A" else 0.0 - for threshold in (0.0, np.inf): - result = tstrait.genetic_value( - ts, trait_df, level="node", _threshold=threshold - ) - np.testing.assert_array_equal( - result["genetic_value"], np.full(ts.num_nodes, expected) - ) + result = tstrait.genetic_value(ts, trait_df, level="node") + np.testing.assert_array_equal( + result["genetic_value"], np.full(ts.num_nodes, expected) + ) def test_site_with_no_mutations(self): # A causal site need not carry any mutation, and when its ancestral @@ -1364,14 +1298,10 @@ def test_site_with_no_mutations(self): "effect_size": [1.0], "trait_id": [0], "causal_allele": ["A"], - "allele_freq": [0.5], } ) - for threshold in (0.0, np.inf): - result = tstrait.genetic_value( - ts, trait_df, level="node", _threshold=threshold - ) - np.testing.assert_array_equal(result["genetic_value"], np.ones(ts.num_nodes)) + result = tstrait.genetic_value(ts, trait_df, level="node") + np.testing.assert_array_equal(result["genetic_value"], np.ones(ts.num_nodes)) def test_site_on_a_breakpoint(self): # A site at a tree boundary belongs to the tree that starts there, @@ -1395,12 +1325,8 @@ def test_site_on_a_breakpoint(self): "effect_size": [1.0, 2.0], "trait_id": [0, 0], "causal_allele": ["T", "T"], - "allele_freq": [0.5, 0.5], } ) expected = naive_genetic_value(ts, trait_df, "node") - for threshold in (0.0, np.inf): - result = tstrait.genetic_value( - ts, trait_df, level="node", _threshold=threshold - ) - pd.testing.assert_frame_equal(result, expected, check_dtype=False) + result = tstrait.genetic_value(ts, trait_df, level="node") + pd.testing.assert_frame_equal(result, expected, check_dtype=False) diff --git a/tests/test_jit.py b/tests/test_jit.py index 7d82282..42604d7 100644 --- a/tests/test_jit.py +++ b/tests/test_jit.py @@ -10,8 +10,7 @@ The examples are small trees from the tskit generators, drawn in the class docstrings so that the expected values can be checked against the topology. Each is turned into a tree sequence with a single causal site, since the -kernel pushes the effects of mutations down the ARG rather than working from -a tree. +kernel works from the trees of a tree sequence rather than from a tree. """ import numba @@ -72,7 +71,7 @@ def node_genetic_value(request): state of "A", and a node's value is ``effect_size`` when the allele it inherits is ``causal_allele``. """ - func = kernel(jit._push_down_arg, request.param) + func = kernel(jit._descend_trees, request.param) def f( tree, @@ -90,11 +89,12 @@ def f( "causal_allele": [causal_allele], } ) - # These test the push down kernel, so keep every row on it rather - # than letting the default threshold send them to the descent. - genetic = _GeneticValue(ts, trait_df, threshold=np.inf) - output = np.zeros(genetic._output_size(level)) - return func(**genetic._descent_arguments(level, 0, output)) + genetic = _GeneticValue(ts, trait_df) + # The kernel accumulates into a row per trait and returns the number + # of nodes it visited, so the values come back through the array. + output = np.zeros((genetic.num_trait, genetic._output_size(level))) + func(**genetic._descend_arguments(level, output)) + return output[0] return f diff --git a/tstrait/genetic_value.py b/tstrait/genetic_value.py index 040af0e..b794d1e 100644 --- a/tstrait/genetic_value.py +++ b/tstrait/genetic_value.py @@ -6,38 +6,15 @@ from . import jit from .base import _check_dataframe, _check_instance, _check_non_decreasing # noreorder -# The optional trait dataframe column that routes a causal site to one -# implementation or the other. -ALLELE_FREQ = "allele_freq" - -# Causal sites at or above this allele frequency go through the descent of the -# trees, and the rest through the push down of the ARG. Zero, because the -# descent measured faster than the push down at every point of the benchmark -# grid on both tree sequences, at every level, for one and three traits, and -# for causal sites drawn both uniformly and from the rare ones: 1.33 to 1.70 -# times on uniformly drawn sites and up to 2.27 times on rare ones with three -# traits, where the push down runs its sweep once per trait and the descent -# takes all of them in one pass. The few cells the push down led were a -# fraction of a millisecond apart. -# -# A threshold of zero needs no allele frequency to compare against, so it -# applies to a trait dataframe assembled by hand as well. Private, and kept -# only so that the two can still be compared, since the intention is to end -# with the descent alone. -_COMMON_THRESHOLD = 0.0 - def _row_mutations(ts, trait_df): """ Expand the trait dataframe to one entry per (row, mutation at that row's - site) pair, returning the row, the mutation, whether the mutation carries - the causal allele and whether the state it replaced did. + site) pair, returning the row and the mutation. - Every mutation at a causal site is returned, including those that leave the - causal allele state unchanged. A descent of the trees needs all of them, + Every mutation at a causal site is returned. The descent needs all of them, because a mutation blocks the inheritance of the allele above it whatever - it changes the state to; ``_causal_mutations`` drops the ones that a sum of - state changes does not need. + it changes the state to. """ site_id = trait_df["site_id"].to_numpy() # Mutations are sorted by site, so each site owns a contiguous run of IDs. @@ -49,13 +26,19 @@ def _row_mutations(ts, trait_df): start - (np.cumsum(count) - count), count ) + return row, mutation + + +def _row_causal_allele(ts, trait_df, row, mutation): + """ + Whether each (row, mutation) pair's mutation carries the row's causal + allele. + """ derived_state = ts.mutations_derived_state causal_allele = np.asarray( trait_df["causal_allele"].to_numpy(), dtype=derived_state.dtype )[row] - had_causal_allele = ts.mutations_inherited_state[mutation] == causal_allele - has_causal_allele = derived_state[mutation] == causal_allele - return row, mutation, has_causal_allele, had_causal_allele + return derived_state[mutation] == causal_allele, causal_allele def _causal_mutations(ts, trait_df): @@ -65,63 +48,22 @@ def _causal_mutations(ts, trait_df): allele state for the mutations that change the state. Mutations that leave the state unchanged contribute nothing and are dropped. """ - row, mutation, has_causal_allele, had_causal_allele = _row_mutations(ts, trait_df) + row, mutation = _row_mutations(ts, trait_df) + has_causal_allele, causal_allele = _row_causal_allele(ts, trait_df, row, mutation) + had_causal_allele = ts.mutations_inherited_state[mutation] == causal_allele state_change = has_causal_allele.astype(np.int8) - had_causal_allele.astype(np.int8) changed = state_change != 0 return row[changed], mutation[changed], state_change[changed] -def _root_runs(ts): - """ - Return the roots of the tree sequence as ``(node, left, right)`` arrays, - where each root spans a maximal interval over which it is a root. - - The roots are taken as the children of the virtual root, so that a node - counts as being in a tree exactly when the tree based implementation would - have reached it from the virtual root. Isolated samples are roots. - """ - node = [] - left = [] - right = [] - tree = tskit.Tree(ts) - left_child_array = tree.left_child_array - right_sib_array = tree.right_sib_array - virtual_root = tree.virtual_root - tree.first() - while True: - interval_left, interval_right = tree.interval - u = left_child_array[virtual_root] - while u != tskit.NULL: - if len(node) > 0 and node[-1] == u and right[-1] == interval_left: - right[-1] = interval_right - else: - node.append(u) - left.append(interval_left) - right.append(interval_right) - u = right_sib_array[u] - if not tree.next(): - break - return ( - np.array(node, dtype=np.int32), - np.array(left, dtype=float), - np.array(right, dtype=float), - ) - - def _check_trait_df(ts, trait_df): """ Check the trait dataframe against the tree sequence, returning the required columns with trait_id cast to int. - - ``allele_freq`` is kept when it is there. It is not required, and a trait - dataframe assembled by hand will not have it, but ``sim_trait`` returns it - and it is what decides which of the two implementations a causal site goes - through. """ - columns = ["site_id", "effect_size", "trait_id", "causal_allele"] - if ALLELE_FREQ in getattr(trait_df, "columns", []): - columns.append(ALLELE_FREQ) - trait_df = _check_dataframe(trait_df, columns, "trait_df") + trait_df = _check_dataframe( + trait_df, ["site_id", "effect_size", "trait_id", "causal_allele"], "trait_df" + ) if len(trait_df) == 0: raise ValueError("trait_df must contain at least one row") _check_non_decreasing(trait_df["site_id"], "site_id") @@ -141,14 +83,13 @@ class _GeneticValue: """ GeneticValue class to compute genetic values of individuals, nodes, or edges. - The genetic values of every causal site are accumulated in a single descent - of the ARG. Each causal mutation changes the causal allele state by - ``[derived == causal] - [inherited == causal]``, and that change applies to - every node below it. These changes telescope down a path, so a node's value - is the sum of the changes on the mutations above it, plus the effect size - of every causal site whose ancestral state is itself the causal allele. - That makes the value additive along a root to node path, which is what the - descent accumulates. Nested and back mutations need no special handling. + The genetic values of every causal site are accumulated in a single pass + over the trees. Each tree is built from the one before it, and for each + causal site it carries the effect is added to the nodes that inherit the + causal allele: those reached by descending from a mutation carrying it, + stopping wherever another mutation at that site replaces the state. Where + the causal allele is the site's ancestral state the roots are descended + from instead, which reaches exactly the nodes of the tree. Parameters ---------- @@ -159,117 +100,32 @@ class _GeneticValue: size, and trait ID. """ - def __init__(self, ts, trait_df, threshold=_COMMON_THRESHOLD): - columns = ["site_id", "effect_size", "trait_id", "causal_allele"] - if ALLELE_FREQ in trait_df.columns: - columns.append(ALLELE_FREQ) - self.trait_df = trait_df[columns] + def __init__(self, ts, trait_df): + self.trait_df = trait_df[["site_id", "effect_size", "trait_id", "causal_allele"]] self.ts = ts self.numba_ts = tskit_numba.jitwrap(ts) - self.child_index = self.numba_ts.child_index() site_id = self.trait_df["site_id"].to_numpy() - effect_size = self.trait_df["effect_size"].to_numpy() self.trait_id = self.trait_df["trait_id"].to_numpy() self.num_trait = np.max(self.trait_id) + 1 - # Causal sites are identified by their rank among the causal sites, so - # that matching an edge to one is an integer comparison and a position - # that cannot matter is never considered. - causal_site = np.unique(site_id) - causal_position = ts.sites_position[causal_site] - rows_site = np.searchsorted(causal_site, site_id) - self.edges_site_start = np.searchsorted(causal_position, ts.edges_left).astype( - np.int32 - ) - self.edges_site_stop = np.searchsorted(causal_position, ts.edges_right).astype( - np.int32 - ) - # tskit does not require the node IDs to be in time order. - self.nodes_by_time = np.argsort(-ts.nodes_time, kind="stable").astype(np.int32) - - # Every mutation at a causal site, which is what the descent needs, - # and the state changing ones, which is what the push down seeds from. - pair_row, pair_mutation, has, had = _row_mutations(ts, self.trait_df) + # One entry per row of the trait dataframe, walked in step with the + # trees, and one per mutation at each row's site. self.row_site = site_id.astype(np.int32) self.row_trait = self.trait_id.astype(np.int32) - self.row_effect = effect_size.astype(float) + self.row_effect = self.trait_df["effect_size"].to_numpy().astype(float) self.row_ancestral = ( ts.sites_ancestral_state[site_id] == self.trait_df["causal_allele"].to_numpy() ) + pair_row, pair_mutation = _row_mutations(ts, self.trait_df) self.pair_offset = np.searchsorted( pair_row, np.arange(len(self.trait_df) + 1) ).astype(np.int64) self.pair_node = ts.mutations_node[pair_mutation].astype(np.int32) - self.pair_carries = has - - # Rows go to the descent of the trees when their causal allele is - # common enough for the tree pass to pay for itself, and when it is the - # ancestral state, where the descent has the roots to hand and the push - # down would need them found for it. Every frequency is at or above a - # threshold of zero, so there is nothing to look up in that case, which - # is what lets the default apply to a trait dataframe that has no - # allele frequency in it. - if threshold <= 0: - self.descent_rows = np.ones(len(self.trait_df), dtype=bool) - elif ALLELE_FREQ in self.trait_df.columns: - self.descent_rows = self.trait_df[ALLELE_FREQ].to_numpy() >= threshold - self.descent_rows |= self.row_ancestral - else: - self.descent_rows = np.zeros(len(self.trait_df), dtype=bool) - - state_change = has.astype(np.int8) - had.astype(np.int8) - changed = state_change != 0 - row = pair_row[changed] - mutation = pair_mutation[changed] - state_change = state_change[changed] - # The row a seed came from, so that a subset of the rows can be - # selected without recomputing the seeds. - self.seed_row = row - self.seed_trait = self.trait_id[row] - self.seed_node = ts.mutations_node[mutation].astype(np.int32) - self.seed_site = rows_site[row].astype(np.int32) - self.seed_weight = effect_size[row] * state_change - # A mutation above a root has no edge, and needs no special handling: - # seeding at its node and pushing down is already right, and the - # missing edge contribution matches the tree based implementation. - self.seed_edge = ts.mutations_edge[mutation] - - # When the ancestral state of a site is the causal allele, every node - # in its tree carries it. There is no mutation to seed from, so the - # roots are seeded instead and the effect reaches the same nodes. - # Only for the rows the push down is taking: the descent seeds the - # roots of the tree it is already standing in. - ancestral = np.flatnonzero(self.row_ancestral & ~self.descent_rows) - if len(ancestral) > 0: - roots, roots_left, roots_right = _root_runs(ts) - start = np.searchsorted(causal_position, roots_left) - stop = np.searchsorted(causal_position, roots_right) - # Each root run takes the ancestral rows spanned by its interval. - ancestral_site = rows_site[ancestral] - first = np.searchsorted(ancestral_site, start) - count = np.searchsorted(ancestral_site, stop) - first - run = np.repeat(np.arange(len(roots)), count) - index = np.arange(count.sum()) + np.repeat( - first - (np.cumsum(count) - count), count - ) - self.seed_row = np.concatenate([self.seed_row, ancestral[index]]) - self.seed_trait = np.concatenate( - [self.seed_trait, self.trait_id[ancestral[index]]] - ) - self.seed_node = np.concatenate([self.seed_node, roots[run]]) - self.seed_site = np.concatenate( - [self.seed_site, ancestral_site[index].astype(np.int32)] - ) - self.seed_weight = np.concatenate( - [self.seed_weight, effect_size[ancestral[index]]] - ) - # A root has no edge above it, so it contributes nothing at the - # edge level, which is what the tree based implementation does too. - self.seed_edge = np.concatenate( - [self.seed_edge, np.full(len(run), tskit.NULL, dtype=np.int32)] - ) + self.pair_carries, _ = _row_causal_allele( + ts, self.trait_df, pair_row, pair_mutation + ) def _output_size(self, level): return { @@ -292,35 +148,27 @@ def _node_output(self, level): return self.ts.nodes_individual return np.arange(self.ts.num_nodes, dtype=np.int32) - def _descent_arguments(self, level, trait, output): + def _descend_arguments(self, level, output): """ - Return the arguments to the push down kernel for one trait, with the - contributions directed at nodes, individuals or edges according to - ``level`` and accumulated into ``output``. + Return the arguments to the descent kernel, with the contributions + directed at nodes, individuals or edges according to ``level`` and + accumulated into ``output``. """ ts = self.ts - in_trait = (self.seed_trait == trait) & ~self.descent_rows[self.seed_row] - if level == "edge": - edges_output = np.arange(ts.num_edges, dtype=np.int32) - # A seed is credited to the edge above the mutation, which does not - # exist when the mutation is above a root. - seed_output = self.seed_edge[in_trait].astype(np.int32) - else: - node_output = self._node_output(level) - edges_output = node_output[ts.edges_child] - seed_output = node_output[self.seed_node[in_trait]] - return { - "child_index": self.child_index, + "numba_ts": self.numba_ts, + "edges_parent": ts.edges_parent, "edges_child": ts.edges_child, - "edges_output": edges_output, - "edges_site_start": self.edges_site_start, - "edges_site_stop": self.edges_site_stop, - "nodes_by_time": self.nodes_by_time, - "seed_node": self.seed_node[in_trait], - "seed_site": self.seed_site[in_trait], - "seed_weight": self.seed_weight[in_trait], - "seed_output": seed_output, + "row_site": self.row_site, + "row_trait": self.row_trait, + "row_effect": self.row_effect, + "row_ancestral": self.row_ancestral, + "pair_offset": self.pair_offset, + "pair_node": self.pair_node, + "pair_carries": self.pair_carries, + "samples": ts.samples().astype(np.int32), + "node_output": self._node_output(level), + "edge_level": level == "edge", "output": output, } @@ -334,34 +182,9 @@ def _run(self, level): pandas.DataFrame Dataframe with trait ID, [individual|node|edge] ID, and genetic value. """ - ts = self.ts N = self._output_size(level) genetic_value_table = np.zeros((self.num_trait, N)) - - if np.any(self.descent_rows): - jit._descend_trees( - self.numba_ts, - ts.edges_parent, - ts.edges_child, - self.row_site, - self.row_trait, - self.row_effect, - self.row_ancestral, - self.descent_rows, - self.pair_offset, - self.pair_node, - self.pair_carries, - ts.samples().astype(np.int32), - self._node_output(level), - level == "edge", - genetic_value_table, - ) - # Compiling the push down is not worth it when nothing is left for it. - if not np.all(self.descent_rows[self.seed_row]): - for trait in range(self.num_trait): - jit._push_down_arg( - **self._descent_arguments(level, trait, genetic_value_table[trait]) - ) + jit._descend_trees(**self._descend_arguments(level, genetic_value_table)) return pd.DataFrame( { @@ -372,7 +195,7 @@ def _run(self, level): ) -def genetic_value(ts, trait_df, level="individual", *, _threshold=_COMMON_THRESHOLD): +def genetic_value(ts, trait_df, level="individual"): """ Compute genetic values for a tree sequence given a trait dataframe. @@ -417,7 +240,7 @@ def genetic_value(ts, trait_df, level="individual", *, _threshold=_COMMON_THRESH raise ValueError("No individuals in the provided tree sequence dataset") trait_df = _check_trait_df(ts, trait_df) - genetic = _GeneticValue(ts=ts, trait_df=trait_df, threshold=_threshold) + genetic = _GeneticValue(ts=ts, trait_df=trait_df) genetic_result = genetic._run(level) diff --git a/tstrait/jit.py b/tstrait/jit.py index c491a3e..a17e3f5 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -3,18 +3,32 @@ Two conventions apply throughout this module: -1. Genetic values are accumulated by pushing the effect of each causal - mutation down the ARG, rather than by working through the trees one at a - time. Positions are expressed as indexes into the sorted array of causal - sites, so that an edge is matched to a causal site with an integer - comparison and a genomic coordinate never appears. +1. Genetic values are accumulated by building each tree in turn and descending + from the mutations of the causal sites that tree carries, so that the cost + is the number of nodes carrying a causal allele rather than the size of the + tree sequence. 2. Numba compiles with bounds checking disabled by default, so an out of bounds access here fails silently rather than raising ``IndexError``. Keeping the compiled functions in one module lets us unit test them through their ``py_func`` attribute in ``tests/test_jit.py``, which runs the untranslated Python and therefore gets numpy's bounds checking and - coverage measurement for free. + coverage measurement for free. ``_descend_trees`` calls the tree building + kernels rather than their ``py_func``, so running it that way interprets + only its own loop; those kernels are covered instead by ``TestTreeState``, + which compares them against ``tskit.Tree`` tree for tree. + +An earlier implementation pushed the effect of each causal mutation down the +ARG in a single sweep of the nodes from the past to the present, and never +built a tree at all. It was measured against this one over the whole benchmark +grid, on both tree sequences, at each level, for one and three traits, and for +causal sites drawn both uniformly and from the rare ones. It was slower +everywhere the cost mattered: 1.33 to 1.70 times on uniformly drawn causal +sites, and up to 2.27 times on rare ones with three traits, since it swept once +per trait where this makes a single pass for all of them. It was ahead only for +a few rare causal sites, which give the tree building here too little to +amortise it against, and then by a millisecond or two on a call taking a few. +It is in the git history if it is ever wanted back. """ from collections import namedtuple @@ -22,102 +36,6 @@ import numba import numpy as np import tskit -from numba.core import types -from numba.typed import List - -# What a node holds while it waits to be swept: the indexes of the seeds whose -# effect has reached it. The causal site and effect size of a seed are looked -# up in arrays that are fixed for the whole sweep, so an item is one integer. -_SEED_LIST = types.ListType(types.int32) - - -@numba.njit -def _push_down_arg( - child_index, - edges_child, - edges_output, - edges_site_start, - edges_site_stop, - nodes_by_time, - seed_node, - seed_site, - seed_weight, - seed_output, - output, -): - """ - Accumulate the genetic value of every causal site in one sweep of the nodes. - - Each node holds the seeds whose effect has reached it, seeded at the node - of each causal mutation. Sweeping the nodes from the past to the present - visits every node after all of its ancestors, because a parent is always - older than its child, so a seed can be passed from a node to its children - and is guaranteed to have arrived before that child is reached. For each - seed the outbound edges spanning its causal site are found, and the seed is - credited to the edge and to the node at the far end before being passed on - to it. - - A node is credited when a seed arrives rather than when it is swept, so a - node that is never a parent never holds anything. On a large tree sequence - that is about half of the work. - - ``child_index`` is the tskit child index, in which a node that is never a - parent has the range ``(-1, -1)``. ``nodes_by_time`` lists the nodes from - the oldest to the youngest. ``edges_output`` and ``seed_output`` give the - index in ``output`` that each contribution is added to, which is how the - same sweep serves both node and edge genetic values; a negative index - discards the contribution. - """ - num_nodes = len(child_index) - # A node's list is made when something first reaches it. Most nodes are - # never reached when the causal alleles are rare, and making a list for - # every one of them up front then costs more than the sweep does. - empty = List.empty_list(types.int32) - pending = List.empty_list(_SEED_LIST) - for _ in range(num_nodes): - pending.append(empty) - reached = np.zeros(num_nodes, dtype=np.bool_) - - for j in range(len(seed_node)): - u = seed_node[j] - if seed_output[j] >= 0: - output[seed_output[j]] += seed_weight[j] - if child_index[u, 0] < 0: - continue - if not reached[u]: - pending[u] = List.empty_list(types.int32) - reached[u] = True - pending[u].append(np.int32(j)) - - for i in range(len(nodes_by_time)): - parent = nodes_by_time[i] - if not reached[parent]: - continue - items = pending[parent] - edge_start = child_index[parent, 0] - edge_stop = child_index[parent, 1] - for k in range(len(items)): - item = items[k] - site = seed_site[item] - weight = seed_weight[item] - for e in range(edge_start, edge_stop): - if edges_site_start[e] <= site and site < edges_site_stop[e]: - if edges_output[e] >= 0: - output[edges_output[e]] += weight - child = edges_child[e] - if child_index[child, 0] < 0: - # Nothing below, so there is nothing to hold. - continue - if not reached[child]: - pending[child] = List.empty_list(types.int32) - reached[child] = True - pending[child].append(np.int32(item)) - # A node is swept once, so dropping the reference here hands its - # storage back for the nodes still to come. - pending[parent] = empty - - return output - # The state of one tree, in the tskit quintuply linked encoding. A tree is # built by applying the edge differences between it and the tree before it, so @@ -236,7 +154,6 @@ def _descend_trees( row_trait, row_effect, row_ancestral, - row_selected, pair_offset, pair_node, pair_carries, @@ -285,9 +202,6 @@ def _descend_trees( _apply_edge_diffs(tree_index, edges_parent, edges_child, tree) site_stop = tree_index.site_range[1] while row < num_rows and row_site[row] < site_stop: - if not row_selected[row]: - row += 1 - continue start = pair_offset[row] stop = pair_offset[row + 1] # Every mutation at the site blocks the allele above it, whatever From 9f0bb966e5f5c1690541838689ad3a35db92ec00 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Thu, 3 Sep 2026 13:27:39 +0100 Subject: [PATCH 12/19] Interpret the tree building kernels under the nojit tests Coverage of jit.py fell to 62% when the push down went. It had called no other kernel, so running it through py_func interpreted everything it did; the descent calls five kernels to build the trees, and a numba function calls whatever the name is bound to in the module, so those five stayed compiled and the coverage of their bodies went with them. Swap the whole set for their py_func together when the nojit parameter asks for it, so that a kernel is interpreted the whole way down rather than only in its own loop, and give the tree state test the same parameter as the rest. That is what the module docstring has always claimed the convention is; correct the caveat that said otherwise. jit.py is back to 100%, and so is every other module, on branches as well as lines. Running only the nojit half covers jit.py completely and only the jit half covers 16% of it, so the measurement is coming from the interpreted path rather than from somewhere incidental. --- tests/test_jit.py | 47 ++++++++++++++++++++++++++++++++++++++--------- tstrait/jit.py | 8 ++++---- 2 files changed, 42 insertions(+), 13 deletions(-) diff --git a/tests/test_jit.py b/tests/test_jit.py index 42604d7..4cbc63c 100644 --- a/tests/test_jit.py +++ b/tests/test_jit.py @@ -36,11 +36,30 @@ ANCESTRAL_STATE = "A" -def kernel(func, param): +# The kernels that _descend_trees and the test driver below call. A numba +# function calls whatever the name is bound to in the module when it runs, so +# swapping these for the Python they were written as makes a caller that is +# itself untranslated interpreted the whole way down rather than only in its +# own loop. +TREE_KERNELS = [ + "tree_state", + "_remove_edge", + "_insert_edge", + "_apply_edge_diffs", + "_tree_roots", +] + + +def kernel(func, param, monkeypatch): """ - Return either the compiled kernel or the Python it was written as. + Return either the compiled kernel or the Python it was written as, in the + latter case untranslating the kernels it calls along with it. """ - return func if param == "jit" else func.py_func + if param == "jit": + return func + for name in TREE_KERNELS: + monkeypatch.setattr(jit, name, getattr(jit, name).py_func) + return func.py_func def one_site(tree, mutations): @@ -62,7 +81,7 @@ def one_site(tree, mutations): @pytest.fixture(params=["jit", "nojit"]) -def node_genetic_value(request): +def node_genetic_value(request, monkeypatch): """ Return a function computing the node genetic values of a tree carrying one causal site. @@ -71,7 +90,7 @@ def node_genetic_value(request): state of "A", and a node's value is ``effect_size`` when the allele it inherits is ``causal_allele``. """ - func = kernel(jit._descend_trees, request.param) + func = kernel(jit._descend_trees, request.param, monkeypatch) def f( tree, @@ -515,10 +534,20 @@ def _walk_trees( return i -def walk_trees(ts): +@pytest.fixture(params=["jit", "nojit"]) +def walk_trees(request, monkeypatch): """ - Return the per tree state arrays that _walk_trees records. + Return a function giving the per tree state arrays that _walk_trees records. """ + driver = kernel(_walk_trees, request.param, monkeypatch) + + def f(ts): + return _walk(driver, ts) + + return f + + +def _walk(driver, ts): numba_ts = tskit_numba.jitwrap(ts) shape = (max(ts.num_trees, 1), ts.num_nodes) got = { @@ -526,7 +555,7 @@ def walk_trees(ts): for name in ("parent", "left_child", "right_sib", "node_edge", "roots") } num_roots = np.zeros(shape[0], dtype=np.int32) - count = _walk_trees( + count = driver( numba_ts, ts.edges_parent, ts.edges_child, @@ -578,7 +607,7 @@ class TestTreeState: empty_ts(), ], ) - def test_matches_tskit(self, ts): + def test_matches_tskit(self, ts, walk_trees): got, num_roots = walk_trees(ts) n = ts.num_nodes for i, tree in enumerate(ts.trees()): diff --git a/tstrait/jit.py b/tstrait/jit.py index a17e3f5..b6af559 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -13,10 +13,10 @@ Keeping the compiled functions in one module lets us unit test them through their ``py_func`` attribute in ``tests/test_jit.py``, which runs the untranslated Python and therefore gets numpy's bounds checking and - coverage measurement for free. ``_descend_trees`` calls the tree building - kernels rather than their ``py_func``, so running it that way interprets - only its own loop; those kernels are covered instead by ``TestTreeState``, - which compares them against ``tskit.Tree`` tree for tree. + coverage measurement for free. A kernel calling another one calls whatever + the name is bound to in this module, so the tests swap the whole set for + their ``py_func`` together and a kernel is interpreted the whole way down + rather than only in its own loop. An earlier implementation pushed the effect of each causal mutation down the ARG in a single sweep of the nodes from the past to the present, and never From aca6627a6bb26b0aefa369a01322c7c5486e2f3e Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Thu, 3 Sep 2026 16:48:11 +0100 Subject: [PATCH 13/19] Divide the causal sites between threads on request The rows of the trait dataframe are independent of each other, so add a num_threads argument that gives each of that many worker threads a contiguous range of them. It defaults to 0, which does the work on the calling thread with no pool created, matching what tskit's divergence_matrix and its neighbours mean by the same argument. Releasing the GIL is what makes it worth anything: a numba kernel holds it for its whole execution by default, so the threads would have taken turns. With nogil=True on the descent, four of them report a busy to wall ratio of 3.8 out of 4. A range needs nothing from any other. The kernel takes row_start and row_stop rather than a slice, so nothing is copied or rebased: pair_offset is indexed by absolute row, and the arrays a thread only reads are built once and shared. Each range accumulates into its own table, since the kernel adds into it and that is not atomic, and the tables are summed at the end, which is where the answer stops being bit for bit what one thread would have produced. The rows are marked with their own index, and two ranges cannot collide. Minimum of two replicates at level="node" on four cores: preset selection num_causal 1 thread 2 3 4 small uniform 10,000 980ms 1.94x 2.86x 3.41x small uniform 300 41ms 1.36x 1.66x 1.69x small rare 10,000 17ms 0.94x 0.91x 0.89x large uniform 10,000 3338ms 1.47x 1.64x 1.64x large rare 10,000 76ms 1.01x 1.00x 0.99x Two things in there are worth saying out loud, since neither is what one would assume. A bigger trait parallelises and a bigger tree sequence does not: the kernel holds around eleven arrays the length of the nodes per thread, 49 bytes a node, which is 3.1MB on 30,000 samples so four threads sit inside a 16MB L3, and 10.6MB on 100,000 samples so four want 42MB and are held up by memory bandwidth instead. Going from 10,000 causal sites to 100,000 on the larger one moves four threads from 1.50x only to 1.61x, so this is cache capacity rather than the division of work. And threads cost more than they save on a rare or a small trait, because a thread walks the whole tree sequence whatever range it takes: there is no seeking to the first tree a range wants, and that pass is nothing against a second of descending and everything against seventeen milliseconds. Equal row counts turned out to need no cost model behind them. They were the thing most likely to want one, since rows differ enormously in how many nodes they reach, but four ranges came out within 1.11x of each other on time and 1.18x on nodes visited. The sequential path is unchanged, within 0.96x to 1.09x of the committed baseline over five replicates. --- CHANGELOG.md | 9 ++ benchmarks/README.md | 37 ++++++ benchmarks/baseline_small.csv | 170 +++++++++++++------------- benchmarks/benchmark_genetic_value.py | 26 +++- docs/genetic.md | 12 ++ tests/test_genetic_value.py | 111 +++++++++++++++++ tstrait/genetic_value.py | 73 +++++++++-- tstrait/jit.py | 18 ++- 8 files changed, 358 insertions(+), 98 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index dd215bd..934697c 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,6 +13,15 @@ In development ### Performance +- `genetic_value` takes a `num_threads` argument, dividing the causal sites + between that many worker threads and defaulting to 0, which does the work on + the calling thread. Each thread holds arrays the length of the nodes, so how + well it scales is set by the size of the tree sequence rather than by the + number of causal sites: on 30,000 samples a trait of 10,000 uniformly drawn + causal sites is 3.4 times faster on four threads, while on 100,000 samples, + where the arrays no longer fit in cache, the same trait is 1.6 times faster. + Threads do not pay for themselves on a trait whose causal sites are few or + rare, since each of them walks the trees. - `genetic_value` descends from the mutations of each causal site instead of making a pass over every node for each of them, so its cost is the number of nodes carrying a causal allele rather than the number of causal sites times diff --git a/benchmarks/README.md b/benchmarks/README.md index 28105bf..f476906 100644 --- a/benchmarks/README.md +++ b/benchmarks/README.md @@ -89,6 +89,43 @@ was tried and rejected: it takes `sites/nodes` from 0.68 to about 2.7 against the large preset's 1.08, and the per-call setup floor grows with the number of sites, so the low end of the curve stops being comparable. +### Threads + +`--num-threads` divides the causal sites between that many worker threads, and +defaults to 0, which does the work on the calling thread. The rows of the trait +dataframe are independent, so a thread takes a contiguous range of them and its +own accumulator, and the accumulators are summed at the end. + +Minimum of two replicates at `level="node"`, on four cores: + +| preset | selection | num_causal | 1 thread | 2 | 3 | 4 | +|---|---|---|---|---|---|---| +| small | uniform | 10,000 | 980ms | 1.94x | 2.86x | **3.41x** | +| small | uniform | 300 | 41ms | 1.36x | 1.66x | 1.69x | +| small | rare | 10,000 | 17ms | 0.94x | 0.91x | 0.89x | +| small | rare | 300 | 12ms | 0.82x | 0.78x | 0.77x | +| large | uniform | 10,000 | 3338ms | 1.47x | 1.64x | 1.64x | +| large | rare | 10,000 | 76ms | 1.01x | 1.00x | 0.99x | + +Two things in that table are worth understanding before reading a number off it +as a bug. + +**A bigger trait parallelises; a bigger tree sequence does not.** The kernel +holds about eleven arrays the length of the nodes per thread, 49 bytes a node. +That is 3.1MB on `small`, so four threads fit inside this machine's 16MB L3 and +scale nearly linearly, and 10.6MB on `large`, where four threads want 42MB and +are limited by memory bandwidth instead. Raising the causal sites from 10,000 to +100,000 on `large` moves four threads from 1.50x only to 1.61x, and shrinking +the two stamp arrays to `int32` gets 1.74x, so this is cache capacity rather +than anything the division of work can fix. A machine with more L3 or more +memory channels should do better on `large`. + +**Threads do not pay for themselves on a rare or a small trait.** A thread walks +the whole tree sequence whatever range of rows it takes, because a tree is built +from the one before it and there is no seeking to the first tree a range wants. +That pass is 2.2ms on `small` and 12.4ms on `large`, so it is nothing against a +second of descending and everything against seventeen milliseconds of it. + ### What else it measures Wall time on its own does not say why a configuration is slow. Four optional diff --git a/benchmarks/baseline_small.csv b/benchmarks/baseline_small.csv index cdbba44..9403205 100644 --- a/benchmarks/baseline_small.csv +++ b/benchmarks/baseline_small.csv @@ -1,85 +1,85 @@ -phase,num_causal,selection,level,replicate,seconds,num_samples,num_individuals,num_nodes,num_edges,num_trees,num_sites -sim_trait,1,uniform,,0,0.0024541959992347984,30000,15000,63287,75586,4108,42794 -sim_trait,1,uniform,,1,0.0020481539995671483,30000,15000,63287,75586,4108,42794 -sim_trait,1,uniform,,2,0.0020018260001961607,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,individual,0,0.012482684000133304,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,individual,1,0.012150070000643609,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,individual,2,0.012130552000598982,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,node,0,0.012228628000229946,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,node,1,0.012219962000017404,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,node,2,0.01234818900047685,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,edge,0,0.012532699000075809,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,edge,1,0.012748904999170918,30000,15000,63287,75586,4108,42794 -genetic_value,1,uniform,edge,2,0.013155032000213396,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,individual,0,0.01225102599892125,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,individual,1,0.011626394998529577,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,individual,2,0.011599434999880032,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,node,0,0.011604778999753762,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,node,1,0.012147759000072256,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,node,2,0.013451222999719903,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,edge,0,0.012008032999801799,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,edge,1,0.012059919001330854,30000,15000,63287,75586,4108,42794 -genetic_value,1,rare,edge,2,0.01325277399882907,30000,15000,63287,75586,4108,42794 -sim_trait,100,uniform,,0,0.006410296000467497,30000,15000,63287,75586,4108,42794 -sim_trait,100,uniform,,1,0.005835321000631666,30000,15000,63287,75586,4108,42794 -sim_trait,100,uniform,,2,0.005629381999824545,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,individual,0,0.021628216998578864,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,individual,1,0.02123565699912433,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,individual,2,0.021466680000230554,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,node,0,0.021798133999254787,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,node,1,0.022633644999586977,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,node,2,0.021683422999558388,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,edge,0,0.021675133999451646,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,edge,1,0.021045367000624537,30000,15000,63287,75586,4108,42794 -genetic_value,100,uniform,edge,2,0.0206327659998351,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,individual,0,0.011630530001639272,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,individual,1,0.011546760999408434,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,individual,2,0.011690916000588913,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,node,0,0.01185587200052396,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,node,1,0.01156106799862755,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,node,2,0.011715804999766988,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,edge,0,0.0116736559994024,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,edge,1,0.014592969999284833,30000,15000,63287,75586,4108,42794 -genetic_value,100,rare,edge,2,0.012491683999542147,30000,15000,63287,75586,4108,42794 -sim_trait,1000,uniform,,0,0.02809734800030128,30000,15000,63287,75586,4108,42794 -sim_trait,1000,uniform,,1,0.026648627999747987,30000,15000,63287,75586,4108,42794 -sim_trait,1000,uniform,,2,0.026458119999006158,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,individual,0,0.11483565899834502,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,individual,1,0.11446014399916749,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,individual,2,0.11398487399856094,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,node,0,0.1097547970002779,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,node,1,0.10957310199955828,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,node,2,0.11040277299980517,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,edge,0,0.10547441699964111,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,edge,1,0.10673052600031951,30000,15000,63287,75586,4108,42794 -genetic_value,1000,uniform,edge,2,0.10701839099965582,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,individual,0,0.013144032000127481,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,individual,1,0.012871730999904685,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,individual,2,0.012051140000039595,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,node,0,0.012279790000320645,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,node,1,0.012358201000097324,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,node,2,0.012161783000919968,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,edge,0,0.012438690000635688,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,edge,1,0.012494279000748065,30000,15000,63287,75586,4108,42794 -genetic_value,1000,rare,edge,2,0.01249857100083318,30000,15000,63287,75586,4108,42794 -sim_trait,10000,uniform,,0,0.22923536500093178,30000,15000,63287,75586,4108,42794 -sim_trait,10000,uniform,,1,0.23691307000080997,30000,15000,63287,75586,4108,42794 -sim_trait,10000,uniform,,2,0.23219180299929576,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,individual,0,0.9973290090001683,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,individual,1,0.9977533809997112,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,individual,2,0.9951367390003725,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,node,0,0.9438101870000537,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,node,1,0.9550599880003574,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,node,2,0.9355915460000688,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,edge,0,0.9174456219989224,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,edge,1,0.9207621040004597,30000,15000,63287,75586,4108,42794 -genetic_value,10000,uniform,edge,2,0.9195124369998666,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,individual,0,0.017908996998812654,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,individual,1,0.017265692000364652,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,individual,2,0.017073662998882355,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,node,0,0.016921062000619713,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,node,1,0.01794522300042445,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,node,2,0.017002496000714018,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,edge,0,0.018678357999306172,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,edge,1,0.017268129999138182,30000,15000,63287,75586,4108,42794 -genetic_value,10000,rare,edge,2,0.017067848000806407,30000,15000,63287,75586,4108,42794 +phase,num_causal,selection,level,replicate,seconds,num_samples,num_individuals,num_nodes,num_edges,num_trees,num_sites,num_threads +sim_trait,1,uniform,,0,0.0024442089998046868,30000,15000,63287,75586,4108,42794,0 +sim_trait,1,uniform,,1,0.00206001099650166,30000,15000,63287,75586,4108,42794,0 +sim_trait,1,uniform,,2,0.0021039189996372443,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,uniform,individual,0,0.01224575399828609,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,uniform,individual,1,0.012782063000486232,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,uniform,individual,2,0.013174051000532927,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,uniform,node,0,0.01399092200153973,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,uniform,node,1,0.012024119001580402,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,uniform,node,2,0.012036211999657098,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,uniform,edge,0,0.012211093002406415,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,uniform,edge,1,0.013050546000158647,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,uniform,edge,2,0.0123963600017305,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,rare,individual,0,0.01127879400155507,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,rare,individual,1,0.012140853999881074,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,rare,individual,2,0.012874144002125831,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,rare,node,0,0.015231686000333866,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,rare,node,1,0.012208834999910323,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,rare,node,2,0.011544229000719497,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,rare,edge,0,0.011383757002477068,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,rare,edge,1,0.012330576999374898,30000,15000,63287,75586,4108,42794,0 +genetic_value,1,rare,edge,2,0.011485152997920522,30000,15000,63287,75586,4108,42794,0 +sim_trait,100,uniform,,0,0.005982647002383601,30000,15000,63287,75586,4108,42794,0 +sim_trait,100,uniform,,1,0.005803553998703137,30000,15000,63287,75586,4108,42794,0 +sim_trait,100,uniform,,2,0.005617545000859536,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,uniform,individual,0,0.021615510999254184,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,uniform,individual,1,0.02126154999859864,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,uniform,individual,2,0.02122037899971474,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,uniform,node,0,0.020638850000977982,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,uniform,node,1,0.021018144998379285,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,uniform,node,2,0.02054931600287091,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,uniform,edge,0,0.020246869000402512,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,uniform,edge,1,0.020713948000775417,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,uniform,edge,2,0.020244728999387007,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,rare,individual,0,0.011508805000630673,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,rare,individual,1,0.011431124999944586,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,rare,individual,2,0.011618645999988075,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,rare,node,0,0.011872388997289818,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,rare,node,1,0.01150943299944629,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,rare,node,2,0.01160151899966877,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,rare,edge,0,0.011911966001207475,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,rare,edge,1,0.013273937001940794,30000,15000,63287,75586,4108,42794,0 +genetic_value,100,rare,edge,2,0.011776721999922302,30000,15000,63287,75586,4108,42794,0 +sim_trait,1000,uniform,,0,0.030374293000932084,30000,15000,63287,75586,4108,42794,0 +sim_trait,1000,uniform,,1,0.027148326000315137,30000,15000,63287,75586,4108,42794,0 +sim_trait,1000,uniform,,2,0.02632547500252258,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,uniform,individual,0,0.11637149899979704,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,uniform,individual,1,0.11538826600008179,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,uniform,individual,2,0.11514968200208386,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,uniform,node,0,0.10923090399955981,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,uniform,node,1,0.10929654799838318,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,uniform,node,2,0.10818059000303037,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,uniform,edge,0,0.105416005000734,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,uniform,edge,1,0.10674772900165408,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,uniform,edge,2,0.11216579600295518,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,rare,individual,0,0.011900954999873647,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,rare,individual,1,0.011856531997182174,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,rare,individual,2,0.012914764000015566,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,rare,node,0,0.014092779001657618,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,rare,node,1,0.012466464999306481,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,rare,node,2,0.012185994000901701,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,rare,edge,0,0.012226915001519956,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,rare,edge,1,0.01296479300071951,30000,15000,63287,75586,4108,42794,0 +genetic_value,1000,rare,edge,2,0.013875839998945594,30000,15000,63287,75586,4108,42794,0 +sim_trait,10000,uniform,,0,0.22894942499988247,30000,15000,63287,75586,4108,42794,0 +sim_trait,10000,uniform,,1,0.23049697800161084,30000,15000,63287,75586,4108,42794,0 +sim_trait,10000,uniform,,2,0.22691678800038062,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,uniform,individual,0,1.005928552000114,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,uniform,individual,1,1.0046196410003176,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,uniform,individual,2,1.0013641219993588,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,uniform,node,0,0.9530282609994174,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,uniform,node,1,0.9479161859999294,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,uniform,node,2,0.946285923000687,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,uniform,edge,0,0.924820839001768,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,uniform,edge,1,0.9280050209999899,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,uniform,edge,2,0.923540199000854,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,rare,individual,0,0.01698478600155795,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,rare,individual,1,0.016901232997042825,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,rare,individual,2,0.017670743000053335,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,rare,node,0,0.018905965000158176,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,rare,node,1,0.017895351000333903,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,rare,node,2,0.01682795199667453,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,rare,edge,0,0.016792753998743137,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,rare,edge,1,0.020035887999256374,30000,15000,63287,75586,4108,42794,0 +genetic_value,10000,rare,edge,2,0.017466116001742193,30000,15000,63287,75586,4108,42794,0 diff --git a/benchmarks/benchmark_genetic_value.py b/benchmarks/benchmark_genetic_value.py index 670e77e..cc78f17 100644 --- a/benchmarks/benchmark_genetic_value.py +++ b/benchmarks/benchmark_genetic_value.py @@ -173,6 +173,8 @@ def warm_up(ts, model, levels, counters): trait_df = tstrait.sim_trait(ts, model=model, num_causal=1, random_seed=1) for level in levels: tstrait.genetic_value(ts, trait_df, level=level) + # The threaded kernel is a specialisation of its own to compile. + tstrait.genetic_value(ts, trait_df, level=levels[0], num_threads=1) if counters: # A kernel of its own to compile, so only pay for it when it is used. count_work(ts, _check_trait_df(ts, trait_df)) @@ -337,7 +339,11 @@ def run_benchmark(ts, args): ) for level in args.levels: call = functools.partial( - tstrait.genetic_value, ts, trait_df, level=level + tstrait.genetic_value, + ts, + trait_df, + level=level, + num_threads=args.num_threads, ) if args.memory: _, memory[(num_causal, selection, level)] = peak_memory(call) @@ -395,7 +401,12 @@ def summarise(rows, counts, memory, completed, ts, args): best[key] = min(best.get(key, seconds), seconds) print(f"\n{describe(ts)}") - print(f"Minimum of {args.replicates} replicates\n") + threads = ( + "on the calling thread" + if args.num_threads <= 0 + else (f"over {args.num_threads} worker threads") + ) + print(f"Minimum of {args.replicates} replicates, {threads}\n") columns = f"{'phase':<14} {'selection':<10} {'level':<11} {'num_causal':>10} " columns += f"{'seconds':>10} {'us/site':>10}" if args.counters: @@ -489,6 +500,7 @@ def write_csv(rows, ts, args): "num_edges", "num_trees", "num_sites", + "num_threads", ] ) for row in rows: @@ -501,6 +513,7 @@ def write_csv(rows, ts, args): ts.num_edges, ts.num_trees, ts.num_sites, + args.num_threads, ] ) print(f"\nWrote {args.output}") @@ -575,6 +588,15 @@ def parse_args(): "taken longer than this" ), ) + parser.add_argument( + "--num-threads", + type=int, + default=0, + help=( + "Worker threads to divide the causal sites between. 0, the " + "default, does the work on the calling thread" + ), + ) parser.add_argument( "--phases", action="store_true", diff --git a/docs/genetic.md b/docs/genetic.md index 6bbff26..fb409c5 100644 --- a/docs/genetic.md +++ b/docs/genetic.md @@ -49,6 +49,18 @@ trait_df : Trait dataframe that is described in [](effect_size_sim). There are some requirements for the trait dataframe input, which are described in [](req_trait_df). +level + +: The entity that genetic values are returned for, one of `"individual"`, `"node"` or `"edge"`. It + defaults to `"individual"`, and the other two are described in [](genetic_individual_node_edge_doc). + +num_threads + +: Number of worker threads to divide the causal sites between, defaulting to 0, which does the work + on the calling thread. Each thread holds arrays the length of the nodes, so how well it scales is + set by the size of the tree sequence rather than by the number of causal sites, and it is worth + measuring on your own data before relying on it. + The details of the parameters and how they influence the genetic value simulation are described in detail below. diff --git a/tests/test_genetic_value.py b/tests/test_genetic_value.py index 20c6d1e..2b2ba13 100644 --- a/tests/test_genetic_value.py +++ b/tests/test_genetic_value.py @@ -1,3 +1,6 @@ +import concurrent.futures +from unittest import mock + import msprime import numpy as np import pandas as pd @@ -5,7 +8,9 @@ import tskit import tstrait +from tstrait import jit from tstrait.base import _check_numeric_array +from tstrait.genetic_value import _check_trait_df, _GeneticValue from .data import ( all_trees_ts, @@ -1330,3 +1335,109 @@ def test_site_on_a_breakpoint(self): expected = naive_genetic_value(ts, trait_df, "node") result = tstrait.genetic_value(ts, trait_df, level="node") pd.testing.assert_frame_equal(result, expected, check_dtype=False) + + +class TestNumThreads: + """ + Dividing the rows between threads must not change the answer, whatever the + number of threads and however few rows there are to divide. + """ + + THREADS = [1, 2, 3, 5] + + @pytest.mark.parametrize("num_threads", THREADS) + @pytest.mark.parametrize("level", ["individual", "node", "edge"]) + @pytest.mark.parametrize("num_trait", [1, 3]) + def test_matches_sequential(self, num_threads, level, num_trait): + ts = multi_allelic_mutations( + msprime.sim_ancestry( + 20, sequence_length=10_000, recombination_rate=1e-4, random_seed=3 + ), + rate=1e-3, + seed=3, + ) + trait_df = random_trait_df(ts, num_trait, seed=5) + # The chunks are summed in a different order from the sequential + # accumulation, so this is equal to tolerance rather than bit for bit. + pd.testing.assert_frame_equal( + tstrait.genetic_value(ts, trait_df, level=level, num_threads=num_threads), + tstrait.genetic_value(ts, trait_df, level=level), + check_dtype=False, + ) + + @pytest.mark.parametrize("num_threads", THREADS) + def test_matches_reference(self, num_threads): + # Against the tree by tree reference rather than against the + # sequential path, so that the threads are not being checked only + # against themselves. + ts = binary_tree_seq() + trait_df = random_trait_df(ts, num_trait=2, seed=1) + for level in ("individual", "node", "edge"): + pd.testing.assert_frame_equal( + tstrait.genetic_value( + ts, trait_df, level=level, num_threads=num_threads + ), + naive_genetic_value(ts, trait_df, level), + check_dtype=False, + ) + + @pytest.mark.parametrize("num_threads", [2, 8, 100]) + def test_more_threads_than_rows(self, num_threads): + # A thread with no rows would walk the trees for nothing, so the count + # is clamped to the rows there are. + ts = binary_tree_seq() + trait_df = random_trait_df(ts, num_trait=1, seed=1).iloc[:1] + assert len(trait_df) == 1 + pd.testing.assert_frame_equal( + tstrait.genetic_value(ts, trait_df, num_threads=num_threads), + tstrait.genetic_value(ts, trait_df), + check_dtype=False, + ) + + def test_zero_is_sequential(self): + ts = binary_tree_seq() + trait_df = random_trait_df(ts, num_trait=1, seed=1) + pd.testing.assert_frame_equal( + tstrait.genetic_value(ts, trait_df, num_threads=0), + tstrait.genetic_value(ts, trait_df), + check_dtype=False, + ) + + def test_zero_starts_no_threads(self): + # The default has to stay a plain synchronous call, not a pool of one. + ts = binary_tree_seq() + trait_df = random_trait_df(ts, num_trait=1, seed=1) + with mock.patch.object( + concurrent.futures, "ThreadPoolExecutor", side_effect=AssertionError + ) as pool: + tstrait.genetic_value(ts, trait_df) + pool.assert_not_called() + + def test_bad_num_threads(self): + # Zero already means sequential, so a negative is a mistake rather + # than another way of asking for it. + ts = binary_tree_seq() + trait_df = random_trait_df(ts, num_trait=1, seed=1) + with pytest.raises(TypeError, match="num_threads must be an integer"): + tstrait.genetic_value(ts, trait_df, num_threads=1.5) + with pytest.raises( + ValueError, match="num_threads must be an integer not less than 0" + ): + tstrait.genetic_value(ts, trait_df, num_threads=-1) + + def test_keyword_only(self): + ts = binary_tree_seq() + trait_df = random_trait_df(ts, num_trait=1, seed=1) + with pytest.raises(TypeError): + tstrait.genetic_value(ts, trait_df, "node", 2) + + def test_exception_in_a_thread_is_raised(self): + # A thread failing must not be swallowed by the pool. + ts = binary_tree_seq() + trait_df = random_trait_df(ts, num_trait=1, seed=1) + genetic = _GeneticValue(ts, _check_trait_df(ts, trait_df)) + with mock.patch.object( + jit, "_descend_trees", side_effect=ValueError("kernel blew up") + ): + with pytest.raises(ValueError, match="kernel blew up"): + genetic._run("node", num_threads=2) diff --git a/tstrait/genetic_value.py b/tstrait/genetic_value.py index b794d1e..9d8ac76 100644 --- a/tstrait/genetic_value.py +++ b/tstrait/genetic_value.py @@ -1,10 +1,17 @@ +import concurrent.futures + import numpy as np import pandas as pd import tskit import tskit.jit.numba as tskit_numba from . import jit -from .base import _check_dataframe, _check_instance, _check_non_decreasing # noreorder +from .base import ( # noreorder + _check_dataframe, + _check_instance, + _check_int, + _check_non_decreasing, +) def _row_mutations(ts, trait_df): @@ -152,7 +159,10 @@ def _descend_arguments(self, level, output): """ Return the arguments to the descent kernel, with the contributions directed at nodes, individuals or edges according to ``level`` and - accumulated into ``output``. + accumulated into ``output``, over every row. + + A caller dividing the rows between threads overrides ``row_start`` and + ``row_stop`` in what it gets back. """ ts = self.ts return { @@ -170,9 +180,46 @@ def _descend_arguments(self, level, output): "node_output": self._node_output(level), "edge_level": level == "edge", "output": output, + "row_start": 0, + "row_stop": len(self.row_site), } - def _run(self, level): + def _run_threaded(self, level, N, num_threads): + """ + Accumulate the genetic values on a pool of threads, each taking a + contiguous range of the rows, and return the sum of what they produced. + + The rows are independent of each other, so a range needs nothing from + any other. Each gets an accumulator of its own rather than sharing one, + since the kernel adds into it and that is not atomic; summing them + afterwards is where the results stop being bit for bit what one thread + would have produced. + """ + # A range with no rows in it would walk the trees for nothing. + num_threads = min(num_threads, len(self.row_site)) + bounds = np.linspace(0, len(self.row_site), num_threads + 1).astype(int) + tables = [np.zeros((self.num_trait, N)) for _ in range(num_threads)] + # Built once, so that the threads share the arrays they only read + # rather than each making its own. All that varies is where a range + # starts and stops and which accumulator it adds into. + shared = self._descend_arguments(level, None) + + def descend(chunk): + return jit._descend_trees( + **{ + **shared, + "output": tables[chunk], + "row_start": bounds[chunk], + "row_stop": bounds[chunk + 1], + } + ) + + with concurrent.futures.ThreadPoolExecutor(max_workers=num_threads) as pool: + # Drain the results, so that an exception in a thread is raised here. + list(pool.map(descend, range(num_threads))) + return sum(tables[1:], tables[0]) + + def _run(self, level, num_threads=0): """ Computes genetic values of individuals, nodes, or edges depending on the value of "level" @@ -183,8 +230,11 @@ def _run(self, level): Dataframe with trait ID, [individual|node|edge] ID, and genetic value. """ N = self._output_size(level) - genetic_value_table = np.zeros((self.num_trait, N)) - jit._descend_trees(**self._descend_arguments(level, genetic_value_table)) + if num_threads <= 0: + genetic_value_table = np.zeros((self.num_trait, N)) + jit._descend_trees(**self._descend_arguments(level, genetic_value_table)) + else: + genetic_value_table = self._run_threaded(level, N, num_threads) return pd.DataFrame( { @@ -195,7 +245,7 @@ def _run(self, level): ) -def genetic_value(ts, trait_df, level="individual"): +def genetic_value(ts, trait_df, level="individual", *, num_threads=0): """ Compute genetic values for a tree sequence given a trait dataframe. @@ -209,6 +259,14 @@ def genetic_value(ts, trait_df, level="individual"): requirements. level : {"individual", "node", "edge"}, default "individual" The level (entity) at which genetic values are returned. + num_threads : int, default 0 + Number of worker threads to divide the causal sites between. The + default of 0 does the work on the calling thread. How well it scales + is set by the size of the tree sequence rather than by the number of + causal sites, because each thread holds arrays the length of the + nodes: a tree sequence small enough for those to stay in cache scales + almost linearly, and a larger one is limited by memory bandwidth well + before it runs out of cores. Returns ------- @@ -238,11 +296,12 @@ def genetic_value(ts, trait_df, level="individual"): raise ValueError("level must be one of 'individual', 'node', or 'edge'") if level == "individual" and ts.num_individuals == 0: raise ValueError("No individuals in the provided tree sequence dataset") + num_threads = _check_int(num_threads, "num_threads", minimum=0) trait_df = _check_trait_df(ts, trait_df) genetic = _GeneticValue(ts=ts, trait_df=trait_df) - genetic_result = genetic._run(level) + genetic_result = genetic._run(level, num_threads) return genetic_result diff --git a/tstrait/jit.py b/tstrait/jit.py index b6af559..9cdb52a 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -145,7 +145,7 @@ def _tree_roots(tree, samples, marked, mark, roots): return num_roots -@numba.njit +@numba.njit(nogil=True) def _descend_trees( numba_ts, edges_parent, @@ -161,6 +161,8 @@ def _descend_trees( node_output, edge_level, output, + row_start, + row_stop, ): """ Accumulate the genetic value of the selected causal sites in one pass over @@ -181,6 +183,15 @@ def _descend_trees( ``output`` is accumulated into, and the count of nodes visited is returned. That count is what the run time is proportional to, so the benchmark divides by it rather than keeping a copy of this loop that only counts. + + Only the rows in ``[row_start, row_stop)`` are accumulated, which is how the + work is divided between threads. The GIL is released, and everything a + thread writes to is allocated here, so the ranges need nothing of each + other: the rows are marked with their own index and the indexes of two + ranges cannot collide. Note that a range still walks the whole tree + sequence, since a tree is built from the one before it and there is no + seeking to the first tree a range wants, so the tree building costs a pass + per range rather than a pass per call. """ num_nodes = numba_ts.num_nodes tree = tree_state(num_nodes) @@ -195,13 +206,12 @@ def _descend_trees( stack = np.empty(num_nodes, dtype=np.int32) visits = 0 - row = 0 - num_rows = len(row_site) + row = row_start tree_index = numba_ts.tree_index() while tree_index.next(): _apply_edge_diffs(tree_index, edges_parent, edges_child, tree) site_stop = tree_index.site_range[1] - while row < num_rows and row_site[row] < site_stop: + while row < row_stop and row_site[row] < site_stop: start = pair_offset[row] stop = pair_offset[row + 1] # Every mutation at the site blocks the allele above it, whatever From 97f0e48202e9460795e40458b2b30b5b27a9757d Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Tue, 22 Sep 2026 14:24:03 +0100 Subject: [PATCH 14/19] Hand num_threads through sim_phenotype genetic_value is the only part of a phenotype simulation that threads, so sim_phenotype takes the same argument and passes it on, with the same default of 0 and the same meaning. It does not check the value itself, which is the convention the function already follows: h2 is not checked until sim_env runs, last of the three. So a bad num_threads is reported by genetic_value, after sim_trait has run. Verified end to end on 30,000 samples with 2,000 causal sites and a heritability of 0.3: the genetic values, the environmental noise and the phenotypes all agree with the sequential run to 7e-14. --- CHANGELOG.md | 2 ++ tests/test_simulate_phenotype.py | 34 ++++++++++++++++++++++++++++++++ tstrait/simulate_phenotype.py | 8 +++++++- 3 files changed, 43 insertions(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 934697c..9d48500 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,6 +13,8 @@ In development ### Performance +- `sim_phenotype` takes a `num_threads` argument and hands it to + `genetic_value`, which is the part of it that threads. - `genetic_value` takes a `num_threads` argument, dividing the causal sites between that many worker threads and defaulting to 0, which does the work on the calling thread. Each thread holds arrays the length of the nodes, so how diff --git a/tests/test_simulate_phenotype.py b/tests/test_simulate_phenotype.py index fac9ba8..03de962 100644 --- a/tests/test_simulate_phenotype.py +++ b/tests/test_simulate_phenotype.py @@ -348,3 +348,37 @@ def test_pleiotropy(self, sample_ts): var_array = grouped.var().values.T[0] np.testing.assert_almost_equal(mean_array, np.zeros(2), decimal=2) np.testing.assert_almost_equal(var_array, np.ones(2), decimal=2) + + +class TestNumThreads: + """ + sim_phenotype hands num_threads to genetic_value, which is the only part + of it that threads. + """ + + @pytest.mark.parametrize("num_threads", [1, 2, 5]) + def test_matches_sequential(self, sample_ts, sample_trait_model, num_threads): + # The same seed gives the same causal sites and the same environmental + # noise, so only the genetic values could differ, and those are summed + # over the chunks rather than accumulated in one go. + sequential = tstrait.sim_phenotype( + ts=sample_ts, model=sample_trait_model, num_causal=30, random_seed=7 + ) + threaded = tstrait.sim_phenotype( + ts=sample_ts, + model=sample_trait_model, + num_causal=30, + random_seed=7, + num_threads=num_threads, + ) + pd.testing.assert_frame_equal(threaded.trait, sequential.trait) + pd.testing.assert_frame_equal(threaded.phenotype, sequential.phenotype) + + def test_bad_num_threads(self, sample_ts, sample_trait_model): + with pytest.raises(TypeError, match="num_threads must be an integer"): + tstrait.sim_phenotype( + ts=sample_ts, + model=sample_trait_model, + num_causal=1, + num_threads=1.5, + ) diff --git a/tstrait/simulate_phenotype.py b/tstrait/simulate_phenotype.py index 92d47c8..185e998 100644 --- a/tstrait/simulate_phenotype.py +++ b/tstrait/simulate_phenotype.py @@ -41,6 +41,7 @@ def sim_phenotype( alpha=None, h2=None, random_seed=None, + num_threads=0, ): """ Simulate quantitative traits. @@ -66,6 +67,11 @@ def sim_phenotype( :param random_seed: Random seed of simulation. If None, simulation will be conducted randomly. :type random_seed: int + :param num_threads: Number of worker threads to divide the causal sites + between when computing genetic values. The default of 0 does the work + on the calling thread. Please see :func:`genetic_value` for what + determines how well it scales. + :type num_threads: int :returns: Dataclass object that includes phenotype and trait dataframe. :rtype: PhenotypeResult :raises ValueError: If the number of mutations in `ts` is smaller than `num_causal`. @@ -123,7 +129,7 @@ def sim_phenotype( alpha=alpha, random_seed=random_seed, ) - genetic_df = tstrait.genetic_value(ts=ts, trait_df=trait_df) + genetic_df = tstrait.genetic_value(ts=ts, trait_df=trait_df, num_threads=num_threads) phenotype_df = tstrait.sim_env(genetic_df=genetic_df, h2=h2, random_seed=random_seed) result = tstrait.PhenotypeResult(trait=trait_df, phenotype=phenotype_df) From 13c0f9fad56a8fcab0d06fd716c88f98d461c0ee Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Tue, 22 Sep 2026 14:27:49 +0100 Subject: [PATCH 15/19] Stop checking in a benchmark baseline The two CSVs were the small preset as measured on one machine, kept so that a change could be diffed against them. That is worth less than it looked: the timings in them only mean anything on the machine they came from, and anyone comparing against them from another one would be reading their own hardware rather than their own change. Take a baseline before a change and compare against that instead. The counts that --counters writes do mean the same thing anywhere, but they are cheap to regenerate and were only ever half of one of the files. --- benchmarks/README.md | 9 +-- benchmarks/baseline_small.csv | 85 -------------------------- benchmarks/baseline_small_counters.csv | 9 --- 3 files changed, 5 insertions(+), 98 deletions(-) delete mode 100644 benchmarks/baseline_small.csv delete mode 100644 benchmarks/baseline_small_counters.csv diff --git a/benchmarks/README.md b/benchmarks/README.md index f476906..862083d 100644 --- a/benchmarks/README.md +++ b/benchmarks/README.md @@ -189,10 +189,11 @@ Results are written to `_output/genetic_value.csv` in long format, one row per replicate, together with the dimensions of the tree sequence they were measured on; `--counters` writes a second file alongside it. `_output/` is gitignored. -`baseline_small.csv` and `baseline_small_counters.csv` are the `small` preset as -it stands, for diffing against. The timings in the first are -specific to the machine they were taken on; the counts in the second are not, -and are the part worth treating as a regression test. +Nothing is checked in to diff against. Timings are specific to the machine +they were taken on, so a baseline from someone else's is not worth much; take +one on yours before a change and compare against that. The counts `--counters` +writes are machine independent, and are the part that would mean the same +thing anywhere. ## `profile_genetic_value.py` diff --git a/benchmarks/baseline_small.csv b/benchmarks/baseline_small.csv deleted file mode 100644 index 9403205..0000000 --- a/benchmarks/baseline_small.csv +++ /dev/null @@ -1,85 +0,0 @@ -phase,num_causal,selection,level,replicate,seconds,num_samples,num_individuals,num_nodes,num_edges,num_trees,num_sites,num_threads -sim_trait,1,uniform,,0,0.0024442089998046868,30000,15000,63287,75586,4108,42794,0 -sim_trait,1,uniform,,1,0.00206001099650166,30000,15000,63287,75586,4108,42794,0 -sim_trait,1,uniform,,2,0.0021039189996372443,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,uniform,individual,0,0.01224575399828609,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,uniform,individual,1,0.012782063000486232,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,uniform,individual,2,0.013174051000532927,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,uniform,node,0,0.01399092200153973,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,uniform,node,1,0.012024119001580402,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,uniform,node,2,0.012036211999657098,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,uniform,edge,0,0.012211093002406415,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,uniform,edge,1,0.013050546000158647,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,uniform,edge,2,0.0123963600017305,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,rare,individual,0,0.01127879400155507,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,rare,individual,1,0.012140853999881074,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,rare,individual,2,0.012874144002125831,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,rare,node,0,0.015231686000333866,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,rare,node,1,0.012208834999910323,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,rare,node,2,0.011544229000719497,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,rare,edge,0,0.011383757002477068,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,rare,edge,1,0.012330576999374898,30000,15000,63287,75586,4108,42794,0 -genetic_value,1,rare,edge,2,0.011485152997920522,30000,15000,63287,75586,4108,42794,0 -sim_trait,100,uniform,,0,0.005982647002383601,30000,15000,63287,75586,4108,42794,0 -sim_trait,100,uniform,,1,0.005803553998703137,30000,15000,63287,75586,4108,42794,0 -sim_trait,100,uniform,,2,0.005617545000859536,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,uniform,individual,0,0.021615510999254184,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,uniform,individual,1,0.02126154999859864,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,uniform,individual,2,0.02122037899971474,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,uniform,node,0,0.020638850000977982,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,uniform,node,1,0.021018144998379285,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,uniform,node,2,0.02054931600287091,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,uniform,edge,0,0.020246869000402512,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,uniform,edge,1,0.020713948000775417,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,uniform,edge,2,0.020244728999387007,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,rare,individual,0,0.011508805000630673,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,rare,individual,1,0.011431124999944586,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,rare,individual,2,0.011618645999988075,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,rare,node,0,0.011872388997289818,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,rare,node,1,0.01150943299944629,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,rare,node,2,0.01160151899966877,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,rare,edge,0,0.011911966001207475,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,rare,edge,1,0.013273937001940794,30000,15000,63287,75586,4108,42794,0 -genetic_value,100,rare,edge,2,0.011776721999922302,30000,15000,63287,75586,4108,42794,0 -sim_trait,1000,uniform,,0,0.030374293000932084,30000,15000,63287,75586,4108,42794,0 -sim_trait,1000,uniform,,1,0.027148326000315137,30000,15000,63287,75586,4108,42794,0 -sim_trait,1000,uniform,,2,0.02632547500252258,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,uniform,individual,0,0.11637149899979704,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,uniform,individual,1,0.11538826600008179,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,uniform,individual,2,0.11514968200208386,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,uniform,node,0,0.10923090399955981,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,uniform,node,1,0.10929654799838318,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,uniform,node,2,0.10818059000303037,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,uniform,edge,0,0.105416005000734,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,uniform,edge,1,0.10674772900165408,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,uniform,edge,2,0.11216579600295518,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,rare,individual,0,0.011900954999873647,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,rare,individual,1,0.011856531997182174,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,rare,individual,2,0.012914764000015566,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,rare,node,0,0.014092779001657618,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,rare,node,1,0.012466464999306481,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,rare,node,2,0.012185994000901701,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,rare,edge,0,0.012226915001519956,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,rare,edge,1,0.01296479300071951,30000,15000,63287,75586,4108,42794,0 -genetic_value,1000,rare,edge,2,0.013875839998945594,30000,15000,63287,75586,4108,42794,0 -sim_trait,10000,uniform,,0,0.22894942499988247,30000,15000,63287,75586,4108,42794,0 -sim_trait,10000,uniform,,1,0.23049697800161084,30000,15000,63287,75586,4108,42794,0 -sim_trait,10000,uniform,,2,0.22691678800038062,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,uniform,individual,0,1.005928552000114,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,uniform,individual,1,1.0046196410003176,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,uniform,individual,2,1.0013641219993588,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,uniform,node,0,0.9530282609994174,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,uniform,node,1,0.9479161859999294,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,uniform,node,2,0.946285923000687,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,uniform,edge,0,0.924820839001768,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,uniform,edge,1,0.9280050209999899,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,uniform,edge,2,0.923540199000854,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,rare,individual,0,0.01698478600155795,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,rare,individual,1,0.016901232997042825,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,rare,individual,2,0.017670743000053335,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,rare,node,0,0.018905965000158176,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,rare,node,1,0.017895351000333903,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,rare,node,2,0.01682795199667453,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,rare,edge,0,0.016792753998743137,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,rare,edge,1,0.020035887999256374,30000,15000,63287,75586,4108,42794,0 -genetic_value,10000,rare,edge,2,0.017466116001742193,30000,15000,63287,75586,4108,42794,0 diff --git a/benchmarks/baseline_small_counters.csv b/benchmarks/baseline_small_counters.csv deleted file mode 100644 index 2bdefaf..0000000 --- a/benchmarks/baseline_small_counters.csv +++ /dev/null @@ -1,9 +0,0 @@ -num_causal,selection,rows,visits,num_nodes,num_edges -1,uniform,1,33819,63287,75586 -1,rare,1,11,63287,75586 -100,uniform,100,514067,63287,75586 -100,rare,100,1554,63287,75586 -1000,uniform,1000,5501454,63287,75586 -1000,rare,1000,12230,63287,75586 -10000,uniform,10000,52900059,63287,75586 -10000,rare,10000,131744,63287,75586 From a0d90d74512b685e873c214e3a64158d8efdbc7d Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Tue, 22 Sep 2026 14:39:25 +0100 Subject: [PATCH 16/19] Correct what the genetic value speedup is attributed to The changelog said genetic_value now descends from the mutations of each causal site "instead of making a pass over every node for each of them". The tree by tree implementation descended from those same mutations, by the same stack walk over left_child and right_sib, pruning at the same mutations. The contrast was wrong, and it took credit for something that was never the difference. What it did per causal site, and this does not, is everything around the descent: constructing a Python Site object and a dict of its mutations, allocating three arrays the length of the nodes, and adding a dense vector into the accumulator. None of that depends on how many nodes the causal allele reaches. Holding the causal sites to fewer than two carrier nodes each, so the descent is negligible either way, and growing the tree sequence: nodes carriers/site old us/site new us/site 4,870 1.7 18.6 0.50 17,143 1.9 22.2 0.57 65,385 1.4 46.7 0.20 257,614 1.7 155.6 0.05 Fifty three times the nodes, the same carriers, and eight times the cost a site: about 16us of Python plus half a nanosecond a node, per causal site, for a descent that touched under two of them. The two headline figures were also stale. Over five replicates the rare trait is 211 to 227 times faster rather than 180, and the uniform one 1.09 to 1.17 times rather than 1.2, on a measurement with a 7% run to run spread, so say close to unchanged for that one and why: uniformly drawn causal sites spend their time in the descent, and the descent is the same. --- CHANGELOG.md | 22 ++++++++++++++-------- 1 file changed, 14 insertions(+), 8 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 9d48500..249a6e6 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -24,14 +24,20 @@ In development where the arrays no longer fit in cache, the same trait is 1.6 times faster. Threads do not pay for themselves on a trait whose causal sites are few or rare, since each of them walks the trees. -- `genetic_value` descends from the mutations of each causal site instead of - making a pass over every node for each of them, so its cost is the number of - nodes carrying a causal allele rather than the number of causal sites times - the size of the tree sequence. On 100,000 samples, a trait with 100,000 rare - causal sites is around 180 times faster, and one whose causal sites are drawn - uniformly, and so are mostly common variants, around 1.2 times faster. Every - trait of a tree sequence is computed in one pass over the trees rather than - one pass each. +- `genetic_value` accumulates every causal site of every trait in one pass over + the trees, into a single output array. Descending from a causal site's + mutations is what it always did and is unchanged; what has gone is the work + around it, which was repeated for every causal site: building Python objects + for that site's mutations, and allocating and accumulating arrays the length + of the nodes. None of that depended on how many nodes the causal allele + actually reached. Measured on causal sites reaching fewer than two nodes + each, so that the descent is negligible either way, the cost per causal site + ran from 19us on a tree sequence of 4,870 nodes to 156us on one of 257,614; + it is now flat. On 100,000 samples a trait with 100,000 rare causal sites is + over 200 times faster. One whose causal sites are drawn uniformly, and so are + mostly common variants, is close to unchanged: a little over 1.1 times + faster, against a run to run spread of 7% on a measurement that size. What + those sites cost has always been the descent, and the descent is the same. ### Breaking changes From bb2dbeb7415606dc13132d2bbda79d51ec1fd023 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Tue, 22 Sep 2026 14:47:30 +0100 Subject: [PATCH 17/19] Say why uniformly drawn causal sites cost so much more The README asserted that uniform and rare differ by two orders of magnitude without saying why, and the obvious reasoning gives the wrong answer: most sites are rare, so most of a uniform draw is rare too, and one might expect the average to follow the majority. It does not. Of 3,000 sites drawn uniformly from the small preset, 79% sit below a frequency of 0.1 and the median one reaches 279 of the 63,287 nodes, but the mean reaches 5,345, and the rarest half of the sites account for half a percent of the work. band sites nodes each share of work 0 - 1e-4 14% 2 0.0% 1e-4 - 1e-3 21% 22 0.1% 1e-3 - 1e-2 24% 243 1.1% 1e-2 - 1e-1 21% 2,238 8.6% 1e-1 - 1 21% 23,073 90.2% The two columns are reciprocal. A coalescent frequency spectrum has density 1/f, so every decade of frequency holds about the same number of sites, and the nodes an allele reaches is its frequency times the tree, so every decade costs ten times the one below it. Equal counts against ten times the cost makes the total a geometric series, and the top decade is nearly all of it. --- benchmarks/README.md | 32 +++++++++++++++++++++++++++----- 1 file changed, 27 insertions(+), 5 deletions(-) diff --git a/benchmarks/README.md b/benchmarks/README.md index 862083d..7601c1d 100644 --- a/benchmarks/README.md +++ b/benchmarks/README.md @@ -16,11 +16,33 @@ uv run --group test benchmarks/benchmark_genetic_value.py causal sites it carries, at around 20ns a node once the tree building is amortised, so its cost is the number of nodes that carry a causal allele and the allele frequency of the causal sites matters more than how many there are. -`--selections` therefore draws them two ways: `uniform` over all sites, which -the common variants in the tail of the frequency spectrum dominate, and `rare`, -restricted to sites below `--rare-threshold`. The two differ by two orders of -magnitude and behave differently, so a single number for "the cost of a causal -site" is meaningless without saying which. +`--selections` therefore draws them two ways: `uniform` over all sites, and +`rare`, restricted to sites below `--rare-threshold`. The two differ by two +orders of magnitude, so a single number for "the cost of a causal site" is +meaningless without saying which. + +It is worth being clear about why, because the obvious reasoning gives the +wrong answer. Most sites are rare, so most of a uniform draw is rare too: of +3,000 drawn from the `small` preset, 79% sit below a frequency of 0.1 and the +median one reaches 279 of the 63,287 nodes. But the mean reaches 5,345, 19 +times the median, and the rarest half of the sites account for half a percent +of the work. Sorting the same sites by frequency: + +| band | sites | nodes reached, each | share of the work | +|---|---|---|---| +| 0 – 1e-4 | 14% | 2 | 0.0% | +| 1e-4 – 1e-3 | 21% | 22 | 0.1% | +| 1e-3 – 1e-2 | 24% | 243 | 1.1% | +| 1e-2 – 1e-1 | 21% | 2,238 | 8.6% | +| 1e-1 – 1 | 21% | 23,073 | 90.2% | + +The two columns are reciprocal. A coalescent frequency spectrum has density +1/f, so each decade of frequency holds about the same number of sites; the +nodes an allele reaches is its frequency times the tree, so each decade costs +ten times the one below. Equal counts against ten times the cost makes the work +a geometric series, and the top decade is nearly all of it. One site at a +frequency of 0.5 costs what ten thousand singletons cost, and there are only +about ten times fewer of them. `sim_trait` is timed separately, because it has a per-site Python loop of its own that we do not want folded into the `genetic_value` numbers. The numba From 41bfb47427c43d111e71211911fed802ff106995 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Tue, 22 Sep 2026 14:55:40 +0100 Subject: [PATCH 18/19] Cut the results out of the benchmark README It had grown into a report: timing tables for both presets, a thread scaling table, sample counter and perf output, peak memory figures, and a derivation of why uniformly drawn causal sites cost more than rare ones. None of that is needed to run the benchmarks, all of it goes stale the moment anything changes, and it is already in the commits that measured it. What is left is how to run the two scripts and what the options do, with numbers only where they are a property of the code rather than of a measurement: the two presets and roughly how long each takes, the mutation rate and why it is what it is. 274 lines down to 126. The gotchas stay, since they are about getting a measurement right rather than about any particular one, with a third added for threads: they are worth having only when the causal sites are many and not rare. --- benchmarks/README.md | 268 ++++++++++--------------------------------- 1 file changed, 60 insertions(+), 208 deletions(-) diff --git a/benchmarks/README.md b/benchmarks/README.md index 7601c1d..dc2ead2 100644 --- a/benchmarks/README.md +++ b/benchmarks/README.md @@ -1,7 +1,8 @@ # Benchmarks Performance benchmarks for tstrait. These are not run in CI; they are here so -that performance work can be repeated and compared. +that performance work can be repeated and compared. Results belong in the +commit that measured them, not in this file. ## `benchmark_genetic_value.py` @@ -12,210 +13,73 @@ a trait. uv run --group test benchmarks/benchmark_genetic_value.py ``` -`genetic_value` builds each tree in turn and descends from the mutations of the -causal sites it carries, at around 20ns a node once the tree building is -amortised, so its cost is the number of nodes that carry a causal allele and the -allele frequency of the causal sites matters more than how many there are. -`--selections` therefore draws them two ways: `uniform` over all sites, and -`rare`, restricted to sites below `--rare-threshold`. The two differ by two -orders of magnitude, so a single number for "the cost of a causal site" is -meaningless without saying which. +Every option has a default and `--help` lists them all. The ones that change +what is being measured rather than how long it takes: -It is worth being clear about why, because the obvious reasoning gives the -wrong answer. Most sites are rare, so most of a uniform draw is rare too: of -3,000 drawn from the `small` preset, 79% sit below a frequency of 0.1 and the -median one reaches 279 of the 63,287 nodes. But the mean reaches 5,345, 19 -times the median, and the rarest half of the sites account for half a percent -of the work. Sorting the same sites by frequency: +`--preset {small,large}` +: Fills in `--samples`, `--length` and `--num-causal`. `small` is the default + and takes about 25 seconds; `large` takes about six minutes. Either is + simulated on first use and cached in `_output/`, keyed by the parameters. -| band | sites | nodes reached, each | share of the work | -|---|---|---|---| -| 0 – 1e-4 | 14% | 2 | 0.0% | -| 1e-4 – 1e-3 | 21% | 22 | 0.1% | -| 1e-3 – 1e-2 | 24% | 243 | 1.1% | -| 1e-2 – 1e-1 | 21% | 2,238 | 8.6% | -| 1e-1 – 1 | 21% | 23,073 | 90.2% | +`--selections {uniform,rare}` +: How the causal sites are drawn: `uniform` over all sites, or `rare`, + restricted to those below `--rare-threshold`. The cost of a causal site is + the number of nodes carrying its allele, so these differ by two orders of + magnitude and a result quoted without saying which is meaningless. -The two columns are reciprocal. A coalescent frequency spectrum has density -1/f, so each decade of frequency holds about the same number of sites; the -nodes an allele reaches is its frequency times the tree, so each decade costs -ten times the one below. Equal counts against ten times the cost makes the work -a geometric series, and the top decade is nearly all of it. One site at a -frequency of 0.5 costs what ten thousand singletons cost, and there are only -about ten times fewer of them. +`--levels {individual,node,edge}` +: Which of the three `genetic_value` levels to time. + +`--num-threads` +: Worker threads to divide the causal sites between. 0, the default, does the + work on the calling thread. + +`--replicates`, `--max-seconds` +: How many times each cell is timed, and the budget after which the larger + numbers of causal sites are skipped. `sim_trait` is timed separately, because it has a per-site Python loop of its -own that we do not want folded into the `genetic_value` numbers. The numba +own that should not be folded into the `genetic_value` numbers, and the numba kernel is compiled by a warm up call that is not timed. The mutation rate defaults to 1e-7, ten times the human rate, so that there are -enough sites in a genome short enough to simulate quickly. The mutation rate -does not affect the allele frequency spectrum, so the causal sites are as -weakly causal as they would be under a realistic rate; only the number of sites -per tree is inflated. - -### Presets - -`--preset small` is the default and takes about 25 seconds; `--preset large` is -the tree sequence this work was written against and takes about six minutes, -with its longest single call at 33 seconds. - -| | small | large | -|---|---|---| -| samples | 30,000 | 100,000 | -| sequence length | 1Mb | 5Mb | -| nodes | 63,287 | 217,091 | -| edges | 75,586 | 285,979 | -| trees | 4,108 | 22,523 | -| sites | 42,794 | 235,458 | -| causal sites | 1 … 10,000 | 1 … 100,000 | -| cached tree sequence | 6.5MB | 27MB | - -A preset only fills in `--samples`, `--length` and `--num-causal`, each of which -still overrides it when given explicitly. The tree sequence is cached in -`_output/` under a name built from the simulation parameters, so `small` -simulates itself in 0.2s the first time and is loaded after that. - -`small` is the default because iterating against a six minute grid is not -practical. It reproduces the patterns the large one shows, at `level="node"`: - -| num_causal | small uniform | small rare | ratio | large uniform | large rare | ratio | -|---|---|---|---|---|---|---| -| 1 | 0.013s | 0.013s | 1.0 | 0.062s | 0.061s | 1.0 | -| 100 | 0.024s | 0.014s | 1.8 | 0.118s | 0.067s | 1.8 | -| 1,000 | 0.125s | 0.017s | 7.4 | 0.443s | 0.072s | 6.1 | -| 10,000 | 1.088s | 0.023s | 47 | 3.848s | 0.091s | 42 | -| 100,000 | — | — | | 32.589s | 0.216s | 151 | - -Both show a flat per call floor, a uniform µs/site that falls to a plateau from -1,000 causal sites upwards, a rare µs/site that is still falling at the top of -the grid, a uniform to rare ratio that grows with the number of causal sites, -and parity between the three levels. The uniform plateau is 109µs/site against -326µs/site, a factor of 3.0 on a node count ratio of 3.4. - -What makes the small preset a fair substitute is not the timings but the -distribution underneath them: the fraction of the nodes that a causal site's -effect reaches, which is what the descent costs. `--structure` measures it. - -| | mean carrier fraction | median | -|---|---|---| -| small | 7.19% | 0.406% | -| large | 5.99% | 0.184% | - -Those are 150 sampled sites from a distribution with a heavy tail, so they -agree about as well as they can. Check this again before trusting a new preset. - -The grid stops at 10,000 causal sites on `small` because the rare pool is only -about 37% of its 42,794 sites. Raising the mutation rate to lift the ceiling -was tried and rejected: it takes `sites/nodes` from 0.68 to about 2.7 against -the large preset's 1.08, and the per-call setup floor grows with the number of -sites, so the low end of the curve stops being comparable. - -### Threads - -`--num-threads` divides the causal sites between that many worker threads, and -defaults to 0, which does the work on the calling thread. The rows of the trait -dataframe are independent, so a thread takes a contiguous range of them and its -own accumulator, and the accumulators are summed at the end. - -Minimum of two replicates at `level="node"`, on four cores: +enough sites in a genome short enough to simulate quickly. It does not affect +the allele frequency spectrum, so the causal sites are as weakly causal as they +would be under a realistic rate; only the number of sites per tree is inflated. -| preset | selection | num_causal | 1 thread | 2 | 3 | 4 | -|---|---|---|---|---|---|---| -| small | uniform | 10,000 | 980ms | 1.94x | 2.86x | **3.41x** | -| small | uniform | 300 | 41ms | 1.36x | 1.66x | 1.69x | -| small | rare | 10,000 | 17ms | 0.94x | 0.91x | 0.89x | -| small | rare | 300 | 12ms | 0.82x | 0.78x | 0.77x | -| large | uniform | 10,000 | 3338ms | 1.47x | 1.64x | 1.64x | -| large | rare | 10,000 | 76ms | 1.01x | 1.00x | 0.99x | +### Modes that say why, not just how long -Two things in that table are worth understanding before reading a number off it -as a bug. +Each roughly doubles the run, `--phases` most of all. -**A bigger trait parallelises; a bigger tree sequence does not.** The kernel -holds about eleven arrays the length of the nodes per thread, 49 bytes a node. -That is 3.1MB on `small`, so four threads fit inside this machine's 16MB L3 and -scale nearly linearly, and 10.6MB on `large`, where four threads want 42MB and -are limited by memory bandwidth instead. Raising the causal sites from 10,000 to -100,000 on `large` moves four threads from 1.50x only to 1.61x, and shrinking -the two stamp arrays to `int32` gets 1.74x, so this is cache capacity rather -than anything the division of work can fix. A machine with more L3 or more -memory channels should do better on `large`. +`--phases` +: Times `_check_trait_df`, `_GeneticValue.__init__`, the kernel and the output + dataframe separately, which is what tells an algorithmic win from a setup one. -**Threads do not pay for themselves on a rare or a small trait.** A thread walks -the whole tree sequence whatever range of rows it takes, because a tree is built -from the one before it and there is no seeking to the first tree a range wants. -That pass is 2.2ms on `small` and 12.4ms on `large`, so it is nothing against a -second of descending and everything against seventeen milliseconds of it. +`--counters` +: Reports the nodes the descent reached. The run time is proportional to that, + so seconds over visits is the constant an optimisation has to move. perf + cannot attribute time inside the kernel (see below), so this is the way to + say where the time goes. -### What else it measures +`--structure` +: Reports the shape of the tree sequence and the distribution of how many nodes + a causal site reaches. This is what makes one preset a fair substitute for + another, so check it before trusting a new one. -Wall time on its own does not say why a configuration is slow. Four optional -modes say more. `--phases` is the expensive one, roughly doubling the run -because it times the same work again a piece at a time; the other three add -seconds. - -`--phases` times `_check_trait_df`, `_GeneticValue.__init__`, the kernel and -the output dataframe separately. The end to end -number stays as the headline. This is what tells an algorithmic win from a -setup win: setup barely grows with the causal sites, going from 8ms to 17ms -across the whole `small` grid and sitting near 100ms on `large`, so at one -causal site the public call is measuring almost nothing else, while at 10,000 -the kernel is 98% of it. Most of setup is `tskit.jit.numba.jitwrap`, which runs -three Python-speed `max(map(len, ...))` passes over the site and mutation -tables. The dataframe is about 1ms and is not worth thinking about. - -`--counters` reports the work the descent does. perf cannot attribute time to -source lines inside a numba kernel here (see below), so counting what the -kernel does and dividing is the way to say where the time goes. The kernel -returns the count itself rather than there being a second copy of the loop to -keep in step with the first. On `small`: - -``` -selection num_causal rows visits visits/row of num_nodes -uniform 1 1 33,819 33819.0 53.44% -uniform 10000 10,000 52,900,059 5290.0 8.36% -rare 1 1 11 11.0 0.02% -rare 10000 10,000 131,744 13.2 0.02% -``` - -`visits` is the number of nodes the descent reached, which is what the run time -is proportional to, so `ns/visit` in the summary table is the constant an -optimisation has to move. Two things fall out of the table: - -- `visits/row` as a fraction of the nodes is the carrier fraction that - `--structure` measures independently, and the two agree: 8.4% for uniformly - drawn causal sites against a measured mean of 7.2%, and 0.02% for rare ones - against a measured median of 0.4%. A single uniformly drawn site reaching - 53% is one draw from a distribution with a long tail. -- `ns/visit` is around 20ns once there are enough causal sites to amortise the - tree pass, and hundreds of nanoseconds below that. The pass is 2.2ms on - `small` and 12.4ms on `large`, so a rare trait of a hundred causal sites is - paying to build a tree sequence it barely touches. - -`--structure` reports the shape of the tree sequence and the carrier fraction -distribution described above. - -`--memory` reports the peak resident set size of each call. VmHWM never falls, -so it is reset before each call by writing to `/proc/self/clear_refs`; on a -kernel without that the column reads `unavailable`. - -The working set is a fixed handful of arrays the length of the nodes however -many causal sites there are: at 100,000 uniformly drawn causal sites on `large` -the peak is 0.42GB, of which 0.01GB is what the call added over the tree -sequence and the interpreter. +`--memory` +: Peak resident set size per call. VmHWM never falls, so it is reset before + each call by writing to `/proc/self/clear_refs`; without that the column + reads `unavailable`. ### Output -Results are written to `_output/genetic_value.csv` in long format, one row per -replicate, together with the dimensions of the tree sequence they were measured -on; `--counters` writes a second file alongside it. `_output/` is gitignored. +Long format CSV to `_output/genetic_value.csv`, one row per replicate, stamped +with the dimensions of the tree sequence; `--counters` and `--memory` write +files alongside it. `_output/` is gitignored. -Nothing is checked in to diff against. Timings are specific to the machine -they were taken on, so a baseline from someone else's is not worth much; take -one on yours before a change and compare against that. The counts `--counters` -writes are machine independent, and are the part that would mean the same -thing anywhere. +Nothing is checked in to diff against. Timings only mean anything on the +machine they came from, so take a baseline on yours before a change and compare +against that. The counts `--counters` writes are machine independent. ## `profile_genetic_value.py` @@ -227,10 +91,8 @@ uv run --group test benchmarks/profile_genetic_value.py --mode python uv run --group test benchmarks/profile_genetic_value.py --mode kernel ``` -`--mode python` is cProfile around the public call. It is the only way to see -the setup, and it shows `jitwrap` and its `builtins.max` rows plainly, which is -now nearly all of what setup costs. The kernel appears in it as one opaque -dispatcher frame. +`--mode python` is cProfile around the public call, which is the only way to +see the setup. The kernel appears in it as one opaque dispatcher frame. `--mode kernel` sets one cell up and runs only the kernel, so that a sampling profile is not swamped by simulating effect sizes for the whole site pool. Run @@ -243,27 +105,14 @@ Two things to know before reading a perf profile of this code. `PerfJITEventListener`, so nothing writes a `/tmp/perf-.map` or a jitdump and the kernel shows up as raw addresses under `[JIT]`. Use `--counters` for attribution inside the kernel. What perf does give is the split between the -kernel, the interpreter, LLVM compilation and the numba runtime: - -``` -59.19% [JIT] tid 23771 <- the kernel -19.24% python3.11 <- setup -11.78% libllvmlite.so <- compilation, not work - 4.14% [kernel.kallsyms] - 1.75% libc.so.6 -``` - -The numba runtime does not appear at all: the descent allocates nothing per -node, so there is no `_helperlib` row. - -That is `--repeats 10`; the setup is a fixed few seconds, so raise `--repeats` -until the `[JIT]` share stops moving. +kernel, the interpreter and LLVM compilation. Setup is a fixed few seconds of +that, so raise `--repeats` until the `[JIT]` share stops moving. **Do not set `NUMBA_ENABLE_PROFILING=1`.** It is the documented way to profile numba and is wrong here: it would only help through the listener llvmlite does not have, and it defaults `NUMBA_DEBUGINFO` to 1, which changes the generated -code and measured the kernel over 60% slower. perf finds the JIT mappings by -itself, so the recipe does not set it. +code and measures a slower kernel than the one that runs. perf finds the JIT +mappings by itself. ## Gotchas @@ -272,3 +121,6 @@ itself, so the recipe does not set it. first call on a given tree sequence pays for all of them. The warm up call covers numba compilation but not this. - The summary takes the minimum over replicates, not the mean. +- Threads are worth having only when the causal sites are many and not rare; + each thread walks the whole tree sequence and holds arrays the length of the + nodes, so a small or rare trait goes slower on more of them. From 5b1334621839deec29a104bf35421d3e696cb1a1 Mon Sep 17 00:00:00 2001 From: Jerome Kelleher Date: Tue, 22 Sep 2026 14:58:14 +0100 Subject: [PATCH 19/19] Cut the changelog entries down to the house style The Performance section had three entries running to twenty five lines, explaining what changed inside the implementation and carrying the measurements that justified it. Everything around it is one or two lines ending in a pr reference. Drop the section and fold its content into Highlights in that style, keeping the comparisons to the two numbers a reader might act on: over 200 times faster on rare causal sites, close to unchanged on common ones, and up to 3.4 times on four threads. Cite {pr}`194`. --- CHANGELOG.md | 37 +++++++++---------------------------- 1 file changed, 9 insertions(+), 28 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 249a6e6..3a81a89 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,34 +10,15 @@ In development and `level="edge"` to return genetic values for the corresponding entities. {pr}`189` - Added `edge_effect` to compute introduced effects on edges {pr}`189` - -### Performance - -- `sim_phenotype` takes a `num_threads` argument and hands it to - `genetic_value`, which is the part of it that threads. -- `genetic_value` takes a `num_threads` argument, dividing the causal sites - between that many worker threads and defaulting to 0, which does the work on - the calling thread. Each thread holds arrays the length of the nodes, so how - well it scales is set by the size of the tree sequence rather than by the - number of causal sites: on 30,000 samples a trait of 10,000 uniformly drawn - causal sites is 3.4 times faster on four threads, while on 100,000 samples, - where the arrays no longer fit in cache, the same trait is 1.6 times faster. - Threads do not pay for themselves on a trait whose causal sites are few or - rare, since each of them walks the trees. -- `genetic_value` accumulates every causal site of every trait in one pass over - the trees, into a single output array. Descending from a causal site's - mutations is what it always did and is unchanged; what has gone is the work - around it, which was repeated for every causal site: building Python objects - for that site's mutations, and allocating and accumulating arrays the length - of the nodes. None of that depended on how many nodes the causal allele - actually reached. Measured on causal sites reaching fewer than two nodes - each, so that the descent is negligible either way, the cost per causal site - ran from 19us on a tree sequence of 4,870 nodes to 156us on one of 257,614; - it is now flat. On 100,000 samples a trait with 100,000 rare causal sites is - over 200 times faster. One whose causal sites are drawn uniformly, and so are - mostly common variants, is close to unchanged: a little over 1.1 times - faster, against a run to run spread of 7% on a measurement that size. What - those sites cost has always been the descent, and the descent is the same. +- `genetic_value` computes every causal site of every trait in one pass over + the trees, instead of taking each causal site on its own. Traits with rare + causal sites are over 200 times faster; traits whose causal sites are mostly + common variants are close to unchanged {pr}`194` +- `genetic_value` and `sim_phenotype` take a `num_threads` argument, dividing + the causal sites between that many worker threads. The default of 0 does the + work on the calling thread. Up to 3.4 times faster on four threads, and less + on a tree sequence big enough that the per-thread arrays leave cache + {pr}`194` ### Breaking changes