← Files NGS Analysis WorkbenchARCHIVED FILE
workflows/bulk_rnaseq_counts_qc/workflow/scripts/aggregate_salmon_quant.py
10.4 KB · Sep 30, 2026 · 23:20 UTC
#!/usr/bin/env python3
"""Aggregate Salmon quant.sf files into transcript- and gene-level matrices."""
from __future__ import annotations
import argparse
import csv
import gzip
import json
from pathlib import Path
MIN_TX2GENE_COVERAGE = 0.95
def parse_gtf_attributes(raw: str) -> dict[str, str]:
values: dict[str, str] = {}
for chunk in raw.strip().split(";"):
part = chunk.strip()
if not part or " " not in part:
continue
key, value = part.split(" ", 1)
values[key] = value.strip().strip('"')
return values
def tx2gene_from_gtf(path: Path | None) -> dict[str, dict[str, str]]:
if path is None:
return {}
mapping: dict[str, dict[str, str]] = {}
opener = gzip.open if path.name.endswith(".gz") else open
with opener(path, "rt", encoding="utf-8", errors="replace") as handle:
for line in handle:
if not line or line.startswith("#"):
continue
fields = line.rstrip("\n").split("\t")
if len(fields) < 9 or fields[2] not in {"transcript", "exon"}:
continue
attrs = parse_gtf_attributes(fields[8])
transcript_id = attrs.get("transcript_id")
gene_id = attrs.get("gene_id")
if not transcript_id or not gene_id:
continue
gene_name = attrs.get("gene_name")
if transcript_id not in mapping or gene_name is not None:
mapping[transcript_id] = {
"gene_id": gene_id,
"gene_name": gene_name or gene_id,
}
return mapping
def canonical_transcript_id(transcript_id: str) -> str:
"""Return the stable transcript identifier used to join Salmon and GTF records."""
primary_id = transcript_id.strip().split("|", 1)[0]
unversioned_id, separator, version = primary_id.rpartition(".")
if separator and unversioned_id and version.isdigit():
return unversioned_id
return primary_id
def resolve_tx2gene(
transcript_ids: list[str], tx2gene: dict[str, dict[str, str]]
) -> dict[str, dict[str, str]]:
"""Resolve quant transcript IDs exactly, then by an unambiguous canonical ID."""
canonical_lookup: dict[str, dict[str, str]] = {}
ambiguous_ids: set[str] = set()
for transcript_id, gene_record in tx2gene.items():
canonical_id = canonical_transcript_id(transcript_id)
existing_record = canonical_lookup.get(canonical_id)
if existing_record and existing_record["gene_id"] != gene_record["gene_id"]:
ambiguous_ids.add(canonical_id)
canonical_lookup.pop(canonical_id, None)
elif canonical_id not in ambiguous_ids:
canonical_lookup[canonical_id] = gene_record
resolved: dict[str, dict[str, str]] = {}
for transcript_id in transcript_ids:
gene_record = tx2gene.get(transcript_id)
if gene_record is None:
gene_record = canonical_lookup.get(canonical_transcript_id(transcript_id))
if gene_record is not None:
resolved[transcript_id] = gene_record
return resolved
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--config", required=True)
parser.add_argument("--outdir", required=True)
parser.add_argument("--quant", action="append", default=[], help="sample=/path/to/quant.sf")
return parser.parse_args()
def read_quant_sf(path: Path) -> dict[str, dict[str, str]]:
with path.open(newline="", encoding="utf-8") as handle:
reader = csv.DictReader(handle, delimiter="\t")
return {row["Name"]: row for row in reader}
def write_matrix(path: Path, header: list[str], rows: list[list[str]]) -> None:
with path.open("w", newline="", encoding="utf-8") as handle:
writer = csv.writer(handle, delimiter="\t")
writer.writerow(header)
writer.writerows(rows)
def main() -> int:
args = parse_args()
config = json.loads(Path(args.config).read_text(encoding="utf-8"))
outdir = Path(args.outdir)
outdir.mkdir(parents=True, exist_ok=True)
sample_to_quant: dict[str, Path] = {}
for item in args.quant:
sample, raw_path = item.split("=", 1)
sample_to_quant[sample] = Path(raw_path)
sample_names = sorted(sample_to_quant)
per_sample = {sample: read_quant_sf(path) for sample, path in sample_to_quant.items()}
transcript_ids = sorted({tx for table in per_sample.values() for tx in table})
gtf_path = (
Path(config["references"]["annotation_gtf"])
if config.get("references", {}).get("annotation_gtf")
else None
)
tx2gene = tx2gene_from_gtf(gtf_path)
resolved_tx2gene = resolve_tx2gene(transcript_ids, tx2gene)
tpm_rows: list[list[str]] = []
num_reads_rows: list[list[str]] = []
effective_length_rows: list[list[str]] = []
for transcript_id in transcript_ids:
tpm_row = [transcript_id]
num_reads_row = [transcript_id]
effective_length_row = [transcript_id]
for sample in sample_names:
record = per_sample[sample].get(transcript_id)
tpm_row.append(record["TPM"] if record else "")
num_reads_row.append(record["NumReads"] if record else "")
effective_length_row.append(record["EffectiveLength"] if record else "")
tpm_rows.append(tpm_row)
num_reads_rows.append(num_reads_row)
effective_length_rows.append(effective_length_row)
write_matrix(outdir / "tpm.tsv", ["transcript_id", *sample_names], tpm_rows)
write_matrix(outdir / "num_reads.tsv", ["transcript_id", *sample_names], num_reads_rows)
write_matrix(
outdir / "effective_length.tsv", ["transcript_id", *sample_names], effective_length_rows
)
mapped_transcripts = [
transcript_id for transcript_id in transcript_ids if transcript_id in resolved_tx2gene
]
unmapped_transcripts = [
transcript_id for transcript_id in transcript_ids if transcript_id not in resolved_tx2gene
]
transcript_coverage = len(mapped_transcripts) / len(transcript_ids) if transcript_ids else 0.0
per_sample_read_coverage: dict[str, float | None] = {}
for sample in sample_names:
total_reads = sum(float(record["NumReads"]) for record in per_sample[sample].values())
mapped_reads = sum(
float(per_sample[sample][transcript_id]["NumReads"])
for transcript_id in mapped_transcripts
if transcript_id in per_sample[sample]
)
per_sample_read_coverage[sample] = mapped_reads / total_reads if total_reads > 0 else None
available_read_coverage = [
coverage for coverage in per_sample_read_coverage.values() if coverage is not None
]
minimum_read_coverage = min(available_read_coverage, default=None)
gene_outputs_eligible = bool(tx2gene) and (
transcript_coverage >= MIN_TX2GENE_COVERAGE
and (minimum_read_coverage is None or minimum_read_coverage >= MIN_TX2GENE_COVERAGE)
)
coverage = {
"threshold": MIN_TX2GENE_COVERAGE,
"total_transcripts": len(transcript_ids),
"mapped_transcripts": len(mapped_transcripts),
"unmapped_transcripts": len(unmapped_transcripts),
"transcript_coverage_fraction": transcript_coverage,
"per_sample_num_reads_coverage_fraction": per_sample_read_coverage,
"minimum_num_reads_coverage_fraction": minimum_read_coverage,
"gene_level_artifacts_created": gene_outputs_eligible,
"unmapped_transcript_preview": unmapped_transcripts[:50],
}
(outdir / "tx2gene_coverage.json").write_text(
json.dumps(coverage, indent=2, sort_keys=True) + "\n", encoding="utf-8"
)
if not gene_outputs_eligible:
for filename in ("tx2gene.tsv", "gene_num_reads.tsv", "gene_tpm.tsv"):
(outdir / filename).unlink(missing_ok=True)
tx2gene_rows: list[list[str]] = []
gene_num_reads: dict[str, list[float]] = {}
gene_tpm: dict[str, list[float]] = {}
for transcript_id in mapped_transcripts:
gene_record = resolved_tx2gene[transcript_id]
gene_id = gene_record["gene_id"]
gene_name = gene_record["gene_name"]
tx2gene_rows.append([transcript_id, gene_id, gene_name])
gene_num_reads.setdefault(gene_id, [0.0] * len(sample_names))
gene_tpm.setdefault(gene_id, [0.0] * len(sample_names))
for idx, sample in enumerate(sample_names):
record = per_sample[sample].get(transcript_id)
if not record:
continue
gene_num_reads[gene_id][idx] += float(record["NumReads"])
gene_tpm[gene_id][idx] += float(record["TPM"])
if gene_outputs_eligible:
write_matrix(
outdir / "tx2gene.tsv", ["transcript_id", "gene_id", "gene_name"], tx2gene_rows
)
write_matrix(
outdir / "gene_num_reads.tsv",
["gene_id", *sample_names],
[
[gene_id, *[f"{value:.6f}" for value in gene_num_reads[gene_id]]]
for gene_id in sorted(gene_num_reads)
],
)
write_matrix(
outdir / "gene_tpm.tsv",
["gene_id", *sample_names],
[
[gene_id, *[f"{value:.6f}" for value in gene_tpm[gene_id]]]
for gene_id in sorted(gene_tpm)
],
)
configured_samples = config.get("rnaseq_salmon_samples") or config.get("samples", {})
sample_rows = []
for sample in sample_names:
info = configured_samples[sample]
r1 = info.get("r1", [])
r2 = info.get("r2", [])
r1 = [r1] if isinstance(r1, str) else r1
r2 = [r2] if isinstance(r2, str) else r2
sample_rows.append(
[
sample,
info.get("layout", "PE" if r2 else "SE"),
info.get("strandedness", "unknown"),
info.get("salmon_libtype", ""),
info.get("salmon_libtype_source", ""),
",".join(str(index) for index in info.get("row_indices", [])),
str(len(r1)),
str(len(r2)),
]
)
write_matrix(
outdir / "samples.tsv",
[
"sample",
"layout",
"strandedness",
"salmon_libtype",
"salmon_libtype_source",
"technical_replicate_rows",
"fastq_1_files",
"fastq_2_files",
],
sample_rows,
)
return 0
if __name__ == "__main__":
raise SystemExit(main())
SHA-256: d5831c9e38cd70fd860ebe77b9bf355d452ebc3b2c5d72e65794dc9a35342dc4