File size: 8,533 Bytes
97d5f73 0d06695 97d5f73 f0190da 97d5f73 f0190da 97d5f73 a312d10 0d06695 a312d10 97d5f73 26b65f6 97d5f73 f0190da 97d5f73 a312d10 ceaecab a312d10 ceaecab a312d10 ceaecab 97d5f73 0d06695 a312d10 97d5f73 0d06695 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 | """Accession lookup and lazy, segment-level access to the local sample."""
from collections import defaultdict
import json
from pathlib import Path
import re
import numpy as np
import pandas as pd
import pyarrow.parquet as pq
ROOT = Path(__file__).resolve().parent
HIST_CHUNK = 1 << 22 # bases per counting pass
PROBS = ["pred_prob_positive_strand_cds", "pred_prob_negative_strand_cds"]
DISPLAY_COLUMNS = ["assembly_accession", "record_name", "organism_name", "division",
"segment_start_bp", "segment_end_bp", "segment_index", "segment_count"]
def normalize(value):
return str(value or "").strip().upper()
def unversioned(value):
return re.sub(r"\.\d+$", "", value)
def segment_metadata(record):
"""Older outputs store a whole aligned contig without segment columns."""
record = dict(record)
if "segment_start_bp" not in record and "segment_end_bp" not in record:
record.update(segment_start_bp=0, segment_end_bp=record["aligned_bp_length"],
segment_bp_length=record["aligned_bp_length"], segment_index=0, segment_count=1)
if not 0 <= record["segment_start_bp"] < record["segment_end_bp"]:
raise ValueError("Invalid segment coordinates in annotation metadata.")
return record
class Catalog:
def __init__(self, directory=ROOT / "data"):
directory = Path(directory)
self.path = directory / "sample.parquet"
self.manifest = json.loads((directory / "manifest.json").read_text())
pf = pq.ParquetFile(self.path)
columns = [c for c in pf.schema_arrow.names if c not in PROBS and c != "sequence"]
self.records = pf.read(columns=columns).to_pylist()
self.locations = [(g, r) for g in range(pf.num_row_groups)
for r in range(pf.metadata.row_group(g).num_rows)]
self.exact, self.base = defaultdict(set), defaultdict(set)
for i, record in enumerate(self.records):
for field in ("assembly_accession", "record_name", "source_key"):
key = normalize(record[field])
if key:
self.exact[key].add(i)
if field != "source_key":
self.base[unversioned(key)].add(i)
def lookup(self, accession):
key = normalize(accession)
if not key:
return []
# A versioned query never silently falls back to another version.
matches = self.exact.get(key, set()) if re.search(r"\.\d+$", key) or "|" in key else self.base.get(key, set())
return sorted(matches, key=lambda i: (self.records[i]["assembly_accession"],
self.records[i]["record_name"],
self.records[i]["segment_index"]))
def table(self, ids):
return pd.DataFrame([self.records[i] for i in ids], columns=DISPLAY_COLUMNS)
def find(self, accession, limit=200):
ids = self.lookup(accession)
return ids[:limit], len(ids)
def browse_ids(self, limit=200):
return list(range(min(limit, len(self.records))))
def fetch(self, index):
import time
started = time.perf_counter()
table = self.segment_table(index)
return table, {"cache_hit": False, "bytes_read": 0, "range_reads": 0,
"seconds": time.perf_counter() - started, "local": True}
def segment_table(self, index):
index = int(index)
if not 0 <= index < len(self.records):
raise ValueError("Choose a loaded segment.")
group, row = self.locations[index]
return pq.ParquetFile(self.path).read_row_group(group).slice(row, 1)
def window(self, index, start=None, end=None, max_points=1200, table=None,
mode="Probabilities", threshold=0.5, max_transitions=20000, hist_rows=40):
if mode not in ("Probabilities", "Binary labels"):
raise ValueError("Choose Probabilities or Binary labels.")
threshold = float(threshold)
if not np.isfinite(threshold) or not 0 <= threshold <= 1:
raise ValueError("Threshold must be between 0 and 1.")
record = self.records[int(index)]
lo, hi = record["segment_start_bp"], record["segment_end_bp"]
start = lo if start is None else int(start)
end = hi if end is None else int(end)
if not lo <= start < end <= hi:
raise ValueError(f"Enter a range within [{lo:,}, {hi:,}) with start < end.")
if table is None:
table = self.segment_table(index)
width = end - start
if mode == "Binary labels":
label = np.zeros(width, dtype=bool)
for column in PROBS:
values = table.column(column)[0].values
if len(values) != hi - lo:
raise ValueError("Probability length does not match segment coordinates.")
values = values.slice(start - lo, width).to_numpy(zero_copy_only=False).astype(np.float32)
if not np.isfinite(values).all():
raise ValueError("Cannot assign binary labels to missing or non-finite probabilities.")
# OR the per-strand decisions: max(P_pos, P_neg) > threshold.
# Strict > keeps exact ties as background, as binary argmax does.
label |= values > threshold
transitions = np.flatnonzero(label[1:] != label[:-1]) + 1
# Preserve every transition when manageable. Only dense regions need an overview.
step = 1 if len(transitions) <= max_transitions else max(1, (width + max_points - 1) // max_points)
offsets = np.r_[0, transitions] if step == 1 else np.arange(0, width, step)
values = label[offsets] if step == 1 else np.logical_or.reduceat(label, offsets)
# The last point closes the final half-open interval; it adds no base.
return pd.DataFrame({"Position (bp)": np.r_[start + offsets, end],
"Predicted CDS": np.r_[values, values[-1]].astype(np.uint8),
"Strand": "CDS (either strand)"}), step
step = max(1, (width + max_points - 1) // max_points)
positions = np.arange(start, end, step)
def strand_values(column):
values = table.column(column)[0].values
if len(values) != hi - lo:
raise ValueError("Probability length does not match segment coordinates.")
return values.slice(start - lo, width).to_numpy(zero_copy_only=False).astype(np.float32)
if step == 1:
frames = [pd.DataFrame({"Position (bp)": positions, "P(CDS)": strand_values(column), "Strand": strand})
for column, strand in zip(PROBS, ["+ strand", "− strand"])]
return pd.concat(frames, ignore_index=True), step
# Averaging a bin destroys what matters here. These probabilities are
# bimodal — a base is confidently coding or confidently not — so the mean
# of a bin that is 20% exons lands near 0.2, a value almost no base holds,
# and the peaks the binary view fires on vanish. Keep the distribution
# instead: one histogram per column over max(P_pos, P_neg), the same value
# the threshold and the binary labels are computed from.
best = np.maximum(strand_values(PROBS[0]), strand_values(PROBS[1]))
offsets = np.arange(0, width, step)
counts = np.minimum(step, width - offsets)
# Counted in chunks: a whole-chromosome window is 100M+ bases, and an
# index array over all of them at once costs more memory than the
# probabilities themselves.
bases = np.zeros(len(offsets) * hist_rows, dtype=np.int64)
for begin in range(0, width, HIST_CHUNK):
piece = best[begin:begin + HIST_CHUNK]
columns = np.minimum(np.arange(begin, begin + len(piece)) // step, len(offsets) - 1)
rows = np.minimum((piece * hist_rows).astype(np.int32), hist_rows - 1)
bases += np.bincount(columns * hist_rows + rows, minlength=bases.size)
means = np.add.reduceat(best, offsets) / counts
centres = (np.arange(hist_rows) + 0.5) / hist_rows
return pd.DataFrame({"Position (bp)": np.repeat(positions, hist_rows),
"P(CDS)": np.tile(centres, len(offsets)),
"Bases": bases,
"Mean P": np.repeat(means, hist_rows)}), step
|