lvwerra's picture
lvwerra HF Staff
Show the distribution of probabilities per column, not their mean (#11)
0d06695
Raw History Blame Contribute Delete
8.53 kB
"""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