diff --git a/CHANGELOG.md b/CHANGELOG.md index b032bda..3a81a89 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -10,6 +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` +- `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 diff --git a/benchmarks/README.md b/benchmarks/README.md new file mode 100644 index 0000000..dc2ead2 --- /dev/null +++ b/benchmarks/README.md @@ -0,0 +1,126 @@ +# Benchmarks + +Performance benchmarks for tstrait. These are not run in CI; they are here so +that performance work can be repeated and compared. Results belong in the +commit that measured them, not in this file. + +## `benchmark_genetic_value.py` + +Measures how `tstrait.genetic_value` scales with the number of causal sites in +a trait. + +``` +uv run --group test benchmarks/benchmark_genetic_value.py +``` + +Every option has a default and `--help` lists them all. The ones that change +what is being measured rather than how long it takes: + +`--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. + +`--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. + +`--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 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. 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. + +### Modes that say why, not just how long + +Each roughly doubles the run, `--phases` most of all. + +`--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. + +`--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. + +`--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. + +`--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 + +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 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` + +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, 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 +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 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 measures a slower kernel than the one that runs. perf finds the JIT +mappings by itself. + +## Gotchas + +- 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. +- 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. diff --git a/benchmarks/benchmark_genetic_value.py b/benchmarks/benchmark_genetic_value.py new file mode 100644 index 0000000..cc78f17 --- /dev/null +++ b/benchmarks/benchmark_genetic_value.py @@ -0,0 +1,658 @@ +""" +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. 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. + +Run with ``uv run --group test benchmarks/benchmark_genetic_value.py``. +""" + +import argparse +import csv +import functools +import pathlib +import statistics +import sys +import time + +import msprime +import numpy as np +import pandas as pd +import tskit + +import tstrait +from tstrait import jit +from tstrait.genetic_value import _check_trait_df, _GeneticValue + +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], + }, +} + + +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. + + 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): + before = time.perf_counter() + result = func() + times.append(time.perf_counter() - before) + 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 + 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 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, counters): + """ + 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) + # 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)) + 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) + + size = genetic._output_size(level) + shape = (genetic.num_trait, size) + _, times = time_call( + # A fresh output array each time, since the kernel accumulates into it. + lambda: jit._descend_trees(**genetic._descend_arguments(level, np.zeros(shape))), + replicates, + ) + phases.append(("kernel", times)) + + 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(), + } + ) + + +COUNTERS = ["rows", "visits"] + + +def count_work(ts, trait_df): + """ + 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 descent serves all + three and only the array the contributions land in differs. + """ + 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): + """ + 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, 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, 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( + 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, "uniform", "", replicate, seconds)) + report(rows[-1]) + slowest = max(times) + for selection in args.selections: + rng = np.random.default_rng(args.seed) + trait_df = select_causal( + pool, selection, num_causal, args.rare_threshold, rng + ) + if trait_df is None: + print(f" too few {selection} sites for num_causal={num_causal}") + continue + if args.counters: + 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, + num_threads=args.num_threads, + ) + 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( + ( + "genetic_value", + num_causal, + selection, + level, + replicate, + seconds, + ) + ) + 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 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: + print( + f" a call took {slowest:.1f}s, over the {args.max_seconds:g}s " + f"budget: skipping {skipped}" + ) + break + return rows, counts, memory, completed + + +def report(row): + 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, 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 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: + key = (phase, selection, level, num_causal) + best[key] = min(best.get(key, seconds), seconds) + + print(f"\n{describe(ts)}") + 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: + columns += f" {'ns/visit':>9}" + print(columns) + print("-" * len(columns)) + phases = ["sim_trait", "genetic_value"] + if args.phases: + phases += ["check", "setup", "kernel", "frame"] + combinations = [("sim_trait", "uniform", "")] + combinations += [ + (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: + key = (phase, selection, level, num_causal) + if key not in best: + continue + seconds = best[key] + 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 descent have a per node cost. + visits = counts.get((num_causal, selection), {}).get("visits") + descends = phase in ("genetic_value", "kernel") + line += ( + f" {seconds / visits * 1e9:>9.2f}" + if descends and visits + else f" {'':>9}" + ) + print(line) + + if args.counters: + print() + header = f"{'selection':<10} {'num_causal':>10} " + header += " ".join(f"{name:>14}" for name in COUNTERS) + header += f" {'visits/row':>11} {'of num_nodes':>13}" + print(header) + print("-" * len(header)) + 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]:>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: + 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): + 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", + "selection", + "level", + "replicate", + "seconds", + "num_samples", + "num_individuals", + "num_nodes", + "num_edges", + "num_trees", + "num_sites", + "num_threads", + ] + ) + 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, + args.num_threads, + ] + ) + 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( + "--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("--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="+") + 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", + 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( + "--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", + 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" + ) + 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(): + 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, 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__": + main() diff --git a/benchmarks/profile_genetic_value.py b/benchmarks/profile_genetic_value.py new file mode 100644 index 0000000..971ab5e --- /dev/null +++ b/benchmarks/profile_genetic_value.py @@ -0,0 +1,223 @@ +""" +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 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. + +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. 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 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``. +""" + +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)) + 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._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 " + 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. Source lines inside the kernel are not available; use\n" + "# --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() 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 9a205d1..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, @@ -36,7 +41,6 @@ def sample_df(): "causal_allele": ["A", "A"], "effect_size": [0.1, 0.1], "trait_id": [0, 0], - "allele_freq": [0.2, 0.3], } ) @@ -1076,3 +1080,364 @@ 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) + + +class TestDescent: + """ + 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. + """ + + @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"], + } + ) + expected = 1.0 if derived_state == "A" else 0.0 + 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 + # 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"], + } + ) + 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, + # 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"], + } + ) + 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/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..4cbc63c 100644 --- a/tests/test_jit.py +++ b/tests/test_jit.py @@ -9,79 +9,127 @@ 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 works from the trees of a tree sequence rather than from 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 +ANCESTRAL_STATE = "A" -def kernel(func, param): - """ - Return either the compiled kernel or the Python it was written as. - """ - return func if param == "jit" else func.py_func +# 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", +] -@pytest.fixture(params=["jit", "nojit"]) -def node_genetic_value(request): - """ - Return a function computing the node genetic values for a tskit tree. - 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. +def kernel(func, param, monkeypatch): + """ + Return either the compiled kernel or the Python it was written as, in the + latter case untranslating the kernels it calls along with it. """ - func = kernel(jit._compute_nodes_genetic_value, request.param) + if param == "jit": + return func + for name in TREE_KERNELS: + monkeypatch.setattr(jit, name, getattr(jit, name).py_func) + return func.py_func - 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, - ) - return f +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 individual_genetic_value(request): +def node_genetic_value(request, monkeypatch): """ - Return a function accumulating node genetic values over the individuals of - a tree sequence. - """ - func = kernel(jit._accumulate_individual_values, request.param) + Return a function computing the node genetic values of a tree carrying one + causal site. - def f(ts, nodes_genetic_value): - return func( - np.asarray(nodes_genetic_value, dtype=float), - ts.nodes_individual, - ts.num_individuals, + 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._descend_trees, request.param, monkeypatch) + + 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) + # 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 -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 +142,18 @@ def multiroot_tree(): return tables.tree_sequence().first() -def empty_tree(): +def empty_ts(): """ - Return the tree of a tree sequence that has no nodes. + Return a tree sequence that has no nodes, but does have a causal site. """ - return tskit.TableCollection(sequence_length=1).tree_sequence().first() + tables = tskit.TableCollection(sequence_length=1) + tables.sites.add_row(position=0, ancestral_state=ANCESTRAL_STATE) + return tables.tree_sequence() class TestBalancedBinaryTree: """ - tskit.Tree.generate_balanced(4), in which node 7 is the virtual root:: + tskit.Tree.generate_balanced(4):: 6 +-+-+ @@ -117,63 +167,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 +247,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 +285,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 +331,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,102 +396,236 @@ 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: +class TestIndividualLevel: """ - 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. + 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 + +-+-+ + 4 5 + +++ +++ + 0 1 2 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_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. + 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() - np.testing.assert_array_equal(ts.nodes_individual, [1, 1, 2, 2, 0, 0, -1]) + site = one_site(ts.first(), [(4, "T")]) np.testing.assert_array_equal( - individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64]), [48, 3, 12] + 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_diff_ind_tree(self, individual_genetic_value): - # The same tree, with the leaves paired up the other way around. + 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() - 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] + 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, individual_genetic_value): - # Two triploids: nodes 0, 2 and 4, and nodes 1, 3 and 5. + 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: + # + # 7 + # +-+-+---+ + # | | | 6 + # | | | +-+-+ + # 0 1 2 3 4 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]) + site = one_site(ts.first(), [(6, "T")]) np.testing.assert_array_equal( - individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64, 128]), [21, 42] + 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_no_individuals(self, individual_genetic_value): - ts = tskit.Tree.generate_balanced(4).tree_sequence - assert ts.num_individuals == 0 + 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( - individual_genetic_value(ts, [1, 2, 4, 8, 16, 32, 64]), [] + node_genetic_value(site, causal_allele="T", level="individual"), expected ) -class TestNodeAndIndividualValues: +@numba.njit +def _walk_trees( + numba_ts, + edges_parent, + edges_child, + samples, + parent, + left_child, + right_sib, + node_edge, + roots, + num_roots, +): """ - 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:: + 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 - 6 - +-+-+ - 4 5 - +++ +++ - 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. +@pytest.fixture(params=["jit", "nojit"]) +def walk_trees(request, monkeypatch): """ + Return a function giving the per tree state arrays that _walk_trees records. + """ + driver = kernel(_walk_trees, request.param, monkeypatch) - def test_internal_node(self, node_genetic_value, individual_genetic_value): - ts = binary_tree() - value = node_genetic_value(ts.first(), [4]) - 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 f(ts): + return _walk(driver, ts) - def test_virtual_root(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]) - 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]) + return f - def test_triploid(self, node_genetic_value, individual_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: - # - # 7 - # +-+-+---+ - # | | | 6 - # | | | +-+-+ - # 0 1 2 3 4 5 - # - ts = triploid_tree() - value = node_genetic_value(ts.first(), [6]) - 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]) + +def _walk(driver, ts): + 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 = driver( + 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, walk_trees): + 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/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/genetic_value.py b/tstrait/genetic_value.py index 382eb61..9d8ac76 100644 --- a/tstrait/genetic_value.py +++ b/tstrait/genetic_value.py @@ -1,24 +1,67 @@ +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 _accumulate_edge_values(nodes_genetic_value, nodes_edge, num_nodes, num_edges): +def _row_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, returning the row and the mutation. + + 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. """ - 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 ) + 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] + return derived_state[mutation] == causal_allele, 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 = _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 _check_trait_df(ts, trait_df): """ @@ -47,6 +90,14 @@ 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 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 ---------- ts : tskit.TreeSequence @@ -59,33 +110,116 @@ 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.numba_ts = tskit_numba.jitwrap(ts) + + site_id = self.trait_df["site_id"].to_numpy() + self.trait_id = self.trait_df["trait_id"].to_numpy() + self.num_trait = np.max(self.trait_id) + 1 + + # 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 = 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, _ = _row_causal_allele( + ts, self.trait_df, pair_row, pair_mutation + ) + + def _output_size(self, level): + return { + "individual": self.ts.num_individuals, + "node": self.ts.num_nodes, + "edge": self.ts.num_edges, + }[level] - def _node_genetic_values(self, tree, site, causal_allele, effect_size): + def _node_output(self, level): """ - Returns a numpy array with node genetic values. + 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. """ - 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, - ) - 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, - ) + if level == "individual": + return self.ts.nodes_individual + return np.arange(self.ts.num_nodes, dtype=np.int32) - def _run(self, level): + 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``, 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 { + "numba_ts": self.numba_ts, + "edges_parent": ts.edges_parent, + "edges_child": ts.edges_child, + "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, + "row_start": 0, + "row_stop": len(self.row_site), + } + + 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" @@ -95,50 +229,23 @@ def _run(self, level): pandas.DataFrame Dataframe with trait ID, [individual|node|edge] ID, and genetic value. """ - - ts = self.ts - size_map = { - "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, - ) - if level == "individual": - genetic_value = jit._accumulate_individual_values( - genetic_value, 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 - - df = pd.DataFrame( + N = self._output_size(level) + 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( { - "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"): +def genetic_value(ts, trait_df, level="individual", *, num_threads=0): """ Compute genetic values for a tree sequence given a trait dataframe. @@ -152,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 ------- @@ -181,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 @@ -231,34 +347,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..9cdb52a 100644 --- a/tstrait/jit.py +++ b/tstrait/jit.py @@ -3,74 +3,259 @@ 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 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. 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 +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 + import numba import numpy as np import tskit +# 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 _compute_nodes_genetic_value( - left_child_array, - right_sib_array, - causal_nodes, - has_mutation, - effect_size, -): +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): """ - 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. + Attach ``child_node`` to ``parent_node`` as its rightmost child, through + ``edge``. """ - 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] + 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 _accumulate_individual_values( - nodes_genetic_value, nodes_individual, num_individuals +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 + + +@numba.njit(nogil=True) +def _descend_trees( + numba_ts, + edges_parent, + edges_child, + row_site, + row_trait, + row_effect, + row_ancestral, + pair_offset, + pair_node, + pair_carries, + samples, + node_output, + edge_level, + output, + row_start, + row_stop, ): """ - Accumulate the individual genetic values by summing their node - contributions. + 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. + + ``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. """ - 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 + 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) + + visits = 0 + 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 < 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 + # 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] + visits += 1 + 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 visits 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)