Combine strand decisions into one binary CDS track with background ties
Browse files- README.md +3 -1
- app.py +8 -6
- catalog.py +12 -14
- 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
|
|
|
|
|
|
|
| 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 |
-
|
|
|
|
| 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
|
| 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 =
|
| 79 |
-
|
| 80 |
-
|
|
|
|
| 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
|
| 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 |
-
|
| 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 |
-
|
| 115 |
-
|
| 116 |
-
|
|
|
|
| 117 |
# Preserve every transition when manageable. Only dense regions need an overview.
|
| 118 |
-
step = 1 if
|
| 119 |
-
|
| 120 |
-
|
| 121 |
-
|
| 122 |
-
|
| 123 |
-
|
| 124 |
-
|
| 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.
|
| 89 |
self.assertEqual(step, 1)
|
| 90 |
-
|
| 91 |
-
|
| 92 |
-
self.assertEqual(
|
| 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.
|
| 100 |
max_points=2, max_transitions=0)
|
| 101 |
self.assertEqual(step, 3)
|
| 102 |
-
|
| 103 |
-
|
| 104 |
-
self.assertEqual(
|
| 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 |
-
|
| 112 |
-
|
| 113 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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)
|