cgeorgiaw HF Staff commited on
Commit
ceaecab
·
verified ·
1 Parent(s): a312d10

Combine strand decisions into one binary CDS track with background ties

Browse files
Files changed (4) hide show
  1. README.md +3 -1
  2. app.py +8 -6
  3. catalog.py +12 -14
  4. tests/test_catalog.py +15 -15
README.md CHANGED
@@ -76,7 +76,9 @@ The benchmark fetches one row group per source, validates its selected probabili
76
 
77
  ## Data and coordinates
78
 
79
- The viewer offers **Probabilities** and **Binary labels** modes, with an adjustable threshold defaulting to **0.5**. Binary labels are computed per base and strand as `P(CDS) >= threshold`. Stepped tracks preserve exact labels using their change positions when a strand has at most 20,000 transitions. For denser regions, a clearly labeled overview shows 1 when any base in a bin meets the threshold; zoom in by changing the start/end fields for exact labels. Thresholding happens before binning, so it never thresholds an averaged probability. Mode and threshold changes update the current region without resetting its coordinates. Downloads continue to contain original probabilities.
 
 
80
 
81
  The bucket contains model predictions: positive- and negative-strand per-base P(CDS), metadata, and sometimes sequence. It does not provide curated gene names or gene models. Coordinates are 0-based, end-exclusive. The initial viewport shows the entire selected segment; update the start/end fields to inspect a smaller region. Large windows are averaged into at most 1,200 bins per strand. Downloads retain original probabilities and source columns without downsampling.
82
 
 
76
 
77
  ## Data and coordinates
78
 
79
+ The viewer offers **Probabilities** and **Binary labels** modes, with an adjustable threshold defaulting to **0.5**. Binary mode shows one combined CDS/background track: `max(P_positive, P_negative) > threshold`. Either strand can make a base CDS; exact ties are background. This combines the per-strand binary decisions with OR. Validation uses per-strand `argmax` between background and CDS and scores those strand labels; the combined viewer track is not a separate reproduction of those validation metrics. Stored float16 probabilities and inference-window averaging also prevent exact reconstruction of original evaluation logits.
80
+
81
+ The stepped track preserves exact labels using their change positions when the combined track has at most 20,000 transitions. For denser regions, a clearly labeled overview shows 1 when any base in a bin exceeds the threshold; zoom in by changing the start/end fields for exact labels. Thresholding happens before binning, so it never thresholds an averaged probability. Mode and threshold changes update the current region without resetting its coordinates. Probability mode retains the two strand tracks. Downloads continue to contain original probabilities.
82
 
83
  The bucket contains model predictions: positive- and negative-strand per-base P(CDS), metadata, and sometimes sequence. It does not provide curated gene names or gene models. Coordinates are 0-based, end-exclusive. The initial viewport shows the entire selected segment; update the start/end fields to inspect a smaller region. Large windows are averaged into at most 1,200 bins per strand. Downloads retain original probabilities and source columns without downsampling.
84
 
app.py CHANGED
@@ -43,7 +43,8 @@ def build_app(catalog=None):
43
  binary = mode == "Binary labels"
44
  column = "Predicted CDS" if binary else "P(CDS)"
45
  figure = go.Figure()
46
- for strand, color in [("+ strand", "#2563eb"), ("− strand", "#ea580c")]:
 
47
  rows = frame[frame["Strand"] == strand]
48
  figure.add_trace(go.Scatter(x=rows["Position (bp)"].tolist(), y=rows[column].tolist(),
49
  name=strand, mode="lines",
@@ -51,7 +52,7 @@ def build_app(catalog=None):
51
  if not binary:
52
  figure.add_hline(y=threshold, line_dash="dot", line_color="#64748b",
53
  annotation_text=f"Threshold {threshold:g}")
54
- figure.update_layout(title="Predicted CDS labels by strand" if binary else "CDS probability by strand",
55
  xaxis_title="Position (bp; 0-based)", yaxis_title=column,
56
  height=360, margin=dict(l=55, r=20, t=50, b=45), template="plotly_white",
57
  hovermode="x unified", legend=dict(orientation="h"))
@@ -75,9 +76,10 @@ def build_app(catalog=None):
75
  def plot_note(step, stats, mode="Probabilities", threshold=0.5):
76
  origin = "local sample" if stats.get("local") else "cache" if stats["cache_hit"] else "bucket"
77
  if mode == "Binary labels":
78
- resolution = (f"**1 = P(CDS) ≥ {threshold:g}; 0 = below threshold**, evaluated separately on each strand. "
79
- + ("The stepped tracks preserve every base label in this region." if step == 1 else
80
- f"**Binned overview:** each interval spans up to {step:,} bases and is 1 if **any** base meets the threshold. "
 
81
  "This does not mean every base in that interval is CDS. Narrow the region for exact labels."))
82
  else:
83
  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."
@@ -141,7 +143,7 @@ def build_app(catalog=None):
141
  view = gr.Button("Update region")
142
  with gr.Row():
143
  mode = gr.Radio(["Probabilities", "Binary labels"], value="Probabilities", label="Viewer mode")
144
- threshold = gr.Slider(0, 1, value=0.5, step=0.01, label="CDS threshold (P ≥ threshold)")
145
  plot = gr.Plot(label="CDS tracks")
146
  note = gr.Markdown("Choose a segment to see its probability tracks.")
147
  with gr.Accordion("Segment metadata and provenance", open=False):
 
43
  binary = mode == "Binary labels"
44
  column = "Predicted CDS" if binary else "P(CDS)"
45
  figure = go.Figure()
46
+ tracks = [("CDS (either strand)", "#2563eb")] if binary else [("+ strand", "#2563eb"), ("− strand", "#ea580c")]
47
+ for strand, color in tracks:
48
  rows = frame[frame["Strand"] == strand]
49
  figure.add_trace(go.Scatter(x=rows["Position (bp)"].tolist(), y=rows[column].tolist(),
50
  name=strand, mode="lines",
 
52
  if not binary:
53
  figure.add_hline(y=threshold, line_dash="dot", line_color="#64748b",
54
  annotation_text=f"Threshold {threshold:g}")
55
+ figure.update_layout(title="Predicted CDS — either strand" if binary else "CDS probability by strand",
56
  xaxis_title="Position (bp; 0-based)", yaxis_title=column,
57
  height=360, margin=dict(l=55, r=20, t=50, b=45), template="plotly_white",
58
  hovermode="x unified", legend=dict(orientation="h"))
 
76
  def plot_note(step, stats, mode="Probabilities", threshold=0.5):
77
  origin = "local sample" if stats.get("local") else "cache" if stats["cache_hit"] else "bucket"
78
  if mode == "Binary labels":
79
+ resolution = (f"**1 = max(P_positive, P_negative) > {threshold:g}; 0 = otherwise.** "
80
+ "This is one CDS/background label per base, combining both strands. "
81
+ + ("The stepped track preserves every base label in this region." if step == 1 else
82
+ f"**Binned overview:** each interval spans up to {step:,} bases and is 1 if **any** base exceeds the threshold. "
83
  "This does not mean every base in that interval is CDS. Narrow the region for exact labels."))
84
  else:
85
  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."
 
143
  view = gr.Button("Update region")
144
  with gr.Row():
145
  mode = gr.Radio(["Probabilities", "Binary labels"], value="Probabilities", label="Viewer mode")
146
+ threshold = gr.Slider(0, 1, value=0.5, step=0.01, label="CDS threshold (max strand P > threshold)")
147
  plot = gr.Plot(label="CDS tracks")
148
  note = gr.Markdown("Choose a segment to see its probability tracks.")
149
  with gr.Accordion("Segment metadata and provenance", open=False):
catalog.py CHANGED
@@ -103,7 +103,7 @@ class Catalog:
103
  table = self.segment_table(index)
104
  width = end - start
105
  if mode == "Binary labels":
106
- labels, changes = [], []
107
  for column in PROBS:
108
  values = table.column(column)[0].values
109
  if len(values) != hi - lo:
@@ -111,20 +111,18 @@ class Catalog:
111
  values = values.slice(start - lo, width).to_numpy(zero_copy_only=False).astype(np.float32)
112
  if not np.isfinite(values).all():
113
  raise ValueError("Cannot assign binary labels to missing or non-finite probabilities.")
114
- label = values >= threshold
115
- labels.append(label)
116
- changes.append(np.flatnonzero(label[1:] != label[:-1]) + 1)
 
117
  # Preserve every transition when manageable. Only dense regions need an overview.
118
- step = 1 if max(map(len, changes)) <= max_transitions else max(1, (width + max_points - 1) // max_points)
119
- frames = []
120
- for label, transitions, strand in zip(labels, changes, ["+ strand", "− strand"]):
121
- offsets = np.r_[0, transitions] if step == 1 else np.arange(0, width, step)
122
- values = label[offsets] if step == 1 else np.logical_or.reduceat(label, offsets)
123
- # The last point closes the final half-open interval; it adds no base.
124
- frames.append(pd.DataFrame({"Position (bp)": np.r_[start + offsets, end],
125
- "Predicted CDS": np.r_[values, values[-1]].astype(np.uint8),
126
- "Strand": strand}))
127
- return pd.concat(frames, ignore_index=True), step
128
  step = max(1, (width + max_points - 1) // max_points)
129
  positions = np.arange(start, end, step)
130
  frames = []
 
103
  table = self.segment_table(index)
104
  width = end - start
105
  if mode == "Binary labels":
106
+ label = np.zeros(width, dtype=bool)
107
  for column in PROBS:
108
  values = table.column(column)[0].values
109
  if len(values) != hi - lo:
 
111
  values = values.slice(start - lo, width).to_numpy(zero_copy_only=False).astype(np.float32)
112
  if not np.isfinite(values).all():
113
  raise ValueError("Cannot assign binary labels to missing or non-finite probabilities.")
114
+ # OR the per-strand decisions: max(P_pos, P_neg) > threshold.
115
+ # Strict > keeps exact ties as background, as binary argmax does.
116
+ label |= values > threshold
117
+ transitions = np.flatnonzero(label[1:] != label[:-1]) + 1
118
  # Preserve every transition when manageable. Only dense regions need an overview.
119
+ step = 1 if len(transitions) <= max_transitions else max(1, (width + max_points - 1) // max_points)
120
+ offsets = np.r_[0, transitions] if step == 1 else np.arange(0, width, step)
121
+ values = label[offsets] if step == 1 else np.logical_or.reduceat(label, offsets)
122
+ # The last point closes the final half-open interval; it adds no base.
123
+ return pd.DataFrame({"Position (bp)": np.r_[start + offsets, end],
124
+ "Predicted CDS": np.r_[values, values[-1]].astype(np.uint8),
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 = []
tests/test_catalog.py CHANGED
@@ -85,32 +85,32 @@ class CoordinateTests(unittest.TestCase):
85
  self.catalog.window(0, start, end)
86
 
87
  def test_binary_labels_preserve_exact_transitions_and_coordinates(self):
88
- frame, step = self.catalog.window(0, 100, 105, mode='Binary labels', threshold=0.4)
89
  self.assertEqual(step, 1)
90
- positive = frame[frame.Strand == '+ strand']
91
- negative = frame[frame.Strand == '− strand']
92
- self.assertEqual(positive['Position (bp)'].tolist(), [100, 102, 105])
93
- self.assertEqual(positive['Predicted CDS'].tolist(), [0, 1, 1])
94
- self.assertEqual(negative['Position (bp)'].tolist(), [100, 104, 105])
95
- self.assertEqual(negative['Predicted CDS'].tolist(), [1, 0, 0])
96
  self.assertTrue(set(frame['Predicted CDS']).issubset({0, 1}))
97
 
98
  def test_binary_overview_thresholds_before_binning(self):
99
- frame, step = self.catalog.window(0, 100, 105, mode='Binary labels', threshold=0.4,
100
  max_points=2, max_transitions=0)
101
  self.assertEqual(step, 3)
102
- positive = frame[frame.Strand == '+ strand']
103
- # First bin's mean is 0.2, but one base reaches 0.4 and must be shown.
104
- self.assertEqual(positive['Position (bp)'].tolist(), [100, 103, 105])
105
- self.assertEqual(positive['Predicted CDS'].tolist(), [1, 1, 1])
106
 
107
  def test_threshold_limits_and_equality(self):
108
  frame, _ = self.catalog.window(0, mode='Binary labels', threshold=0)
109
  self.assertEqual(set(frame['Predicted CDS']), {1})
110
  frame, _ = self.catalog.window(0, mode='Binary labels', threshold=1)
111
- negative = frame[frame.Strand == '− strand']
112
- self.assertEqual(negative['Position (bp)'].tolist(), [100, 101, 105])
113
- self.assertEqual(negative['Predicted CDS'].tolist(), [1, 0, 0])
 
 
 
 
114
  for threshold in [-0.1, 1.1, float('nan')]:
115
  with self.assertRaises(ValueError):
116
  self.catalog.window(0, mode='Binary labels', threshold=threshold)
 
85
  self.catalog.window(0, start, end)
86
 
87
  def test_binary_labels_preserve_exact_transitions_and_coordinates(self):
88
+ frame, step = self.catalog.window(0, 100, 105, mode='Binary labels', threshold=0.75)
89
  self.assertEqual(step, 1)
90
+ self.assertEqual(frame.Strand.unique().tolist(), ['CDS (either strand)'])
91
+ self.assertEqual(frame['Position (bp)'].tolist(), [100, 102, 104, 105])
92
+ self.assertEqual(frame['Predicted CDS'].tolist(), [1, 0, 1, 1])
 
 
 
93
  self.assertTrue(set(frame['Predicted CDS']).issubset({0, 1}))
94
 
95
  def test_binary_overview_thresholds_before_binning(self):
96
+ frame, step = self.catalog.window(0, 100, 105, mode='Binary labels', threshold=0.85,
97
  max_points=2, max_transitions=0)
98
  self.assertEqual(step, 3)
99
+ # First bin's combined mean is 0.8, but one base exceeds 0.85.
100
+ self.assertEqual(frame['Position (bp)'].tolist(), [100, 103, 105])
101
+ self.assertEqual(frame['Predicted CDS'].tolist(), [1, 0, 0])
 
102
 
103
  def test_threshold_limits_and_equality(self):
104
  frame, _ = self.catalog.window(0, mode='Binary labels', threshold=0)
105
  self.assertEqual(set(frame['Predicted CDS']), {1})
106
  frame, _ = self.catalog.window(0, mode='Binary labels', threshold=1)
107
+ self.assertEqual(set(frame['Predicted CDS']), {0})
108
+ table = pa.Table.from_pylist([{
109
+ 'pred_prob_positive_strand_cds': [0.5, 0.1, 0.51, 0.49, 0.5],
110
+ 'pred_prob_negative_strand_cds': [0.1, 0.5, 0.1, 0.1, 0.5]}])
111
+ frame, _ = self.catalog.window(0, mode='Binary labels', threshold=0.5, table=table)
112
+ self.assertEqual(frame['Position (bp)'].tolist(), [100, 102, 103, 105])
113
+ self.assertEqual(frame['Predicted CDS'].tolist(), [0, 1, 0, 0])
114
  for threshold in [-0.1, 1.1, float('nan')]:
115
  with self.assertRaises(ValueError):
116
  self.catalog.window(0, mode='Binary labels', threshold=threshold)