parallel: auto threads (n-1, capped at 8)
- --threads/-t (default auto = detected cores -1, cap 8) - per-chromosome multiprocessing.Pool, FASTA caching - robust nested dict handling for merged profiles
This commit is contained in:
+257
-22
@@ -39,8 +39,9 @@ Outputs (in --outdir):
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import argparse
|
import argparse
|
||||||
|
import multiprocessing
|
||||||
import os
|
import os
|
||||||
from collections import defaultdict
|
from collections import Counter, defaultdict
|
||||||
|
|
||||||
import pysam
|
import pysam
|
||||||
import pandas as pd
|
import pandas as pd
|
||||||
@@ -53,6 +54,200 @@ SUBSTITUTIONS = [f"{r}>{a}" for r in "ACGT" for a in "ACGT" if r != a]
|
|||||||
END_SIZES = [1, 3, 5, 10, 20]
|
END_SIZES = [1, 3, 5, 10, 20]
|
||||||
|
|
||||||
|
|
||||||
|
def get_default_threads(cap=8):
|
||||||
|
try:
|
||||||
|
n = len(os.sched_getaffinity(0))
|
||||||
|
except AttributeError:
|
||||||
|
n = os.cpu_count() or 1
|
||||||
|
if n is None:
|
||||||
|
n = 1
|
||||||
|
n = max(1, n - 1)
|
||||||
|
return min(n, cap)
|
||||||
|
|
||||||
|
|
||||||
|
def _process_shard(task):
|
||||||
|
bam_path, ref_path, chrom, mapq, baseq, include_dup = task
|
||||||
|
bam = pysam.AlignmentFile(bam_path, "rb")
|
||||||
|
fasta = pysam.FastaFile(ref_path)
|
||||||
|
try:
|
||||||
|
chrom_seq = fasta.fetch(chrom)
|
||||||
|
except Exception:
|
||||||
|
chrom_seq = None
|
||||||
|
|
||||||
|
sub_counts = Counter()
|
||||||
|
ref_base_counts = Counter()
|
||||||
|
profile = {}
|
||||||
|
profile3 = {}
|
||||||
|
strand_profile = {}
|
||||||
|
strand_profile3 = {}
|
||||||
|
genomic = {}
|
||||||
|
total_reads = 0
|
||||||
|
used_reads = 0
|
||||||
|
usable_bases = 0
|
||||||
|
|
||||||
|
def _inc(store, k1, k2, k3, k4):
|
||||||
|
a = store.setdefault(k1, {})
|
||||||
|
b = a.setdefault(k2, {})
|
||||||
|
c = b.setdefault(k3, {})
|
||||||
|
c[k4] = c.get(k4, 0) + 1
|
||||||
|
|
||||||
|
for read in bam.fetch(chrom):
|
||||||
|
total_reads += 1
|
||||||
|
if read.is_secondary or read.is_supplementary:
|
||||||
|
continue
|
||||||
|
if read.mapping_quality < mapq:
|
||||||
|
continue
|
||||||
|
if read.is_duplicate and not include_dup:
|
||||||
|
continue
|
||||||
|
seq = read.query_sequence
|
||||||
|
quals = read.query_qualities
|
||||||
|
if seq is None or quals is None:
|
||||||
|
continue
|
||||||
|
used_reads += 1
|
||||||
|
read_label = "R1" if read.is_read1 else "R2" if read.is_read2 else "single"
|
||||||
|
strand = "-" if read.is_reverse else "+"
|
||||||
|
read_len = read.query_length
|
||||||
|
chrom_name = bam.get_reference_name(read.reference_id)
|
||||||
|
seq_chrom = chrom_seq if chrom_name == chrom else None
|
||||||
|
if seq_chrom is None:
|
||||||
|
try:
|
||||||
|
seq_chrom = fasta.fetch(chrom_name)
|
||||||
|
except Exception:
|
||||||
|
seq_chrom = None
|
||||||
|
for query_pos, ref_pos in read.get_aligned_pairs(matches_only=True):
|
||||||
|
if quals[query_pos] < baseq:
|
||||||
|
continue
|
||||||
|
if seq_chrom is not None and 0 <= ref_pos < len(seq_chrom):
|
||||||
|
ref_base = seq_chrom[ref_pos].upper()
|
||||||
|
else:
|
||||||
|
try:
|
||||||
|
ref_base = fasta.fetch(chrom_name, ref_pos, ref_pos + 1).upper()
|
||||||
|
except Exception:
|
||||||
|
continue
|
||||||
|
if ref_base not in "ACGT":
|
||||||
|
continue
|
||||||
|
alt_base = seq[query_pos].upper()
|
||||||
|
if alt_base not in "ACGT":
|
||||||
|
continue
|
||||||
|
usable_bases += 1
|
||||||
|
pos5 = query_pos + 1
|
||||||
|
pos3 = read_len - query_pos
|
||||||
|
ref_base_counts[ref_base] += 1
|
||||||
|
key = (chrom_name, ref_pos)
|
||||||
|
g = genomic.get(key)
|
||||||
|
if g is None:
|
||||||
|
g = {"ref": ref_base, "depth": 0, "depth_plus": 0, "depth_minus": 0, "alts": {}}
|
||||||
|
genomic[key] = g
|
||||||
|
else:
|
||||||
|
g["ref"] = ref_base
|
||||||
|
g["depth"] += 1
|
||||||
|
if strand == "+":
|
||||||
|
g["depth_plus"] += 1
|
||||||
|
else:
|
||||||
|
g["depth_minus"] += 1
|
||||||
|
_inc(profile, read_label, pos5, ref_base, "total")
|
||||||
|
_inc(profile3, read_label, pos3, ref_base, "total")
|
||||||
|
_inc(strand_profile, strand, pos5, ref_base, "total")
|
||||||
|
_inc(strand_profile3, strand, pos3, ref_base, "total")
|
||||||
|
if ref_base == alt_base:
|
||||||
|
continue
|
||||||
|
sub = f"{ref_base}>{alt_base}"
|
||||||
|
sub_counts[sub] += 1
|
||||||
|
_inc(profile, read_label, pos5, ref_base, sub)
|
||||||
|
_inc(profile3, read_label, pos3, ref_base, sub)
|
||||||
|
_inc(strand_profile, strand, pos5, ref_base, sub)
|
||||||
|
_inc(strand_profile3, strand, pos3, ref_base, sub)
|
||||||
|
alts = g["alts"]
|
||||||
|
alt = alts.get(alt_base)
|
||||||
|
if alt is None:
|
||||||
|
alt = {"count": 0, "bq_sum": 0, "pos5_sum": 0, "pos3_sum": 0, "plus": 0, "minus": 0, "r1": 0, "r2": 0}
|
||||||
|
alts[alt_base] = alt
|
||||||
|
alt["count"] += 1
|
||||||
|
alt["bq_sum"] += quals[query_pos]
|
||||||
|
alt["pos5_sum"] += pos5
|
||||||
|
alt["pos3_sum"] += pos3
|
||||||
|
if strand == "+":
|
||||||
|
alt["plus"] += 1
|
||||||
|
else:
|
||||||
|
alt["minus"] += 1
|
||||||
|
if read_label == "R1":
|
||||||
|
alt["r1"] += 1
|
||||||
|
elif read_label == "R2":
|
||||||
|
alt["r2"] += 1
|
||||||
|
bam.close()
|
||||||
|
fasta.close()
|
||||||
|
return {
|
||||||
|
"total_reads": total_reads,
|
||||||
|
"used_reads": used_reads,
|
||||||
|
"usable_bases": usable_bases,
|
||||||
|
"sub_counts": dict(sub_counts),
|
||||||
|
"ref_base_counts": dict(ref_base_counts),
|
||||||
|
"profile": profile,
|
||||||
|
"profile3": profile3,
|
||||||
|
"strand_profile": strand_profile,
|
||||||
|
"strand_profile3": strand_profile3,
|
||||||
|
"genomic": genomic,
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
def _merge_nested(target, source):
|
||||||
|
for k1, v1 in source.items():
|
||||||
|
t1 = target.setdefault(k1, {})
|
||||||
|
for k2, v2 in v1.items():
|
||||||
|
t2 = t1.setdefault(k2, {})
|
||||||
|
for k3, v3 in v2.items():
|
||||||
|
t3 = t2.setdefault(k3, {})
|
||||||
|
for k4, cnt in v3.items():
|
||||||
|
t3[k4] = t3.get(k4, 0) + cnt
|
||||||
|
|
||||||
|
|
||||||
|
def _merge_results(partials):
|
||||||
|
merged = {
|
||||||
|
"total_reads": 0,
|
||||||
|
"used_reads": 0,
|
||||||
|
"usable_bases": 0,
|
||||||
|
"sub_counts": Counter(),
|
||||||
|
"ref_base_counts": Counter(),
|
||||||
|
"profile": {},
|
||||||
|
"profile3": {},
|
||||||
|
"strand_profile": {},
|
||||||
|
"strand_profile3": {},
|
||||||
|
"genomic": {},
|
||||||
|
}
|
||||||
|
for p in partials:
|
||||||
|
merged["total_reads"] += p["total_reads"]
|
||||||
|
merged["used_reads"] += p["used_reads"]
|
||||||
|
merged["usable_bases"] += p["usable_bases"]
|
||||||
|
merged["sub_counts"].update(p["sub_counts"])
|
||||||
|
merged["ref_base_counts"].update(p["ref_base_counts"])
|
||||||
|
_merge_nested(merged["profile"], p["profile"])
|
||||||
|
_merge_nested(merged["profile3"], p["profile3"])
|
||||||
|
_merge_nested(merged["strand_profile"], p["strand_profile"])
|
||||||
|
_merge_nested(merged["strand_profile3"], p["strand_profile3"])
|
||||||
|
for key, g in p["genomic"].items():
|
||||||
|
mg = merged["genomic"].get(key)
|
||||||
|
if mg is None:
|
||||||
|
merged["genomic"][key] = g
|
||||||
|
else:
|
||||||
|
mg["depth"] += g["depth"]
|
||||||
|
mg["depth_plus"] += g["depth_plus"]
|
||||||
|
mg["depth_minus"] += g["depth_minus"]
|
||||||
|
for alt_base, alt in g["alts"].items():
|
||||||
|
malt = mg["alts"].get(alt_base)
|
||||||
|
if malt is None:
|
||||||
|
mg["alts"][alt_base] = alt
|
||||||
|
else:
|
||||||
|
malt["count"] += alt["count"]
|
||||||
|
malt["bq_sum"] += alt["bq_sum"]
|
||||||
|
malt["pos5_sum"] += alt["pos5_sum"]
|
||||||
|
malt["pos3_sum"] += alt["pos3_sum"]
|
||||||
|
malt["plus"] += alt["plus"]
|
||||||
|
malt["minus"] += alt["minus"]
|
||||||
|
malt["r1"] += alt["r1"]
|
||||||
|
malt["r2"] += alt["r2"]
|
||||||
|
return merged
|
||||||
|
|
||||||
|
|
||||||
def parse_args():
|
def parse_args():
|
||||||
ap = argparse.ArgumentParser(description="FFPE damage analysis")
|
ap = argparse.ArgumentParser(description="FFPE damage analysis")
|
||||||
ap.add_argument("--bam", required=True)
|
ap.add_argument("--bam", required=True)
|
||||||
@@ -74,6 +269,9 @@ def parse_args():
|
|||||||
"counts as end-associated (default 10)")
|
"counts as end-associated (default 10)")
|
||||||
ap.add_argument("--filter-strand-p", type=float, default=0.05,
|
ap.add_argument("--filter-strand-p", type=float, default=0.05,
|
||||||
help="Fisher p-value threshold for strand bias (default 0.05)")
|
help="Fisher p-value threshold for strand bias (default 0.05)")
|
||||||
|
ap.add_argument("--threads", "-t", type=int, default=None,
|
||||||
|
help="threads for BAM processing (default auto = "
|
||||||
|
"detected cores -1, capped at 8; 1 = single-thread)")
|
||||||
return ap.parse_args()
|
return ap.parse_args()
|
||||||
|
|
||||||
|
|
||||||
@@ -229,14 +427,14 @@ def write_read_position_profiles(d, args):
|
|||||||
for label in store:
|
for label in store:
|
||||||
for pos in sorted(store[label]):
|
for pos in sorted(store[label]):
|
||||||
for ref in "ACGT":
|
for ref in "ACGT":
|
||||||
total = store[label][pos][ref]["total"]
|
total = _get_nested(store, label, pos, ref, "total")
|
||||||
if total == 0:
|
if total == 0:
|
||||||
continue
|
continue
|
||||||
for alt in "ACGT":
|
for alt in "ACGT":
|
||||||
if alt == ref:
|
if alt == ref:
|
||||||
continue
|
continue
|
||||||
sub = f"{ref}>{alt}"
|
sub = f"{ref}>{alt}"
|
||||||
count = store[label][pos][ref][sub]
|
count = _get_nested(store, label, pos, ref, sub)
|
||||||
out.append({
|
out.append({
|
||||||
"read": label,
|
"read": label,
|
||||||
pos_key: pos,
|
pos_key: pos,
|
||||||
@@ -261,14 +459,18 @@ def write_read_position_profiles(d, args):
|
|||||||
return df5, df3
|
return df5, df3
|
||||||
|
|
||||||
|
|
||||||
|
def _get_nested(store, label, pos, ref, field):
|
||||||
|
return store.get(label, {}).get(pos, {}).get(ref, {}).get(field, 0)
|
||||||
|
|
||||||
|
|
||||||
def cg_rows(store, pos_key):
|
def cg_rows(store, pos_key):
|
||||||
out = []
|
out = []
|
||||||
for label in store:
|
for label in store:
|
||||||
for pos in sorted(store[label]):
|
for pos in sorted(store[label]):
|
||||||
c_total = store[label][pos]["C"]["total"]
|
c_total = _get_nested(store, label, pos, "C", "total")
|
||||||
ct = store[label][pos]["C"]["C>T"]
|
ct = _get_nested(store, label, pos, "C", "C>T")
|
||||||
g_total = store[label][pos]["G"]["total"]
|
g_total = _get_nested(store, label, pos, "G", "total")
|
||||||
ga = store[label][pos]["G"]["G>A"]
|
ga = _get_nested(store, label, pos, "G", "G>A")
|
||||||
out.append({
|
out.append({
|
||||||
"read": label,
|
"read": label,
|
||||||
pos_key: pos,
|
pos_key: pos,
|
||||||
@@ -310,10 +512,10 @@ def write_end_enrichment(d, args):
|
|||||||
for size in END_SIZES:
|
for size in END_SIZES:
|
||||||
ct_num = ct_den = ga_num = ga_den = 0
|
ct_num = ct_den = ga_num = ga_den = 0
|
||||||
for pos in range(1, size + 1):
|
for pos in range(1, size + 1):
|
||||||
ct_den += store[label][pos]["C"]["total"]
|
ct_den += _get_nested(store, label, pos, "C", "total")
|
||||||
ct_num += store[label][pos]["C"]["C>T"]
|
ct_num += _get_nested(store, label, pos, "C", "C>T")
|
||||||
ga_den += store[label][pos]["G"]["total"]
|
ga_den += _get_nested(store, label, pos, "G", "total")
|
||||||
ga_num += store[label][pos]["G"]["G>A"]
|
ga_num += _get_nested(store, label, pos, "G", "G>A")
|
||||||
rows.append({
|
rows.append({
|
||||||
"read": label,
|
"read": label,
|
||||||
"end": side,
|
"end": side,
|
||||||
@@ -592,27 +794,31 @@ def main():
|
|||||||
args = parse_args()
|
args = parse_args()
|
||||||
os.makedirs(args.outdir, exist_ok=True)
|
os.makedirs(args.outdir, exist_ok=True)
|
||||||
|
|
||||||
bam = pysam.AlignmentFile(args.bam, "rb")
|
threads = args.threads
|
||||||
if not bam.has_index():
|
if threads is None:
|
||||||
|
threads = get_default_threads()
|
||||||
|
if threads < 1:
|
||||||
|
threads = 1
|
||||||
|
|
||||||
|
bam_probe = pysam.AlignmentFile(args.bam, "rb")
|
||||||
|
has_index = bam_probe.has_index()
|
||||||
|
bam_probe.close()
|
||||||
|
if not has_index:
|
||||||
try:
|
try:
|
||||||
print(f"[info] BAM index not found, creating {args.bam}.bai ...")
|
print(f"[info] BAM index not found, creating {args.bam}.bai ...")
|
||||||
pysam.index(args.bam)
|
pysam.index(args.bam)
|
||||||
bam.close()
|
|
||||||
bam = pysam.AlignmentFile(args.bam, "rb")
|
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
bam.close()
|
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"BAM index (.bai) not found and auto-creation failed: {e}. "
|
f"BAM index (.bai) not found and auto-creation failed: {e}. "
|
||||||
f"Run: samtools index {args.bam}") from e
|
f"Run: samtools index {args.bam}") from e
|
||||||
try:
|
try:
|
||||||
fasta = pysam.FastaFile(args.reference)
|
fasta_probe = pysam.FastaFile(args.reference)
|
||||||
|
fasta_probe.close()
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
bam.close()
|
|
||||||
if "fai" in str(e).lower() or "index" in str(e).lower():
|
if "fai" in str(e).lower() or "index" in str(e).lower():
|
||||||
try:
|
try:
|
||||||
print(f"[info] FASTA index not found, creating {args.reference}.fai ...")
|
print(f"[info] FASTA index not found, creating {args.reference}.fai ...")
|
||||||
pysam.faidx(args.reference)
|
pysam.faidx(args.reference)
|
||||||
fasta = pysam.FastaFile(args.reference)
|
|
||||||
except Exception as e2:
|
except Exception as e2:
|
||||||
raise RuntimeError(
|
raise RuntimeError(
|
||||||
f"FASTA index (.fai) not found and auto-creation failed: {e2}. "
|
f"FASTA index (.fai) not found and auto-creation failed: {e2}. "
|
||||||
@@ -620,9 +826,38 @@ def main():
|
|||||||
else:
|
else:
|
||||||
raise
|
raise
|
||||||
|
|
||||||
d = process_bam(bam, fasta, args)
|
if threads == 1:
|
||||||
bam.close()
|
bam = pysam.AlignmentFile(args.bam, "rb")
|
||||||
fasta.close()
|
fasta = pysam.FastaFile(args.reference)
|
||||||
|
d = process_bam(bam, fasta, args)
|
||||||
|
bam.close()
|
||||||
|
fasta.close()
|
||||||
|
else:
|
||||||
|
print(f"[info] using {threads} threads (auto, capped at 8, 1 reserved)")
|
||||||
|
bam_tmp = pysam.AlignmentFile(args.bam, "rb")
|
||||||
|
chroms = list(bam_tmp.references)
|
||||||
|
bam_tmp.close()
|
||||||
|
if not chroms:
|
||||||
|
bam = pysam.AlignmentFile(args.bam, "rb")
|
||||||
|
fasta = pysam.FastaFile(args.reference)
|
||||||
|
d = process_bam(bam, fasta, args)
|
||||||
|
bam.close()
|
||||||
|
fasta.close()
|
||||||
|
else:
|
||||||
|
tasks = [(args.bam, args.reference, c, args.mapq, args.baseq,
|
||||||
|
args.include_duplicates) for c in chroms]
|
||||||
|
with multiprocessing.Pool(threads) as pool:
|
||||||
|
partials = pool.map(_process_shard, tasks)
|
||||||
|
d = _merge_results(partials)
|
||||||
|
try:
|
||||||
|
bam_cnt = pysam.AlignmentFile(args.bam, "rb")
|
||||||
|
total_via_idx = sum(s.mapped for s in bam_cnt.get_index_statistics())
|
||||||
|
total_via_idx += bam_cnt.unmapped
|
||||||
|
bam_cnt.close()
|
||||||
|
if total_via_idx > d["total_reads"]:
|
||||||
|
d["total_reads"] = total_via_idx
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
sub_df = write_substitution_summary(d, args)
|
sub_df = write_substitution_summary(d, args)
|
||||||
write_read_position_profiles(d, args)
|
write_read_position_profiles(d, args)
|
||||||
|
|||||||
Reference in New Issue
Block a user