Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
22 changes: 18 additions & 4 deletions CLAUDE.md
Original file line number Diff line number Diff line change
Expand Up @@ -17,8 +17,6 @@ The package includes a C extension (`_tsinfer`) built from `lib/` sources via se
uv run pytest tests/ -v # Run all tests
uv run pytest tests/test_matching.py # Run a single test file
uv run pytest tests/test_matching.py::TestFoo::test_bar -v # Run a single test
uv run pytest --skip-slow # Skip slow tests

uv run ruff check --fix # Lint Python code (auto-fix)
uv run ruff format # Format Python code
```
Expand Down Expand Up @@ -81,7 +79,7 @@ Source in `lib/`. Three main classes exposed to Python:
- `AncestorBuilder` — builds inferred ancestors from genotype data
- `AncestorMatcher` — Li & Stephens HMM matching algorithm

When changes are made to the C library, ensure that the ``_tskit`` module is rebuilt
When changes are made to the C library, ensure that the ``_tsinfer`` module is rebuilt
before running Python tests.

Vendored dependencies in `lib/subprojects/`: tskit C library and kastore.
Expand All @@ -98,9 +96,22 @@ Sample VCZ → `infer_ancestors` → Ancestor VCZ → `match` → raw `tskit.Tre

- Do not be overly defensive - defend only against circumstances that can
occur within the current codebase.
- Do not make production code more complex for the sake of minimising
changes to the test suite. Simplicity and clarity of the production code
is imperative.
- Do not combine multiple complex operations in a single statement. Prefer
to keep a single operation per statement, and use intermediate variables
as a form of documentation. For example:
```python
# Bad — multiple operations in one expression
result = sorted(k for k, v in mapping.items() if v in set(x.name for x in sources))

# Good — intermediate variable makes intent clear
source_names = {x.name for x in sources}
result = sorted(k for k, v in mapping.items() if v in source_names)
```
- Prefer dataclasses over tuples when returning multiple values.
- Use explicit `None` comparisons: `if x is not None` not `if x`.
- Zarr v3 is now used (dependency: `zarr>=3`).
- Import all modules at the top of the file, not inside functions or methods.
- Prefer importing a module and using module.function instead of
using ``from module import function``. This applies to intra-package
Expand All @@ -112,6 +123,7 @@ Sample VCZ → `infer_ancestors` → Ancestor VCZ → `match` → raw `tskit.Tre
- When a parameter has a computed default derived from another parameter,
compute it once at the point of use (the leaf function), not at every
layer in the call chain. Pass `None` through intermediate layers.
- Zarr v3 is used (dependency: `zarr>=3`). Do not use Zarr v2 APIs.
- Use PEP 604 union syntax: `int | None`, not `Optional[int]`.
- One `logger = logging.getLogger(__name__)` per module at top level.

Expand All @@ -125,3 +137,5 @@ Sample VCZ → `infer_ancestors` → Ancestor VCZ → `match` → raw `tskit.Tre
- Test helpers are in `tests/helpers.py` (e.g., `make_sample_vcz`, `make_ancestor_vcz`)
- `tests/algorithm.py` contains Python reference implementations used to verify C code
- `msprime` is used to simulate test data
- Run the test suite with coverage before committing to ensure that new code is
fully covered by tests.
9 changes: 5 additions & 4 deletions docs/inference.md
Original file line number Diff line number Diff line change
Expand Up @@ -225,9 +225,10 @@ paths, "path compression".

Matching ancestors is dependent on the time allocated to each ancestor; an
ancestor can only copy from any older ancestor. For each ancestor,
we find the most likely path through older ancestors: that is the path that
maximises the product of the probabilities of recombination and mismatch
over all sites.
we find a copying path through older ancestors using a Li & Stephens
hidden Markov model. This determines, at each site, which older ancestor
the current ancestor copies from, allowing for both recombination
(switching between ancestors) and mismatch (mutations).

:::{todo}
Schematic of the ancestors copying process.
Expand Down Expand Up @@ -279,4 +280,4 @@ The final phase of a `tsinfer` inference consists of a number steps:
section
2. Describe the structure of the output tree sequences; how the
nodes are mapped, what the time values mean, etc.
:::
:::
2 changes: 1 addition & 1 deletion tests/test_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -159,7 +159,7 @@ def test_basic(self):
},
output="out.trees",
)
assert m.path_compression is True
assert m.path_compression is False
assert m.reference_ts is None
assert m.sources["cohort"].node_flags == 1
assert m.sources["cohort"].create_individuals is True
Expand Down
6 changes: 3 additions & 3 deletions tests/test_grouping.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,7 +29,7 @@
import numpy as np
import pytest

from tsinfer import grouping
from tsinfer import config, grouping


@dataclasses.dataclass
Expand Down Expand Up @@ -558,7 +558,7 @@ class TestAssignGroups:
"""Test that assign_groups takes ungrouped MatchJobs and assigns groups."""

def _make_job(self, haplotype_index, time, start=0, end=100):
return grouping.MatchJob(
return config.MatchJob(
haplotype_index=haplotype_index,
source="test",
sample_id=f"s{haplotype_index}",
Expand Down Expand Up @@ -617,7 +617,7 @@ def test_sorted_by_group_then_haplotype_index(self):
)

def test_preserves_job_fields(self):
job = grouping.MatchJob(
job = config.MatchJob(
haplotype_index=0,
source="my_source",
sample_id="my_sample",
Expand Down
21 changes: 16 additions & 5 deletions tests/test_matcher_fixtures.py
Original file line number Diff line number Diff line change
Expand Up @@ -29,8 +29,7 @@
import numpy as np
import pytest

from tsinfer import matching, vcz
from tsinfer.grouping import MatchJob
from tsinfer import config, matching, vcz

# ---------------------------------------------------------------------------
# Helpers
Expand All @@ -52,8 +51,20 @@ def read_haplotype(self, job):

def _match(ts, positions, query, allele_mapper):
"""Run the current Matcher and return (path, mutations) as tuples."""
matcher = matching.Matcher(ts, positions)
results = list(matcher.match([None], _ArrayReader([query])))
job = config.MatchJob(
haplotype_index=0,
source="test",
sample_id="s0",
ploidy_index=0,
time=0,
start_position=0,
end_position=int(ts.sequence_length),
group=0,
)
matcher = matching.Matcher(
ts, positions, source_parameters={"test": config.MatchSourceConfig()}
)
results = list(matcher.match([job], _ArrayReader([query])))
_, r = results[0]
path = [(s.left, s.right, s.parent) for s in r.path]
mutations = [(m.position, m.derived_state) for m in r.mutations]
Expand All @@ -62,7 +73,7 @@ def _match(ts, positions, query, allele_mapper):

def _add_node(ts, time, path_segments, mutations, allele_mapper):
"""Add a single ancestor node to *ts* via extend_ts."""
job = MatchJob(
job = config.MatchJob(
haplotype_index=0,
source="test",
sample_id="s0",
Expand Down
46 changes: 32 additions & 14 deletions tests/test_matching.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
import numpy as np
import tskit

from tsinfer import grouping, matching, vcz
from tsinfer import config, matching, vcz

# ---------------------------------------------------------------------------
# Helpers
Expand Down Expand Up @@ -92,8 +92,8 @@ def read_haplotype(self, job):


def _jobs(n):
"""Return *n* opaque dummy jobs (the reader ignores them)."""
return [None] * n
"""Return *n* dummy MatchJobs (the reader ignores their content)."""
return [_make_job(haplotype_index=i) for i in range(n)]


def _make_job(
Expand All @@ -110,7 +110,7 @@ def _make_job(
population_id=None,
):
"""Create a MatchJob with sensible defaults for testing."""
return grouping.MatchJob(
return config.MatchJob(
haplotype_index=haplotype_index,
source=source,
sample_id=sample_id,
Expand All @@ -125,6 +125,9 @@ def _make_job(
)


_SOURCE_PARAMS = {"test": config.MatchSourceConfig()}


# ---------------------------------------------------------------------------
# TestMakeRootTs
# ---------------------------------------------------------------------------
Expand Down Expand Up @@ -252,7 +255,7 @@ def test_match_identical_haplotype(self):
"""Matching an identical haplotype returns a MatchResult with path edges."""
hap = np.array([0, 1, 0], dtype=np.int8)
ts = self._make_ts_with_one_ancestor(hap)
matcher = matching.Matcher(ts, self.positions)
matcher = matching.Matcher(ts, self.positions, source_parameters=_SOURCE_PARAMS)
results = list(matcher.match(_jobs(1), _ArrayReader([hap])))
assert len(results) == 1
_, r = results[0]
Expand All @@ -265,7 +268,7 @@ def test_match_with_mutation(self):
ts = self._make_ts_with_one_ancestor(ancestor_hap)
# Query haplotype has derived state at site 1
query_hap = np.array([0, 1, 0], dtype=np.int8)
matcher = matching.Matcher(ts, self.positions)
matcher = matching.Matcher(ts, self.positions, source_parameters=_SOURCE_PARAMS)
results = list(matcher.match(_jobs(1), _ArrayReader([query_hap])))
_, r = results[0]
# There should be a mutation at position 20 (site index 1)
Expand All @@ -277,7 +280,7 @@ def test_match_all_missing(self):
hap = np.array([0, 1, 0], dtype=np.int8)
ts = self._make_ts_with_one_ancestor(hap)
missing_hap = np.array([-1, -1, -1], dtype=np.int8)
matcher = matching.Matcher(ts, self.positions)
matcher = matching.Matcher(ts, self.positions, source_parameters=_SOURCE_PARAMS)
results = list(matcher.match(_jobs(1), _ArrayReader([missing_hap])))
assert len(results) == 1

Expand All @@ -286,7 +289,7 @@ def test_match_partial_missing(self):
hap = np.array([0, 1, 0], dtype=np.int8)
ts = self._make_ts_with_one_ancestor(hap)
query_hap = np.array([-1, 1, 0], dtype=np.int8)
matcher = matching.Matcher(ts, self.positions)
matcher = matching.Matcher(ts, self.positions, source_parameters=_SOURCE_PARAMS)
results = list(matcher.match(_jobs(1), _ArrayReader([query_hap])))
assert len(results) == 1
_, r = results[0]
Expand All @@ -299,15 +302,15 @@ def test_match_returns_one_result_per_haplotype(self):
hap = np.array([0, 1, 0], dtype=np.int8)
ts = self._make_ts_with_one_ancestor(hap)
haplotypes = np.array([[0, 1, 0], [0, 0, 1], [1, 0, 0]], dtype=np.int8)
matcher = matching.Matcher(ts, self.positions)
matcher = matching.Matcher(ts, self.positions, source_parameters=_SOURCE_PARAMS)
results = list(matcher.match(_jobs(3), _ArrayReader(haplotypes)))
assert len(results) == 3

def test_match_result_types(self):
"""MatchResult should contain PathSegment and Mutation objects."""
hap = np.array([0, 1, 0], dtype=np.int8)
ts = self._make_ts_with_one_ancestor(hap)
matcher = matching.Matcher(ts, self.positions)
matcher = matching.Matcher(ts, self.positions, source_parameters=_SOURCE_PARAMS)
results = list(matcher.match(_jobs(1), _ArrayReader([hap])))
_, r = results[0]
assert isinstance(r.path, list)
Expand All @@ -322,7 +325,7 @@ def test_path_segments_have_valid_coordinates(self):
"""PathSegment left < right for every segment."""
hap = np.array([0, 1, 0], dtype=np.int8)
ts = self._make_ts_with_one_ancestor(hap)
matcher = matching.Matcher(ts, self.positions)
matcher = matching.Matcher(ts, self.positions, source_parameters=_SOURCE_PARAMS)
results = list(matcher.match(_jobs(1), _ArrayReader([hap])))
_, r = results[0]
for seg in r.path:
Expand All @@ -333,7 +336,7 @@ def test_mutations_have_valid_fields(self):
hap = np.array([0, 0, 0], dtype=np.int8)
ts = self._make_ts_with_one_ancestor(hap)
query = np.array([1, 1, 0], dtype=np.int8)
matcher = matching.Matcher(ts, self.positions)
matcher = matching.Matcher(ts, self.positions, source_parameters=_SOURCE_PARAMS)
results = list(matcher.match(_jobs(1), _ArrayReader([query])))
_, r = results[0]
assert len(r.mutations) > 0
Expand All @@ -342,6 +345,19 @@ def test_mutations_have_valid_fields(self):
assert isinstance(m.position, float)
assert isinstance(m.derived_state, int)

def test_match_with_custom_recombination_and_mismatch(self):
"""Custom per-source recombination/mismatch overrides are used."""
hap = np.array([0, 1, 0], dtype=np.int8)
ts = self._make_ts_with_one_ancestor(hap)
source_params = {
"test": config.MatchSourceConfig(recombination=0.5, mismatch=0.1),
}
matcher = matching.Matcher(ts, self.positions, source_parameters=source_params)
results = list(matcher.match(_jobs(1), _ArrayReader([hap])))
assert len(results) == 1
_, r = results[0]
assert len(r.path) > 0


# ---------------------------------------------------------------------------
# TestExtendTs
Expand Down Expand Up @@ -673,7 +689,7 @@ def _build_root_ts(self):

def _match_and_pair(self, ts, haplotypes):
"""Match haplotypes against ts and return paired_results list."""
matcher = matching.Matcher(ts, self.positions)
matcher = matching.Matcher(ts, self.positions, source_parameters=_SOURCE_PARAMS)
results = list(matcher.match(_jobs(len(haplotypes)), _ArrayReader(haplotypes)))
return results

Expand Down Expand Up @@ -705,7 +721,9 @@ def test_matching_after_two_extends(self):

# Now match a sample against ts2
sample_hap = np.array([0, 0, 1, 1, 0], dtype=np.int8)
matcher2 = matching.Matcher(ts2, self.positions)
matcher2 = matching.Matcher(
ts2, self.positions, source_parameters=_SOURCE_PARAMS
)
results2 = list(matcher2.match(_jobs(1), _ArrayReader([sample_hap])))
_, r = results2[0]
assert len(r.path) > 0
Expand Down
20 changes: 10 additions & 10 deletions tests/test_vcz.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,7 +30,7 @@
import pytest
import zarr

from tsinfer import config, grouping, vcz
from tsinfer import config, vcz

# ---------------------------------------------------------------------------
# Fixtures
Expand Down Expand Up @@ -346,7 +346,7 @@ def test_ancestor_haplotype(self):
schedule=schedule,
)
# a0 has genotype [0, 1, 0]
job = grouping.MatchJob(
job = config.MatchJob(
haplotype_index=1,
source="ancestors",
sample_id="a0",
Expand All @@ -371,7 +371,7 @@ def test_ancestor_haplotype_second(self):
schedule=schedule,
)
# a1 has genotype [1, 0, 1]
job = grouping.MatchJob(
job = config.MatchJob(
haplotype_index=2,
source="ancestors",
sample_id="a1",
Expand All @@ -396,7 +396,7 @@ def test_sample_haplotype_encoding(self):
schedule=schedule,
)
# sample_0: gt [0, 1, 1] → encoded [0=anc, 1=derived, 1=derived]
job = grouping.MatchJob(
job = config.MatchJob(
haplotype_index=3,
source="test",
sample_id="sample_0",
Expand All @@ -421,7 +421,7 @@ def test_sample_haplotype_second(self):
schedule=schedule,
)
# sample_1: gt [1, 0, 0] → encoded [1=derived, 0=anc, 0=anc]
job = grouping.MatchJob(
job = config.MatchJob(
haplotype_index=4,
source="test",
sample_id="sample_1",
Expand Down Expand Up @@ -467,7 +467,7 @@ def test_missing_genotype_encoded_as_minus_one(self):
anc_alleles,
schedule=schedule,
)
job = grouping.MatchJob(
job = config.MatchJob(
haplotype_index=2,
source="test",
sample_id="sample_0",
Expand Down Expand Up @@ -519,7 +519,7 @@ def test_multiple_sources(self):
schedule = [("src_a", "sample_0", 0), ("src_b", "sample_0", 0)]
reader = vcz.HaplotypeReader(sources, positions, anc_alleles, schedule=schedule)

job_a = grouping.MatchJob(
job_a = config.MatchJob(
haplotype_index=2,
source="src_a",
sample_id="sample_0",
Expand All @@ -529,7 +529,7 @@ def test_multiple_sources(self):
end_position=200,
group=2,
)
job_b = grouping.MatchJob(
job_b = config.MatchJob(
haplotype_index=3,
source="src_b",
sample_id="sample_0",
Expand Down Expand Up @@ -1475,7 +1475,7 @@ def test_facade_returns_alleles_after_reads(self):
reader = vcz.HaplotypeReader(sources, positions, anc_alleles, schedule=schedule)

# Read from both sources to trigger allele discovery
job_a = grouping.MatchJob(
job_a = config.MatchJob(
haplotype_index=0,
source="src_a",
sample_id="sample_0",
Expand All @@ -1485,7 +1485,7 @@ def test_facade_returns_alleles_after_reads(self):
end_position=200,
group=0,
)
job_b = grouping.MatchJob(
job_b = config.MatchJob(
haplotype_index=1,
source="src_b",
sample_id="sample_0",
Expand Down
Loading
Loading