import argparse
import gzip
import os
import re
import shutil
import subprocess
import sys
from dataclasses import dataclass, field
from itertools import permutations
from pathlib import Path
from typing import Optional
PBSIM = Path("~/install/pbsim3/src/pbsim").expanduser()
PBMODELS = Path("~/install/pbsim3/data").expanduser()
POA_BIN = Path(__file__).parent.parent / "target/release/poa-consensus"
_U = "ACGTACGTCGATCGATTAGCTAGCGCTAGCTA" _LEFT_FLANK_FULL = _U * 300 _RIGHT_FLANK_FULL = _U[::-1] * 300
FLANK_LEN = 2_048 LONG_FLANK_LEN = 8_000 LEFT_FLANK = _LEFT_FLANK_FULL[:FLANK_LEN] RIGHT_FLANK = _RIGHT_FLANK_FULL[:FLANK_LEN]
_FLANK_ROTATE_STEP = 11
def _rotated_flank(pool: str, flank_len: int, allele_idx: int) -> str:
offset = allele_idx * _FLANK_ROTATE_STEP
return pool[offset: offset + flank_len]
ANCHOR_PAD = 100
ERRMODELS: dict[str, dict] = {
"ont_r9": dict(
errhmm="ERRHMM-ONT.model",
accuracy_mean=0.85, difference_ratio="39:24:36",
mm2_preset="map-ont",
label="ONT R9 (85% acc)",
),
"ont_r10": dict(
errhmm="ERRHMM-ONT-HQ.model",
accuracy_mean=0.92, difference_ratio="39:24:36",
mm2_preset="map-ont",
label="ONT R10 HQ (92% acc)",
),
"hifi": dict(
errhmm="ERRHMM-SEQUEL.model",
accuracy_mean=0.89, difference_ratio="22:45:33",
mm2_preset="map-hifi",
label="HiFi/Sequel (89% acc)",
),
}
@dataclass
class Allele:
unit: str count: int depth: int desc: Optional[str] = None
@property
def seq(self) -> str:
return self.unit * self.count
@property
def id(self) -> str:
return f"{self.unit.lower()}{self.count}"
@property
def display(self) -> str:
return self.desc if self.desc is not None else f"{self.unit}×{self.count}"
@dataclass
class Scenario:
name: str
alleles: list[Allele] model: str multi: bool = False poa_flags: list = field(default_factory=list)
seed: int = 42
length_mean: Optional[int] = None flank_len: int = FLANK_LEN long_read: bool = False partial: bool = False
@property
def is_multi(self) -> bool:
return self.multi
SCENARIOS: list[Scenario] = [
Scenario("cag20_d05_r10", [Allele("CAG", 20, 5)], "ont_r10"),
Scenario("cag20_d10_r10", [Allele("CAG", 20, 10)], "ont_r10"),
Scenario("cag20_d20_r10", [Allele("CAG", 20, 20)], "ont_r10"),
Scenario("cag20_d30_r10", [Allele("CAG", 20, 30)], "ont_r10"),
Scenario("cag20_d20_r9", [Allele("CAG", 20, 20)], "ont_r9"),
Scenario("cag20_d20_hifi", [Allele("CAG", 20, 20)], "hifi"),
Scenario("cag5_d20", [Allele("CAG", 5, 20)], "ont_r10"),
Scenario("cag10_d20", [Allele("CAG", 10, 20)], "ont_r10"),
Scenario("cag50_d20", [Allele("CAG", 50, 20)], "ont_r10"),
Scenario("cag100_d20", [Allele("CAG", 100, 20)], "ont_r10"),
Scenario("cag200_d20", [Allele("CAG", 200, 20)], "ont_r10"),
Scenario("gaa50_d20", [Allele("GAA", 50, 20)], "ont_r10"),
Scenario("gaa100_d20", [Allele("GAA", 100, 20)], "ont_r10"),
Scenario("gaa200_d10", [Allele("GAA", 200, 10)], "ont_r10"),
Scenario("multi_cag15_25",
[Allele("CAG", 15, 20), Allele("CAG", 25, 20)],
"ont_r10", multi=True),
Scenario("multi_cag20_50",
[Allele("CAG", 20, 20), Allele("CAG", 50, 20)],
"ont_r10", multi=True),
Scenario("multi_gaa30_100",
[Allele("GAA", 30, 20), Allele("GAA", 100, 20)],
"ont_r10", multi=True),
Scenario("sv_cag20_out60",
[Allele("CAG", 20, 20), Allele("CAG", 60, 2)],
"ont_r10", multi=False),
Scenario("sv_gaa50_out200",
[Allele("GAA", 50, 20), Allele("GAA", 200, 2)],
"ont_r10", multi=False),
Scenario("multi_skew_cag20_40",
[Allele("CAG", 20, 24), Allele("CAG", 40, 8)],
"ont_r10", multi=True),
Scenario("lr_cag100_15k",
[Allele("CAG", 100, 20)], "ont_r10",
length_mean=15_000, flank_len=LONG_FLANK_LEN, long_read=True),
Scenario("lr_cag200_15k",
[Allele("CAG", 200, 20)], "ont_r10",
length_mean=15_000, flank_len=LONG_FLANK_LEN, long_read=True),
Scenario("lr_gaa200_15k",
[Allele("GAA", 200, 20)], "ont_r10",
length_mean=15_000, flank_len=LONG_FLANK_LEN, long_read=True),
Scenario("lr_gaa500_15k",
[Allele("GAA", 500, 20)], "ont_r10",
length_mean=15_000, flank_len=LONG_FLANK_LEN, long_read=True),
Scenario("lr_rfc1_100_15k",
[Allele("CTTTT", 100, 20)], "ont_r10",
length_mean=15_000, flank_len=LONG_FLANK_LEN, long_read=True),
Scenario("lr_rfc1_500_15k",
[Allele("CTTTT", 500, 20)], "ont_r10",
length_mean=15_000, flank_len=LONG_FLANK_LEN, long_read=True),
Scenario("lr_multi_cag100_200_15k",
[Allele("CAG", 100, 20), Allele("CAG", 200, 20)], "ont_r10",
multi=True, length_mean=15_000, flank_len=LONG_FLANK_LEN, long_read=True),
]
def run(cmd: list, *, cwd=None, capture=False, check=True) -> subprocess.CompletedProcess:
kwargs = dict(cwd=cwd, check=check)
if capture:
kwargs["capture_output"] = True
else:
kwargs["stdout"] = subprocess.DEVNULL
kwargs["stderr"] = subprocess.DEVNULL
return subprocess.run([str(c) for c in cmd], **kwargs)
def levenshtein(a: bytes, b: bytes, max_len: int = 2000) -> int:
if len(a) > max_len or len(b) > max_len:
return abs(len(a) - len(b))
dp = list(range(len(b) + 1))
for ca in a:
dp2 = [dp[0] + 1] + [0] * len(b)
for j, cb in enumerate(b):
dp2[j + 1] = min(dp[j] + (0 if ca == cb else 1), dp2[j] + 1, dp[j + 1] + 1)
dp = dp2
return dp[-1]
def max_units_in_reads(fasta_path: Path, unit: str) -> int:
unit_bytes = unit.upper().encode()
unit_len = len(unit_bytes)
max_count = 0
with open(fasta_path, "rb") as fh:
seq = b""
for line in fh:
line = line.rstrip()
if line.startswith(b">"):
if seq:
n, i = 0, 0
s = seq.upper()
while True:
i = s.find(unit_bytes, i)
if i == -1:
break
n += 1
i += unit_len
max_count = max(max_count, n)
seq = b""
else:
seq += line
if seq:
n, i = 0, 0
s = seq.upper()
while True:
i = s.find(unit_bytes, i)
if i == -1:
break
n += 1
i += unit_len
max_count = max(max_count, n)
return max_count
def parse_header_reads(header: str) -> Optional[int]:
m = re.search(r'\breads=(\d+)', header)
return int(m.group(1)) if m else None
def count_repeat(seq: bytes, unit: str) -> int:
u = unit.upper().encode()
s = seq.upper()
n, i = 0, 0
while True:
i = s.find(u, i)
if i == -1:
break
n += 1
i += len(u)
return n
def extract_allele_by_unit(cons: bytes, unit: str) -> Optional[bytes]:
u = unit.upper().encode()
s = cons.upper()
n = len(u)
positions: list[int] = []
i = 0
while True:
i = s.find(u, i)
if i == -1:
break
positions.append(i)
i += n
if not positions:
return None
if len(positions) == 1:
return cons[positions[0]: positions[0] + n]
clusters: list[list[int]] = [[positions[0]]]
for p in positions[1:]:
if p - clusters[-1][-1] <= 4 * n:
clusters[-1].append(p)
else:
clusters.append([p])
best = max(clusters, key=len)
return cons[best[0]: best[-1] + n]
def find_allele_in_consensus(cons: bytes, anchor_k: int = 20) -> Optional[bytes]:
la = LEFT_FLANK[-anchor_k:].encode()
ra = RIGHT_FLANK[:anchor_k].encode()
upper = cons.upper()
ri = upper.find(ra.upper())
if ri == -1:
return None
li = upper[:ri].rfind(la.upper())
if li == -1 or ri <= li + anchor_k:
return None
return cons[li + anchor_k: ri]
def parse_fasta(path: Path) -> list[tuple[str, bytes]]:
records, header, seq = [], None, []
with open(path, "rb") as fh:
for line in fh:
line = line.rstrip()
if line.startswith(b">"):
if header is not None:
records.append((header, b"".join(seq)))
header = line[1:].decode()
seq = []
else:
seq.append(line)
if header is not None:
records.append((header, b"".join(seq)))
return records
def build_reference(scenario: Scenario, work: Path) -> Path:
fl = scenario.flank_len
ref = work / "reference.fa"
with open(ref, "w") as fh:
for i, allele in enumerate(scenario.alleles):
left = _rotated_flank(_LEFT_FLANK_FULL, fl, i)
right = _rotated_flank(_RIGHT_FLANK_FULL, fl, i)
seq = left + allele.seq + right
fh.write(f">allele_{i} unit={allele.unit} count={allele.count}\n{seq}\n")
return ref
def simulate_reads(scenario: Scenario, ref: Path, work: Path, model: dict) -> Path:
allele_len_max = max(len(a.seq) for a in scenario.alleles)
ref_len = scenario.flank_len * 2 + allele_len_max
if scenario.length_mean is not None:
length_mean = scenario.length_mean
else:
length_mean = max(500, int(ref_len * 0.75))
length_mean = min(length_mean, ref_len - 100) length_sd = max(50, length_mean // 5)
prefix = str(work / "sim")
cmd = [
PBSIM,
"--strategy", "wgs",
"--genome", ref,
"--method", "errhmm",
"--errhmm", PBMODELS / model["errhmm"],
"--accuracy-mean", str(model["accuracy_mean"]),
"--difference-ratio", model["difference_ratio"],
"--depth", "1", "--length-mean", str(length_mean),
"--length-sd", str(length_sd),
"--seed", str(scenario.seed),
"--prefix", prefix,
]
merged = work / "reads.fq"
with open(merged, "wb") as out:
for i, allele in enumerate(scenario.alleles):
single_ref = work / f"ref_allele{i}.fa"
with open(single_ref, "w") as fh:
left = _rotated_flank(_LEFT_FLANK_FULL, FLANK_LEN, i)
right = _rotated_flank(_RIGHT_FLANK_FULL, FLANK_LEN, i)
seq = left + allele.seq + right
fh.write(f">allele_{i}\n{seq}\n")
ap = str(work / f"sim_a{i}")
sim_cmd = [
PBSIM,
"--strategy", "wgs",
"--genome", single_ref,
"--method", "errhmm",
"--errhmm", PBMODELS / model["errhmm"],
"--accuracy-mean", str(model["accuracy_mean"]),
"--difference-ratio", model["difference_ratio"],
"--depth", str(allele.depth),
"--length-mean", str(length_mean),
"--length-sd", str(length_sd),
"--seed", str(scenario.seed + i),
"--prefix", ap,
]
run(sim_cmd)
fq_gz = work / f"sim_a{i}_0001.fq.gz"
if fq_gz.exists():
with gzip.open(fq_gz) as fq_in:
out.write(fq_in.read())
return merged
def align_and_index(reads: Path, ref: Path, work: Path, mm2_preset: str) -> Path:
bam = work / "aligned.bam"
mm2 = subprocess.Popen(
["minimap2", "-a", f"-x{mm2_preset}", "--secondary=no", str(ref), str(reads)],
stdout=subprocess.PIPE, stderr=subprocess.DEVNULL,
)
with open(bam, "wb") as bam_fh:
subprocess.run(
["samtools", "sort", "-"],
stdin=mm2.stdout, stdout=bam_fh, stderr=subprocess.DEVNULL, check=True,
)
mm2.wait()
run(["samtools", "index", bam])
return bam
def extract_reads(scenario: Scenario, bam: Path, work: Path) -> Path:
bed = work / "targets.bed"
min_lens: dict[str, int] = {}
_flank_unit_len = len(_U) fl = scenario.flank_len
with open(bed, "w") as fh:
for i, allele in enumerate(scenario.alleles):
start = max(0, fl - ANCHOR_PAD)
end = fl + len(allele.seq) + ANCHOR_PAD
fh.write(f"allele_{i}\t{start}\t{end}\n")
if len(allele.seq) < _flank_unit_len:
min_lens[f"allele_{i}"] = ANCHOR_PAD + len(allele.seq)
else:
min_lens[f"allele_{i}"] = ANCHOR_PAD // 2
raw = work / "extracted_raw.fa"
bedpull_cmd = ["bedpull", "-b", bam, "-r", bed, "-o", raw]
if scenario.partial:
bedpull_cmd.append("--partial")
run(bedpull_cmd)
extracted = work / "extracted.fa"
kept = total = 0
with open(raw) as fin, open(extracted, "w") as fout:
header = seq = ""
for line in fin:
line = line.rstrip()
if line.startswith(">"):
if header and seq:
total += 1
allele_id = header.split("|")[1].split(":")[0] if "|" in header else "allele_0"
min_len = min_lens.get(allele_id, ANCHOR_PAD)
if len(seq) >= min_len:
fout.write(f">{header}\n{seq}\n")
kept += 1
header = line[1:]
seq = ""
else:
seq += line
if header and seq:
total += 1
allele_id = header.split("|")[1].split(":")[0] if "|" in header else "allele_0"
min_len = min_lens.get(allele_id, ANCHOR_PAD)
if len(seq) >= min_len:
fout.write(f">{header}\n{seq}\n")
kept += 1
return extracted
def run_consensus(extracted: Path, work: Path, scenario: Scenario) -> Optional[Path]:
consensus = work / "consensus.fa"
band_args = ["--band-width", "0"] if scenario.multi else ["--band-width", "50"]
cmd = (
[str(POA_BIN)] + band_args
+ (["--multi"] if scenario.multi else [])
+ scenario.poa_flags
+ ["--min-reads", "2", str(extracted)]
)
try:
result = subprocess.run(cmd, capture_output=True, check=True)
consensus.write_bytes(result.stdout)
if result.stderr:
sys.stderr.buffer.write(result.stderr)
return consensus
except subprocess.CalledProcessError as exc:
if exc.stderr:
sys.stderr.buffer.write(exc.stderr)
return None
def evaluate(scenario: Scenario, consensus_path: Optional[Path],
extracted_path: Optional[Path] = None) -> dict:
result = {
"scenario": scenario.name,
"model": scenario.model,
"alleles": "+".join(f"{a.unit}×{a.count}" for a in scenario.alleles),
"total_depth": sum(a.depth for a in scenario.alleles),
}
if consensus_path is None or not consensus_path.exists():
result["status"] = "FAILED (no consensus)"
return result
records = parse_fasta(consensus_path)
if not records:
result["status"] = "FAILED (empty output)"
return result
if not scenario.is_multi:
cons_seq = records[0][1]
truth = scenario.alleles[0]
extracted = extract_allele_by_unit(cons_seq, truth.unit)
result["units_truth"] = truth.count
if extracted is None:
result["anchor_found"] = False
result["cons_len"] = len(cons_seq)
result["truth_len"] = len(truth.seq)
result["delta_len"] = len(cons_seq) - len(truth.seq)
result["units_found"] = 0
result["delta_units"] = -truth.count
ok = False
else:
result["anchor_found"] = True
result["cons_len"] = len(extracted)
result["truth_len"] = len(truth.seq)
result["delta_len"] = len(extracted) - len(truth.seq)
result["units_found"] = count_repeat(extracted, truth.unit)
result["delta_units"] = result["units_found"] - truth.count
result["edit_dist"] = levenshtein(extracted.upper(), truth.seq.upper().encode())
unit_tol = max(1, truth.count // 50)
edit_tol = max(3, len(truth.seq) // 25)
ok = (abs(result["delta_units"]) <= unit_tol) and result["edit_dist"] <= edit_tol
if ok:
result["status"] = "OK"
elif extracted_path is not None:
max_vis = max_units_in_reads(extracted_path, truth.unit)
result["max_visible_units"] = max_vis
if max_vis < truth.count:
result["status"] = f"OK (read_limit: max_visible={max_vis}/{truth.count})"
else:
result["status"] = "FAIL"
else:
result["status"] = "FAIL"
else:
truths = scenario.alleles
if len(records) != len(truths):
result["status"] = f"FAIL (expected {len(truths)} alleles, got {len(records)})"
result["n_alleles_found"] = len(records)
result["n_alleles_truth"] = len(truths)
return result
unit = truths[0].unit
found_units_list = [count_repeat(cons_seq, unit) for _, cons_seq in records]
best_assignment = None
best_cost = float('inf')
for perm in permutations(range(len(truths))):
cost = sum(abs(found_units_list[i] - truths[perm[i]].count) for i in range(len(records)))
if cost < best_cost:
best_cost = cost
best_assignment = perm
allele_read_counts = [parse_header_reads(hdr) for hdr, _ in records]
matched = []
for i, truth_idx in enumerate(best_assignment):
t = truths[truth_idx]
found = found_units_list[i]
m = {
"truth_units": t.count,
"found_units": found,
"delta_units": found - t.count,
}
if allele_read_counts[i] is not None:
m["allele_reads"] = allele_read_counts[i]
matched.append(m)
result["allele_results"] = matched
unit_tol = max(1, max(t.count for t in truths) // 50)
all_ok = all(abs(m["delta_units"]) <= max(1, m["truth_units"] // 50) for m in matched)
if all_ok:
result["status"] = "OK"
else:
_MIN_RELIABLE_ALLELE_DEPTH = 15
failing_are_all_depth_limited = all(
abs(m["delta_units"]) <= max(1, m["truth_units"] // 50)
or m.get("allele_reads", _MIN_RELIABLE_ALLELE_DEPTH) < _MIN_RELIABLE_ALLELE_DEPTH
for m in matched
)
if failing_are_all_depth_limited:
low_depth_notes = [
f"{m['truth_units']}× ({m.get('allele_reads', '?')} reads)"
for m in matched
if abs(m["delta_units"]) > max(1, m["truth_units"] // 50)
]
result["status"] = f"OK (depth_limit: {', '.join(low_depth_notes)})"
else:
result["status"] = "FAIL"
return result
def print_result(r: dict):
if "allele_results" in r:
allele_str = " ".join(
f"({m['truth_units']}→{m['found_units']} Δ{m['delta_units']:+d})"
for m in r["allele_results"]
)
print(f" {r['scenario']:<35} {r['model']:<10} d={r['total_depth']:<3} "
f"multi: {allele_str} {r['status']}")
else:
unit_info = f"units: {r.get('units_truth','?')}→{r.get('units_found','?')} (Δ{r.get('delta_units',0):+d})"
edit_info = f"edit={r.get('edit_dist', '?')}" if "edit_dist" in r else f"Δlen={r.get('delta_len',0):+d}"
anchor = "⚓" if r.get("anchor_found") else "~"
print(f" {r['scenario']:<35} {r['model']:<10} d={r['total_depth']:<3} "
f"{unit_info} {edit_info} {anchor} {r['status']}")
def main():
ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter)
ap.add_argument("--scenarios", nargs="*", help="Scenario names to run (default: all)")
ap.add_argument("--workdir", default="bench/work", help="Working directory for temp files")
ap.add_argument("--keep", action="store_true", help="Keep working files after run")
ap.add_argument("--long-reads", action="store_true",
help="Include 15 kb long-read scenarios (slow; requires long pbsim runs)")
args = ap.parse_args()
for binary, name in [(POA_BIN, "poa-consensus"), ("minimap2", "minimap2"),
("samtools", "samtools"), ("bedpull", "bedpull"), (PBSIM, "pbsim")]:
path = shutil.which(str(binary)) or (Path(str(binary)).exists() and str(binary))
if not path:
sys.exit(f"error: {name} not found ({binary}); build CLI with: cargo build --release --features cli")
to_run = SCENARIOS
if not args.long_reads:
to_run = [s for s in to_run if not s.long_read]
if args.scenarios:
to_run = [s for s in to_run if s.name in args.scenarios]
unknown = set(args.scenarios) - {s.name for s in SCENARIOS}
if unknown:
print(f"warning: unknown scenarios: {', '.join(sorted(unknown))}", file=sys.stderr)
workdir = Path(args.workdir)
workdir.mkdir(parents=True, exist_ok=True)
print(f"\n{'='*80}")
print(f" poa-consensus validation ({len(to_run)} scenarios)")
print(f"{'='*80}")
print(f" {'SCENARIO':<35} {'MODEL':<10} {'DEPTH':<6} ACCURACY STATUS")
print(f" {'-'*35} {'-'*10} {'-'*5} {'─'*30} ──────")
results, n_ok, n_fail = [], 0, 0
for scenario in to_run:
work = workdir / scenario.name
if work.exists():
shutil.rmtree(work)
work.mkdir(parents=True)
model = ERRMODELS[scenario.model]
try:
ref = build_reference(scenario, work)
reads = simulate_reads(scenario, ref, work, model)
bam = align_and_index(reads, ref, work, model["mm2_preset"])
extracted = extract_reads(scenario, bam, work)
consensus = run_consensus(extracted, work, scenario)
r = evaluate(scenario, consensus, extracted)
except Exception as exc:
r = {"scenario": scenario.name, "model": scenario.model,
"total_depth": sum(a.depth for a in scenario.alleles),
"alleles": "+".join(f"{a.unit}×{a.count}" for a in scenario.alleles),
"status": f"ERROR: {exc}"}
results.append(r)
print_result(r)
if r["status"].startswith("OK"):
n_ok += 1
else:
n_fail += 1
if not args.keep:
shutil.rmtree(work, ignore_errors=True)
print(f"\n Result: {n_ok} passed, {n_fail} failed out of {len(to_run)} scenarios")
print()
if __name__ == "__main__":
main()