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