Show the distribution of probabilities per column, not their mean

#11
by lvwerra HF Staff - opened
Files changed (3) hide show
  1. app.py +37 -1
  2. catalog.py +35 -8
  3. 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 f"Each point is the mean of up to **{step:,} bases**; short peaks can be smoothed."
 
 
 
 
 
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
- frames = []
129
- for column, strand in zip(PROBS, ["+ strand", "− strand"]):
130
  values = table.column(column)[0].values
131
  if len(values) != hi - lo:
132
  raise ValueError("Probability length does not match segment coordinates.")
133
- values = values.slice(start - lo, width).to_numpy(zero_copy_only=False).astype(np.float32)
134
- counts = np.minimum(step, width - np.arange(0, width, step))
135
- means = np.add.reduceat(values, np.arange(0, width, step)) / counts
136
- frames.append(pd.DataFrame({"Position (bp)": positions, "P(CDS)": means, "Strand": strand}))
137
- return pd.concat(frames, ignore_index=True), step
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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 test_absolute_coordinates_and_partial_last_bin(self):
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)]: