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
227 changes: 227 additions & 0 deletions tests/test_arg_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,3 +17,230 @@ def test_empty_tree_sequence(self):
ts = tables.tree_sequence()
with pytest.raises(ValueError, match="Emtpy trees"):
arg_ops.add_vestigial_root(ts)


def _make_ts(nodes, edges, sites=None, mutations=None, sequence_length=100):
"""Helper to build a tree sequence from lists of (flags, time) and
(left, right, parent, child) tuples."""
tables = tskit.TableCollection(sequence_length=sequence_length)
for flags, time in nodes:
tables.nodes.add_row(flags=flags, time=time)
for left, right, parent, child in edges:
tables.edges.add_row(left=left, right=right, parent=parent, child=child)
if sites is not None:
for pos, ancestral in sites:
tables.sites.add_row(position=pos, ancestral_state=ancestral)
if mutations is not None:
for site, node, derived in mutations:
tables.mutations.add_row(site=site, node=node, derived_state=derived)
tables.sort()
return tables.tree_sequence()


class TestIsPcAncestor:
def test_zero_flags(self):
assert not arg_ops.is_pc_ancestor(0)

def test_sample_flag(self):
assert not arg_ops.is_pc_ancestor(tskit.NODE_IS_SAMPLE)

def test_pc_flag(self):
assert arg_ops.is_pc_ancestor(arg_ops.NODE_IS_PC_ANCESTOR)

def test_pc_flag_combined(self):
assert arg_ops.is_pc_ancestor(arg_ops.NODE_IS_PC_ANCESTOR | tskit.NODE_IS_SAMPLE)

def test_other_bits(self):
for bit in range(32):
flags = 1 << bit
if bit == 16:
assert arg_ops.is_pc_ancestor(flags)
else:
assert not arg_ops.is_pc_ancestor(flags)


class TestCompressPaths:
def test_identical_single_edge_path(self):
"""Two children with identical single-edge path — compressed."""
ts = _make_ts(
nodes=[(1, 0), (1, 0), (0, 1.0)],
edges=[(0, 100, 2, 0), (0, 100, 2, 1)],
)
result = arg_ops.compress_paths(ts)
assert result.num_nodes == 4 # 3 + 1 PC
# 3 edges: parent->PC, PC->0, PC->1
assert result.num_edges == 3
pc_node = result.node(3)
assert arg_ops.is_pc_ancestor(pc_node.flags)
assert pc_node.time == 1.0 - arg_ops.PC_ANCESTOR_INCREMENT

def test_identical_two_edge_path(self):
"""Two children with identical two-edge path — one PC node."""
ts = _make_ts(
nodes=[(1, 0), (1, 0), (0, 1.0), (0, 2.0), (0, 2.0)],
edges=[
(0, 50, 3, 0),
(50, 100, 4, 0),
(0, 50, 3, 1),
(50, 100, 4, 1),
],
)
result = arg_ops.compress_paths(ts)
assert result.num_nodes == 6 # 5 + 1 PC
# 6 edges: 3->PC, 4->PC, PC->0 x2, PC->1 x2
assert result.num_edges == 6
pc_node = result.node(5)
assert arg_ops.is_pc_ancestor(pc_node.flags)
assert pc_node.time == 2.0 - arg_ops.PC_ANCESTOR_INCREMENT

def test_three_children_identical_path(self):
"""Three children with identical path — one PC node."""
ts = _make_ts(
nodes=[(1, 0), (1, 0), (1, 0), (0, 2.0)],
edges=[(0, 100, 3, 0), (0, 100, 3, 1), (0, 100, 3, 2)],
)
result = arg_ops.compress_paths(ts)
assert result.num_nodes == 5 # 4 + 1 PC
# 4 edges: parent->PC, PC->0, PC->1, PC->2
assert result.num_edges == 4

def test_different_paths_not_compressed(self):
"""Children with different edge sets are not compressed."""
ts = _make_ts(
nodes=[(1, 0), (1, 0), (0, 1.0), (0, 2.0), (0, 2.0), (0, 2.0)],
edges=[
(0, 50, 3, 0),
(50, 100, 5, 0), # parent 5 for child 0
(0, 50, 3, 1),
(50, 100, 4, 1), # parent 4 for child 1
],
)
result = arg_ops.compress_paths(ts)
assert result.num_nodes == ts.num_nodes
assert result.num_edges == ts.num_edges

def test_no_shared_edges(self):
"""Each child has unique edges — nothing to compress."""
ts = _make_ts(
nodes=[(1, 0), (1, 0), (0, 1.0)],
edges=[(0, 100, 2, 0), (0, 50, 2, 1)], # different intervals
)
result = arg_ops.compress_paths(ts)
assert result.num_nodes == ts.num_nodes
assert result.num_edges == ts.num_edges

def test_partial_path_overlap_not_compressed(self):
"""Children share some edges but not all — not compressed."""
# Child 0 has edges: (0,50,P3) and (50,100,P4)
# Child 1 has edges: (0,50,P3) and (50,100,P5) ← differs
ts = _make_ts(
nodes=[(1, 0), (1, 0), (0, 1.0), (0, 2.0), (0, 2.0), (0, 2.0)],
edges=[
(0, 50, 3, 0),
(50, 100, 4, 0),
(0, 50, 3, 1),
(50, 100, 5, 1),
],
)
result = arg_ops.compress_paths(ts)
assert result.num_nodes == ts.num_nodes
assert result.num_edges == ts.num_edges

def test_preserves_other_edges(self):
"""Edges for non-compressed children are preserved."""
ts = _make_ts(
nodes=[
(1, 0),
(1, 0),
(1, 0),
(0, 1.0),
(0, 2.0),
(0, 2.0),
],
edges=[
(0, 50, 4, 0),
(50, 100, 5, 0),
(0, 50, 4, 1),
(50, 100, 5, 1),
(0, 100, 3, 2), # different path, not shared
],
)
result = arg_ops.compress_paths(ts)
assert result.num_nodes == 7 # 6 + 1 PC
# 6 compressed edges + 1 preserved = 7
assert result.num_edges == 7

def test_multiple_groups(self):
"""Two groups of children with different identical paths."""
# Children 0,1 share path through parent 4
# Children 2,3 share path through parent 5
ts = _make_ts(
nodes=[(1, 0), (1, 0), (1, 0), (1, 0), (0, 2.0), (0, 2.0)],
edges=[
(0, 100, 4, 0),
(0, 100, 4, 1),
(0, 100, 5, 2),
(0, 100, 5, 3),
],
)
result = arg_ops.compress_paths(ts)
assert result.num_nodes == 8 # 6 + 2 PC
# 6 edges: 2 * (parent->PC, PC->child, PC->child)
assert result.num_edges == 6

def test_time_too_close_skipped(self):
"""Group is skipped when PC ancestor time would not be valid."""
inc = arg_ops.PC_ANCESTOR_INCREMENT
ts = _make_ts(
nodes=[(1, 1.0), (1, 1.0), (0, 1.0 + inc / 2)],
edges=[(0, 100, 2, 0), (0, 100, 2, 1)],
)
result = arg_ops.compress_paths(ts)
# No compression possible, tree unchanged
assert result.num_nodes == ts.num_nodes
assert result.num_edges == ts.num_edges

def test_with_mutations(self):
"""Mutations on compressed edges still reference correct nodes."""
ts = _make_ts(
nodes=[(1, 0), (1, 0), (0, 1.0)],
edges=[(0, 100, 2, 0), (0, 100, 2, 1)],
sites=[(50, "A")],
mutations=[(0, 0, "T")],
)
result = arg_ops.compress_paths(ts)
assert result.mutation(0).node == 0
assert result.mutation(0).derived_state == "T"

def test_no_edges(self):
"""Tree sequence with no edges — nothing to compress."""
ts = _make_ts(nodes=[(1, 0)], edges=[])
result = arg_ops.compress_paths(ts)
assert result.num_nodes == 1
assert result.num_edges == 0

def test_min_parent_time_used(self):
"""PC ancestor time uses minimum parent time across the path."""
ts = _make_ts(
nodes=[(1, 0), (1, 0), (0, 1.0), (0, 3.0), (0, 5.0)],
edges=[
(0, 50, 3, 0),
(50, 100, 4, 0),
(0, 50, 3, 1),
(50, 100, 4, 1),
],
)
result = arg_ops.compress_paths(ts)
pc_node = result.node(5)
assert pc_node.time == 3.0 - arg_ops.PC_ANCESTOR_INCREMENT

def test_pc_ancestor_flags(self):
"""New PC ancestor nodes have NODE_IS_PC_ANCESTOR flag set."""
ts = _make_ts(
nodes=[(1, 0), (1, 0), (0, 1.0)],
edges=[(0, 100, 2, 0), (0, 100, 2, 1)],
)
result = arg_ops.compress_paths(ts)
for i in range(3):
assert not arg_ops.is_pc_ancestor(result.node(i).flags)
assert arg_ops.is_pc_ancestor(result.node(3).flags)
102 changes: 102 additions & 0 deletions tsinfer/arg_ops.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,10 +20,15 @@
Operations on ARG (Ancestral Recombination Graph) topology.
"""

import collections
import logging
import time as time_

logger = logging.getLogger(__name__)

PC_ANCESTOR_INCREMENT = 1.0 / (1 << 32)
NODE_IS_PC_ANCESTOR = 1 << 16


def add_vestigial_root(ts):
"""
Expand Down Expand Up @@ -59,3 +64,100 @@ def add_vestigial_root(ts):
# we can just sort almost the end of the table.
tables.sort()
return tables.tree_sequence()


def is_pc_ancestor(flags):
"""
Returns True if the node flags indicate a path compression ancestor.
"""
return (flags & NODE_IS_PC_ANCESTOR) != 0


def compress_paths(ts):
"""
Find groups of child nodes whose entire edge sets are identical
(same set of (left, right, parent) tuples) and create an intermediate
path compression ancestor node for each group.

For each group of 2+ children with identical paths, a new node is
created with time slightly less than the minimum parent time in the
path. The children's edges are rewritten to go through the new node.
"""
start_time = time_.time()
tables = ts.dump_tables()
node_time = tables.nodes.time
original_num_nodes = ts.num_nodes
original_num_edges = ts.num_edges

# Build path key for each child: the full sorted tuple of its edges
edges_by_child = collections.defaultdict(list)
for edge in tables.edges:
edges_by_child[edge.child].append((edge.left, edge.right, edge.parent))

path_map = collections.defaultdict(list)
for child, edges in edges_by_child.items():
key = tuple(sorted(edges))
path_map[key].append(child)

# Groups with 2+ children sharing identical paths get a PC ancestor
compressed_children = set()
new_edges = []
num_pc_nodes = 0
for path, children in path_map.items():
if len(children) < 2:
continue

# Compute PC ancestor time
min_parent_time = min(node_time[parent] for _, _, parent in path)
pc_time = min_parent_time - PC_ANCESTOR_INCREMENT
max_child_time = max(node_time[child] for child in children)
if pc_time <= max_child_time:
logger.debug(
"Skipping path compression group: computed time %f "
"is not greater than maximum child time %f",
pc_time,
max_child_time,
)
continue

# Create PC ancestor node
pc_node = tables.nodes.add_row(flags=NODE_IS_PC_ANCESTOR, time=pc_time)
num_pc_nodes += 1

# Add edges from each original parent to the PC node
for left, right, parent in path:
new_edges.append((left, right, parent, pc_node))
# Add edges from PC node to each child
for child in children:
for left, right, _ in path:
new_edges.append((left, right, pc_node, child))
compressed_children.update(children)

# Rebuild edge table: drop edges for compressed children, add new ones
tables.edges.clear()
for edge in ts.edges():
if edge.child not in compressed_children:
tables.edges.add_row(
left=edge.left,
right=edge.right,
parent=edge.parent,
child=edge.child,
)
for left, right, parent, child in new_edges:
tables.edges.add_row(left=left, right=right, parent=parent, child=child)

tables.sort()
result = tables.tree_sequence()
elapsed = time_.time() - start_time
logger.info(
"Path compression: %d PC nodes added, %d children compressed, "
"edges %d -> %d, nodes %d -> %d (%.2fs)",
num_pc_nodes,
len(compressed_children),
original_num_edges,
result.num_edges,
original_num_nodes,
result.num_nodes,
elapsed,
)
return result
2 changes: 0 additions & 2 deletions tsinfer/matching.py
Original file line number Diff line number Diff line change
Expand Up @@ -155,13 +155,11 @@ def __init__(
self,
ts: tskit.TreeSequence,
positions: np.ndarray, # (num_sites,) int32 — inference site positions
path_compression: bool = True,
num_alleles: np.ndarray | None = None,
):
self._positions = np.asarray(positions, dtype=np.int32)
self._num_sites = len(positions)
self._sequence_length = ts.sequence_length
self._path_compression = path_compression

logger.info(
"Creating Matcher: ts=%d nodes, %d edges, %.1f MiB, RSS=%.1f MiB",
Expand Down
8 changes: 5 additions & 3 deletions tsinfer/pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -35,7 +35,7 @@
import tqdm
import tskit

from . import ancestors, config, grouping, matching, provenance
from . import ancestors, arg_ops, config, grouping, matching, provenance
from . import vcz as vcz_mod

logger = logging.getLogger(__name__)
Expand Down Expand Up @@ -362,7 +362,6 @@ def _process_group(
m = matching.Matcher(
ts,
positions,
path_compression=path_compression,
num_alleles=reader.get_num_alleles(),
)
job_list = [job for _, job in group_jobs]
Expand Down Expand Up @@ -420,12 +419,15 @@ def _process_group(
resources=tm.metrics.asdict(),
)

return matching.extend_ts(
result = matching.extend_ts(
ts,
paired_results=paired_results,
allele_mapper=allele_mapper,
provenance_record=json.dumps(prov_dict),
)
if path_compression:
result = arg_ops.compress_paths(result)
return result


def match(
Expand Down
Loading