Show the distribution of probabilities per column, not their mean
#11
by lvwerra HF Staff - opened
- app.py +37 -1
- catalog.py +35 -8
- tests/test_catalog.py +26 -6
app.py
CHANGED
|
@@ -11,10 +11,13 @@ import time
|
|
| 11 |
os.environ.setdefault("GRADIO_TEMP_DIR", str(Path(tempfile.gettempdir()) / f"genbank-explorer-{os.getuid()}"))
|
| 12 |
|
| 13 |
import gradio as gr
|
|
|
|
| 14 |
import pyarrow.parquet as pq
|
| 15 |
import plotly.graph_objects as go
|
| 16 |
|
| 17 |
from taxonomy import build_taxonomy_tab
|
|
|
|
|
|
|
| 18 |
from style import APP_CSS, atlas_theme, section_header
|
| 19 |
from catalog import Catalog
|
| 20 |
from remote_catalog import RemoteCatalog, RemoteReadError
|
|
@@ -47,6 +50,34 @@ def build_app(catalog=None):
|
|
| 47 |
binary = mode == "Binary labels"
|
| 48 |
column = "Predicted CDS" if binary else "P(CDS)"
|
| 49 |
figure = go.Figure()
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 50 |
tracks = [("CDS (either strand)", "#287557")] if binary else [("+ strand", "#287557"), ("− strand", "#ae754b")]
|
| 51 |
for strand, color in tracks:
|
| 52 |
rows = frame[frame["Strand"] == strand]
|
|
@@ -93,7 +124,12 @@ def build_app(catalog=None):
|
|
| 93 |
f"**Binned overview:** each interval spans up to {step:,} bases and is 1 if **any** base exceeds the threshold. "
|
| 94 |
"This does not mean every base in that interval is CDS. Narrow the region for exact labels."))
|
| 95 |
else:
|
| 96 |
-
resolution = "Each point is one base." if step == 1 else
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 97 |
return ("Coordinates are **0-based, end-exclusive**. "
|
| 98 |
+ resolution
|
| 99 |
+ " Download the segment for the original per-base probabilities.\n\n"
|
|
|
|
| 11 |
os.environ.setdefault("GRADIO_TEMP_DIR", str(Path(tempfile.gettempdir()) / f"genbank-explorer-{os.getuid()}"))
|
| 12 |
|
| 13 |
import gradio as gr
|
| 14 |
+
import numpy as np
|
| 15 |
import pyarrow.parquet as pq
|
| 16 |
import plotly.graph_objects as go
|
| 17 |
|
| 18 |
from taxonomy import build_taxonomy_tab
|
| 19 |
+
|
| 20 |
+
HIST_ROWS = 40
|
| 21 |
from style import APP_CSS, atlas_theme, section_header
|
| 22 |
from catalog import Catalog
|
| 23 |
from remote_catalog import RemoteCatalog, RemoteReadError
|
|
|
|
| 50 |
binary = mode == "Binary labels"
|
| 51 |
column = "Predicted CDS" if binary else "P(CDS)"
|
| 52 |
figure = go.Figure()
|
| 53 |
+
if "Bases" in frame.columns:
|
| 54 |
+
# A column of the heatmap is the distribution of per-base
|
| 55 |
+
# probabilities there, so a region that is part exon and part intron
|
| 56 |
+
# shows both bands instead of an average lying between them.
|
| 57 |
+
positions = frame["Position (bp)"].to_numpy()[::HIST_ROWS]
|
| 58 |
+
centres = frame["P(CDS)"].to_numpy()[:HIST_ROWS]
|
| 59 |
+
bases = frame["Bases"].to_numpy().reshape(len(positions), HIST_ROWS).T
|
| 60 |
+
figure.add_trace(go.Heatmap(
|
| 61 |
+
x=positions, y=centres, z=np.log10(bases + 1), customdata=bases,
|
| 62 |
+
colorscale=[[0, "#fffefa"], [0.25, "#cfe0cd"], [0.6, "#63a07f"], [1, "#173c30"]],
|
| 63 |
+
colorbar=dict(title=dict(text="bases", side="right"), thickness=12,
|
| 64 |
+
tickvals=[0, 1, 2, 3, 4], ticktext=["1", "10", "100", "1k", "10k"]),
|
| 65 |
+
hovertemplate="%{customdata:,} bases near P=%{y:.2f}<br>from %{x:,}<extra></extra>"))
|
| 66 |
+
figure.add_trace(go.Scatter(
|
| 67 |
+
x=positions, y=frame["Mean P"].to_numpy()[::HIST_ROWS], mode="lines", name="mean",
|
| 68 |
+
line=dict(color="#c98b5b", width=1), hovertemplate="mean %{y:.3f}<extra></extra>"))
|
| 69 |
+
figure.add_hline(y=threshold, line_dash="dot", line_color="#8b9b7b",
|
| 70 |
+
annotation_text=f"Threshold {threshold:g}")
|
| 71 |
+
figure.update_layout(title="CDS probability distribution", xaxis_title="Position (bp; 0-based)",
|
| 72 |
+
yaxis_title="P(CDS)", height=380, margin=dict(l=60, r=25, t=75, b=50),
|
| 73 |
+
template="plotly_white", paper_bgcolor="#fffefa", plot_bgcolor="#fffefa",
|
| 74 |
+
font=dict(family="Arial, Helvetica, sans-serif", color="#315641", size=12),
|
| 75 |
+
title_font=dict(family="Georgia, Times New Roman, serif", size=22, color="#173c30"),
|
| 76 |
+
hoverlabel=dict(bgcolor="#173c30", font_color="#ffffff", bordercolor="#173c30"),
|
| 77 |
+
showlegend=False)
|
| 78 |
+
figure.update_xaxes(gridcolor="#e9edde", zerolinecolor="#dbe3d3", linecolor="#dbe3d3")
|
| 79 |
+
figure.update_yaxes(range=[0, 1], gridcolor="#e9edde", zerolinecolor="#dbe3d3", linecolor="#dbe3d3")
|
| 80 |
+
return figure
|
| 81 |
tracks = [("CDS (either strand)", "#287557")] if binary else [("+ strand", "#287557"), ("− strand", "#ae754b")]
|
| 82 |
for strand, color in tracks:
|
| 83 |
rows = frame[frame["Strand"] == strand]
|
|
|
|
| 124 |
f"**Binned overview:** each interval spans up to {step:,} bases and is 1 if **any** base exceeds the threshold. "
|
| 125 |
"This does not mean every base in that interval is CDS. Narrow the region for exact labels."))
|
| 126 |
else:
|
| 127 |
+
resolution = ("Each point is one base, per strand." if step == 1 else
|
| 128 |
+
f"Each column covers up to **{step:,} bases** and shows how their probabilities are "
|
| 129 |
+
"distributed, over max(P_positive, P_negative) — the value the threshold tests. "
|
| 130 |
+
"Colour is how many bases fall in that band, so a part-coding region shows a high "
|
| 131 |
+
"band and a low one rather than an average between them. The thin line is the mean. "
|
| 132 |
+
"Narrow the region for exact per-strand values.")
|
| 133 |
return ("Coordinates are **0-based, end-exclusive**. "
|
| 134 |
+ resolution
|
| 135 |
+ " Download the segment for the original per-base probabilities.\n\n"
|
catalog.py
CHANGED
|
@@ -9,6 +9,7 @@ import pandas as pd
|
|
| 9 |
import pyarrow.parquet as pq
|
| 10 |
|
| 11 |
ROOT = Path(__file__).resolve().parent
|
|
|
|
| 12 |
PROBS = ["pred_prob_positive_strand_cds", "pred_prob_negative_strand_cds"]
|
| 13 |
DISPLAY_COLUMNS = ["assembly_accession", "record_name", "organism_name", "division",
|
| 14 |
"segment_start_bp", "segment_end_bp", "segment_index", "segment_count"]
|
|
@@ -87,7 +88,7 @@ class Catalog:
|
|
| 87 |
return pq.ParquetFile(self.path).read_row_group(group).slice(row, 1)
|
| 88 |
|
| 89 |
def window(self, index, start=None, end=None, max_points=1200, table=None,
|
| 90 |
-
mode="Probabilities", threshold=0.5, max_transitions=20000):
|
| 91 |
if mode not in ("Probabilities", "Binary labels"):
|
| 92 |
raise ValueError("Choose Probabilities or Binary labels.")
|
| 93 |
threshold = float(threshold)
|
|
@@ -125,13 +126,39 @@ class Catalog:
|
|
| 125 |
"Strand": "CDS (either strand)"}), step
|
| 126 |
step = max(1, (width + max_points - 1) // max_points)
|
| 127 |
positions = np.arange(start, end, step)
|
| 128 |
-
|
| 129 |
-
|
| 130 |
values = table.column(column)[0].values
|
| 131 |
if len(values) != hi - lo:
|
| 132 |
raise ValueError("Probability length does not match segment coordinates.")
|
| 133 |
-
|
| 134 |
-
|
| 135 |
-
|
| 136 |
-
frames
|
| 137 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 9 |
import pyarrow.parquet as pq
|
| 10 |
|
| 11 |
ROOT = Path(__file__).resolve().parent
|
| 12 |
+
HIST_CHUNK = 1 << 22 # bases per counting pass
|
| 13 |
PROBS = ["pred_prob_positive_strand_cds", "pred_prob_negative_strand_cds"]
|
| 14 |
DISPLAY_COLUMNS = ["assembly_accession", "record_name", "organism_name", "division",
|
| 15 |
"segment_start_bp", "segment_end_bp", "segment_index", "segment_count"]
|
|
|
|
| 88 |
return pq.ParquetFile(self.path).read_row_group(group).slice(row, 1)
|
| 89 |
|
| 90 |
def window(self, index, start=None, end=None, max_points=1200, table=None,
|
| 91 |
+
mode="Probabilities", threshold=0.5, max_transitions=20000, hist_rows=40):
|
| 92 |
if mode not in ("Probabilities", "Binary labels"):
|
| 93 |
raise ValueError("Choose Probabilities or Binary labels.")
|
| 94 |
threshold = float(threshold)
|
|
|
|
| 126 |
"Strand": "CDS (either strand)"}), step
|
| 127 |
step = max(1, (width + max_points - 1) // max_points)
|
| 128 |
positions = np.arange(start, end, step)
|
| 129 |
+
|
| 130 |
+
def strand_values(column):
|
| 131 |
values = table.column(column)[0].values
|
| 132 |
if len(values) != hi - lo:
|
| 133 |
raise ValueError("Probability length does not match segment coordinates.")
|
| 134 |
+
return values.slice(start - lo, width).to_numpy(zero_copy_only=False).astype(np.float32)
|
| 135 |
+
|
| 136 |
+
if step == 1:
|
| 137 |
+
frames = [pd.DataFrame({"Position (bp)": positions, "P(CDS)": strand_values(column), "Strand": strand})
|
| 138 |
+
for column, strand in zip(PROBS, ["+ strand", "− strand"])]
|
| 139 |
+
return pd.concat(frames, ignore_index=True), step
|
| 140 |
+
|
| 141 |
+
# Averaging a bin destroys what matters here. These probabilities are
|
| 142 |
+
# bimodal — a base is confidently coding or confidently not — so the mean
|
| 143 |
+
# of a bin that is 20% exons lands near 0.2, a value almost no base holds,
|
| 144 |
+
# and the peaks the binary view fires on vanish. Keep the distribution
|
| 145 |
+
# instead: one histogram per column over max(P_pos, P_neg), the same value
|
| 146 |
+
# the threshold and the binary labels are computed from.
|
| 147 |
+
best = np.maximum(strand_values(PROBS[0]), strand_values(PROBS[1]))
|
| 148 |
+
offsets = np.arange(0, width, step)
|
| 149 |
+
counts = np.minimum(step, width - offsets)
|
| 150 |
+
# Counted in chunks: a whole-chromosome window is 100M+ bases, and an
|
| 151 |
+
# index array over all of them at once costs more memory than the
|
| 152 |
+
# probabilities themselves.
|
| 153 |
+
bases = np.zeros(len(offsets) * hist_rows, dtype=np.int64)
|
| 154 |
+
for begin in range(0, width, HIST_CHUNK):
|
| 155 |
+
piece = best[begin:begin + HIST_CHUNK]
|
| 156 |
+
columns = np.minimum(np.arange(begin, begin + len(piece)) // step, len(offsets) - 1)
|
| 157 |
+
rows = np.minimum((piece * hist_rows).astype(np.int32), hist_rows - 1)
|
| 158 |
+
bases += np.bincount(columns * hist_rows + rows, minlength=bases.size)
|
| 159 |
+
means = np.add.reduceat(best, offsets) / counts
|
| 160 |
+
centres = (np.arange(hist_rows) + 0.5) / hist_rows
|
| 161 |
+
return pd.DataFrame({"Position (bp)": np.repeat(positions, hist_rows),
|
| 162 |
+
"P(CDS)": np.tile(centres, len(offsets)),
|
| 163 |
+
"Bases": bases,
|
| 164 |
+
"Mean P": np.repeat(means, hist_rows)}), step
|
tests/test_catalog.py
CHANGED
|
@@ -69,15 +69,35 @@ class CoordinateTests(unittest.TestCase):
|
|
| 69 |
self.assertEqual(self.catalog.lookup('TEST000001'), [0, 1])
|
| 70 |
self.assertEqual(self.catalog.lookup('GCA_000001.3'), [])
|
| 71 |
|
| 72 |
-
def
|
| 73 |
-
frame, step = self.catalog.window(0, 100, 105, max_points=2)
|
| 74 |
-
self.assertEqual(step, 3)
|
| 75 |
-
positive = frame[frame.Strand == '+ strand']
|
| 76 |
-
self.assertEqual(positive['Position (bp)'].tolist(), [100, 103])
|
| 77 |
-
np.testing.assert_allclose(positive['P(CDS)'], [0.2, 0.7], atol=1e-6)
|
| 78 |
frame, step = self.catalog.window(0, 103, 105)
|
| 79 |
self.assertEqual(step, 1)
|
|
|
|
| 80 |
np.testing.assert_allclose(frame[frame.Strand == '+ strand']['P(CDS)'], [0.6, 0.8])
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 81 |
|
| 82 |
def test_invalid_ranges(self):
|
| 83 |
for start, end in [(99, 103), (100, 106), (103, 103), (104, 102)]:
|
|
|
|
| 69 |
self.assertEqual(self.catalog.lookup('TEST000001'), [0, 1])
|
| 70 |
self.assertEqual(self.catalog.lookup('GCA_000001.3'), [])
|
| 71 |
|
| 72 |
+
def test_exact_window_keeps_both_strands_per_base(self):
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 73 |
frame, step = self.catalog.window(0, 103, 105)
|
| 74 |
self.assertEqual(step, 1)
|
| 75 |
+
self.assertEqual(frame['Position (bp)'].tolist(), [103, 104, 103, 104])
|
| 76 |
np.testing.assert_allclose(frame[frame.Strand == '+ strand']['P(CDS)'], [0.6, 0.8])
|
| 77 |
+
np.testing.assert_allclose(frame[frame.Strand == '− strand']['P(CDS)'], [0.4, 0.2])
|
| 78 |
+
|
| 79 |
+
def test_binned_window_keeps_the_distribution_not_the_mean(self):
|
| 80 |
+
"""A mean hides a part-coding bin; the histogram has to keep both bands."""
|
| 81 |
+
rows = 4
|
| 82 |
+
frame, step = self.catalog.window(0, 100, 105, max_points=2, hist_rows=rows)
|
| 83 |
+
self.assertEqual(step, 3)
|
| 84 |
+
self.assertNotIn('Strand', frame.columns)
|
| 85 |
+
self.assertEqual(frame['Position (bp)'].tolist()[::rows], [100, 103])
|
| 86 |
+
bases = frame['Bases'].to_numpy().reshape(-1, rows)
|
| 87 |
+
# Every base is counted exactly once, and only once.
|
| 88 |
+
self.assertEqual(bases.sum(), 5)
|
| 89 |
+
self.assertEqual(bases.sum(axis=1).tolist(), [3, 2])
|
| 90 |
+
# max(P_pos, P_neg) over 100..104 is [1.0, 0.8, 0.6, 0.6, 0.8]; with four
|
| 91 |
+
# bands those land in the top band, top, third, third, top.
|
| 92 |
+
self.assertEqual(bases.tolist(), [[0, 0, 1, 2], [0, 0, 1, 1]])
|
| 93 |
+
np.testing.assert_allclose(frame['Mean P'].to_numpy()[::rows], [0.8, 0.7], atol=1e-6)
|
| 94 |
+
|
| 95 |
+
def test_binned_window_conserves_every_base(self):
|
| 96 |
+
frame, step = self.catalog.window(0, max_points=2, hist_rows=8)
|
| 97 |
+
width = self.catalog.records[0]['segment_end_bp'] - self.catalog.records[0]['segment_start_bp']
|
| 98 |
+
self.assertGreater(step, 1)
|
| 99 |
+
self.assertEqual(frame['Bases'].sum(), width)
|
| 100 |
+
self.assertTrue((frame['P(CDS)'] > 0).all() and (frame['P(CDS)'] < 1).all())
|
| 101 |
|
| 102 |
def test_invalid_ranges(self):
|
| 103 |
for start, end in [(99, 103), (100, 106), (103, 103), (104, 102)]:
|