"""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