Skip to content

Commit a0ea863

Browse files
authored
add select_representive_sample_uids.py (#662)
1 parent f02b0b9 commit a0ea863

9 files changed

Lines changed: 3202 additions & 3 deletions
File renamed without changes.
File renamed without changes.
File renamed without changes.

sqlite/util/select_evaluation_subset.py renamed to graph_net/sqlite_util/select_evaluation_subset.py

Lines changed: 14 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -4,6 +4,7 @@
44

55
import math
66
from collections import Counter
7+
from typing import Callable
78

89
import numpy as np
910
from scipy import sparse
@@ -88,12 +89,17 @@ def _compute_metrics(
8889
return result
8990

9091

91-
def _greedy_select(metrics: list[dict], k: int, rarity_weight: float) -> list[int]:
92+
def _greedy_select(
93+
metrics: list[dict], k: int, rarity_weight: float, selected: list[int] = None
94+
) -> list[int]:
9295
"""Greedily maximise edge coverage, weighted by rarity score."""
9396
target = min(k, len(metrics))
9497
rarity_norm = _min_max_normalize([m["rarity_score"] for m in metrics])
9598

96-
selected: list[int] = []
99+
if selected is None:
100+
selected: list[int] = []
101+
else:
102+
assert isinstance(selected, list)
97103
selected_set: set[int] = set()
98104
covered_edges: set[tuple[str, str]] = set()
99105

@@ -123,6 +129,7 @@ def select_evaluation_subset(
123129
*,
124130
smoothing_alpha: float = 1e-3,
125131
rarity_weight: float = 1,
132+
is_selected: Callable[tuple[str, ...], bool] = lambda x: False,
126133
) -> list[tuple[str, ...]]:
127134
"""Select k sequences from op_seqs using Markov-based greedy coverage.
128135
@@ -140,4 +147,8 @@ def select_evaluation_subset(
140147

141148
op_to_id, count_matrix, row_sums = _build_markov_model(seqs)
142149
metrics = _compute_metrics(seqs, op_to_id, count_matrix, row_sums, smoothing_alpha)
143-
return [seqs[i] for i in _greedy_select(metrics, k, rarity_weight)]
150+
selected_indexes = [i for i, seq in enumerate(seqs) if is_selected(seq)]
151+
return [
152+
seqs[i]
153+
for i in _greedy_select(metrics, k, rarity_weight, selected=selected_indexes)
154+
]
Lines changed: 108 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,108 @@
1+
#!/usr/bin/env python3
2+
"""
3+
select_representative_sample_uids.py
4+
5+
Reads a TSV file of (sample_uid, op_seq) pairs and a file listing selected op_seqs.
6+
For each selected op_seq, picks the sample_uid with the maximum string length from its group,
7+
and outputs them in the same order as the selected op_seqs.
8+
"""
9+
10+
import argparse
11+
import sys
12+
from collections import defaultdict
13+
14+
15+
def read_tsv_pairs(file_path):
16+
"""
17+
Read a tab-separated file where each line contains sample_uid and op_seq.
18+
Returns a list of (sample_uid, op_seq) tuples.
19+
"""
20+
pairs = []
21+
with open(file_path, "r", encoding="utf-8") as f:
22+
for line_num, line in enumerate(f, 1):
23+
line = line.strip()
24+
if not line:
25+
continue
26+
parts = line.split("\t")
27+
if len(parts) != 2:
28+
print(
29+
f"Warning: {file_path}:{line_num} invalid format, skipped: {line}",
30+
file=sys.stderr,
31+
)
32+
continue
33+
sample_uid, op_seq = parts
34+
pairs.append((sample_uid, op_seq))
35+
return pairs
36+
37+
38+
def read_op_seq_list(file_path):
39+
"""
40+
Read a file with one op_seq per line. Returns a list preserving order.
41+
"""
42+
op_seqs = []
43+
with open(file_path, "r", encoding="utf-8") as f:
44+
for line in f:
45+
line = line.strip()
46+
if line:
47+
op_seqs.append(line)
48+
return op_seqs
49+
50+
51+
def get_max_len_string(strings):
52+
"""
53+
Return the string with the maximum length from a list.
54+
If multiple strings have the same max length, return the first encountered.
55+
"""
56+
if not strings:
57+
return None
58+
# max with key returns the first element in case of ties (Python's max is stable)
59+
return max(strings, key=lambda s: len(s))
60+
61+
62+
def main():
63+
parser = argparse.ArgumentParser(
64+
description="Select representative sample UIDs by max string length per op_seq group."
65+
)
66+
parser.add_argument(
67+
"pairs_file", help="TSV file with two columns: sample_uid and op_seq"
68+
)
69+
parser.add_argument(
70+
"opseq_file", help="File with one op_seq per line (order matters)"
71+
)
72+
parser.add_argument("-o", "--output", help="Output file path (default: stdout)")
73+
args = parser.parse_args()
74+
75+
# 1. Load input data
76+
pairs = read_tsv_pairs(args.pairs_file)
77+
selected_op_seqs = read_op_seq_list(args.opseq_file)
78+
79+
# 2. Group sample_uids by op_seq
80+
groups = defaultdict(list)
81+
for sample_uid, op_seq in pairs:
82+
groups[op_seq].append(sample_uid)
83+
84+
# 3. For each op_seq, find the longest sample_uid
85+
op_seq_to_max_uid = {}
86+
for op_seq, uid_list in groups.items():
87+
max_uid = get_max_len_string(uid_list)
88+
if max_uid is not None:
89+
op_seq_to_max_uid[op_seq] = max_uid
90+
91+
# 4. Collect results in the order of selected_op_seqs
92+
result_uids = []
93+
for op_seq in selected_op_seqs:
94+
if op_seq in op_seq_to_max_uid:
95+
result_uids.append(op_seq_to_max_uid[op_seq])
96+
97+
# 5. Write output
98+
out_fh = open(args.output, "w", encoding="utf-8") if args.output else sys.stdout
99+
try:
100+
for uid in result_uids:
101+
print(uid, file=out_fh)
102+
finally:
103+
if args.output:
104+
out_fh.close()
105+
106+
107+
if __name__ == "__main__":
108+
main()

0 commit comments

Comments
 (0)