diff --git a/tests/test_arg_ops.py b/tests/test_arg_ops.py index 904ff407..6698a23e 100644 --- a/tests/test_arg_ops.py +++ b/tests/test_arg_ops.py @@ -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) diff --git a/tsinfer/arg_ops.py b/tsinfer/arg_ops.py index 05e337bc..d9deb078 100644 --- a/tsinfer/arg_ops.py +++ b/tsinfer/arg_ops.py @@ -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): """ @@ -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 diff --git a/tsinfer/matching.py b/tsinfer/matching.py index 9c002216..477a22f6 100644 --- a/tsinfer/matching.py +++ b/tsinfer/matching.py @@ -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", diff --git a/tsinfer/pipeline.py b/tsinfer/pipeline.py index 93cde3d2..8fcbe2ec 100644 --- a/tsinfer/pipeline.py +++ b/tsinfer/pipeline.py @@ -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__) @@ -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] @@ -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(