File size: 19,130 Bytes
97d5f73
 
f0190da
386fb8f
97d5f73
1dadabb
97d5f73
f0190da
97d5f73
 
 
 
 
0d06695
97d5f73
a312d10
97d5f73
35d4af4
07f2e7f
97d5f73
f0190da
97d5f73
07f2e7f
19eb242
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
07f2e7f
 
 
 
97d5f73
 
f0190da
 
 
 
 
7034da5
 
97d5f73
07f2e7f
 
 
 
 
97d5f73
f0190da
 
 
97d5f73
07f2e7f
97d5f73
7034da5
97d5f73
f0190da
07f2e7f
 
97d5f73
a312d10
 
 
 
0d06695
 
 
 
 
 
 
 
 
 
 
 
 
 
07f2e7f
0d06695
 
 
 
 
 
 
 
 
07f2e7f
0d06695
 
 
648cd8a
ceaecab
a312d10
 
 
648cd8a
 
a312d10
648cd8a
a312d10
07f2e7f
a312d10
648cd8a
 
 
 
 
 
 
 
a312d10
 
 
 
07f2e7f
97d5f73
07f2e7f
97d5f73
 
26b65f6
f0190da
 
a312d10
 
 
07f2e7f
 
97d5f73
a312d10
f0190da
a312d10
07f2e7f
 
a312d10
07f2e7f
 
 
 
97d5f73
a312d10
97d5f73
07f2e7f
97d5f73
f0190da
a312d10
 
97d5f73
a312d10
 
97d5f73
 
 
 
f0190da
 
 
 
1dadabb
 
f0190da
 
 
1dadabb
 
 
07f2e7f
97d5f73
 
be172fd
7d9eeed
19eb242
be172fd
447bf4f
19eb242
07f2e7f
 
53c8150
07f2e7f
 
 
 
 
53c8150
386fb8f
53c8150
7034da5
 
 
386fb8f
 
 
 
 
 
 
07f2e7f
 
 
 
 
648cd8a
53c8150
07f2e7f
 
 
 
 
 
 
 
7f2efda
07f2e7f
 
 
 
53c8150
07f2e7f
53c8150
07f2e7f
 
 
 
53c8150
 
7034da5
53c8150
07f2e7f
372db56
 
 
 
 
07f2e7f
 
 
 
 
 
 
be172fd
 
 
 
 
 
 
 
07f2e7f
 
97d5f73
07f2e7f
 
 
a312d10
 
 
7f2efda
 
 
 
 
 
 
 
 
 
 
 
 
 
 
97d5f73
 
 
 
648cd8a
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
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
"""A small Hugging Face Space for exploring genome annotations by accession."""
from pathlib import Path
from contextlib import closing
import json
import os
import re
import tempfile
import time

# Shared compute hosts can have /tmp/gradio owned by a different user.
os.environ.setdefault("GRADIO_TEMP_DIR", str(Path(tempfile.gettempdir()) / f"genbank-explorer-{os.getuid()}"))

import gradio as gr
import numpy as np
import pyarrow.parquet as pq
import plotly.graph_objects as go

from taxonomy import build_taxonomy_tab
from style import APP_CSS, atlas_theme
from catalog import Catalog
from remote_catalog import RemoteCatalog, RemoteReadError

HIST_ROWS = 40
ABOUT = Path(__file__).parent / "content/about.md"
MORE = "<!-- more -->"


def tab_intro():
    """What was annotated and how, at the top of a tab, from content/about.md.

    The text lives in one Markdown file so both tabs say the same thing and
    rewording it needs no code change. Above the "more" marker is a short lead;
    below it, the method, in a closed disclosure.
    """
    if not ABOUT.exists():
        return
    lead, _, more = ABOUT.read_text().partition(MORE)
    lead, more = (re.sub(r"<!--.*?-->", "", part, flags=re.S).strip() for part in (lead, more))
    with gr.Column(elem_classes="tab-intro"):
        if lead:
            gr.Markdown(lead, elem_classes="tab-intro-lead")
        if more:
            with gr.Accordion("How the annotations were made", open=False, elem_classes="atlas-disclosure"):
                gr.Markdown(more)
# Readable headers for the results table; the catalog keeps its own names.
HEADERS = {"assembly_accession": "Assembly", "record_name": "Record", "organism_name": "Organism",
           "division": "Division", "segment_start_bp": "Start", "segment_end_bp": "End"}


def build_app(catalog=None):
    if catalog is None:
        catalog = Catalog() if os.environ.get("GENBANK_DATA_MODE") == "sample" else RemoteCatalog()
    remote_mode = isinstance(catalog, RemoteCatalog)
    all_ids = catalog.browse_ids()
    first = catalog.records[all_ids[0]]
    full_snapshot = remote_mode and catalog.manifest.get("full_snapshot", False)
    scope = "published annotation snapshot" if full_snapshot else "indexed subset" if remote_mode else "sample"

    def display_table(ids):
        frame = catalog.table(ids)
        frame["Segment"] = [f"{i + 1} of {n}" for i, n in zip(frame.pop("segment_index"), frame.pop("segment_count"))]
        return frame.rename(columns=HEADERS)

    def search(accession):
        began = time.perf_counter()
        ids, total = catalog.find(accession)
        elapsed = time.perf_counter() - began
        if not str(accession or "").strip():
            message = "Enter an assembly or contig accession, or try an example."
        elif not ids:
            message = f"No match in this {scope}. Newer bucket publications may not be indexed yet." if full_snapshot else f"No match in this {scope}. This does not mean the accession is absent from the full bucket."
        else:
            message = f"Found **{total:,} indexed segment(s)** in {elapsed * 1000:.1f} ms. Showing {len(ids):,}. Assembly coverage may be partial."
        return (gr.Markdown(message, visible=True), gr.Dataframe(value=display_table(ids), visible=bool(ids)),
                ids, ids[0] if ids else None, gr.DownloadButton(visible=False))

    def make_plot(frame, mode, threshold):
        binary = mode == "Binary labels"
        column = "Predicted CDS" if binary else "P(CDS)"
        figure = go.Figure()
        if "Bases" in frame.columns:
            # A column of the heatmap is the distribution of per-base
            # probabilities there, so a region that is part exon and part intron
            # shows both bands instead of an average lying between them.
            positions = frame["Position (bp)"].to_numpy()[::HIST_ROWS]
            centres = frame["P(CDS)"].to_numpy()[:HIST_ROWS]
            bases = frame["Bases"].to_numpy().reshape(len(positions), HIST_ROWS).T
            figure.add_trace(go.Heatmap(
                x=positions, y=centres, z=np.log10(bases + 1), customdata=bases,
                colorscale=[[0, "#fffefa"], [0.25, "#cfe0cd"], [0.6, "#63a07f"], [1, "#173c30"]],
                colorbar=dict(title=dict(text="bases", side="right"), thickness=12,
                              tickvals=[0, 1, 2, 3, 4], ticktext=["1", "10", "100", "1k", "10k"]),
                hovertemplate="%{customdata:,} bases near P=%{y:.2f}<br>from %{x:,}<extra></extra>"))
            figure.add_trace(go.Scatter(
                x=positions, y=frame["Mean P"].to_numpy()[::HIST_ROWS], mode="lines", name="mean per column",
                line=dict(color="#c98b5b", width=1), hovertemplate="mean %{y:.3f}<extra></extra>"))
            figure.add_hline(y=threshold, line_dash="dot", line_color="#8b9b7b",
                             annotation_text=f"Threshold {threshold:g}")
            figure.update_layout(title="CDS probability distribution", xaxis_title="Position (bp; 0-based)",
                                 yaxis_title="P(CDS)", height=380, margin=dict(l=60, r=25, t=75, b=50),
                                 template="plotly_white", paper_bgcolor="#fffefa", plot_bgcolor="#fffefa",
                                 font=dict(family="Arial, Helvetica, sans-serif", color="#315641", size=12),
                                 title_font=dict(family="Georgia, Times New Roman, serif", size=22, color="#173c30"),
                                 hoverlabel=dict(bgcolor="#173c30", font_color="#ffffff", bordercolor="#173c30"),
                                 showlegend=True, legend=dict(orientation="h", y=1.12, x=1, xanchor="right"))
            figure.update_xaxes(gridcolor="#e9edde", zerolinecolor="#dbe3d3", linecolor="#dbe3d3")
            figure.update_yaxes(range=[0, 1], gridcolor="#e9edde", zerolinecolor="#dbe3d3", linecolor="#dbe3d3")
            return figure
        tracks = [("CDS (either strand)", "#287557")] if binary else [("+ strand", "#287557"), ("− strand", "#ae754b")]
        for strand, color in tracks:
            rows = frame[frame["Strand"] == strand]
            figure.add_trace(go.Scatter(x=rows["Position (bp)"].tolist(), y=rows[column].tolist(),
                                       name=strand, mode="lines",
                                       line=dict(color=color, width=2, dash="dash" if strand == "− strand" else "solid",
                                                 shape="hv" if binary else "linear")))
        if not binary:
            figure.add_hline(y=threshold, line_dash="dot", line_color="#8b9b7b",
                            annotation_text=f"Threshold {threshold:g}")
        figure.update_layout(title="Predicted CDS, either strand" if binary else "CDS probability by strand",
                             xaxis_title="Position (bp; 0-based)", yaxis_title=column,
                             height=380, margin=dict(l=60, r=25, t=75, b=50), template="plotly_white",
                             paper_bgcolor="#fffefa", plot_bgcolor="#fffefa",
                             font=dict(family="Arial, Helvetica, sans-serif", color="#315641", size=12),
                             title_font=dict(family="Georgia, Times New Roman, serif", size=22, color="#173c30"),
                             hoverlabel=dict(bgcolor="#173c30", font_color="#ffffff", bordercolor="#173c30"),
                             hovermode="x unified", legend=dict(orientation="h", y=1.12, x=1, xanchor="right"))
        figure.update_xaxes(gridcolor="#e9edde", zerolinecolor="#dbe3d3", linecolor="#dbe3d3")
        figure.update_yaxes(gridcolor="#e9edde", zerolinecolor="#dbe3d3", linecolor="#dbe3d3")
        figure.update_yaxes(range=[-0.05, 1.05], tickvals=[0, 1] if binary else None)
        return figure

    def select_segment(index, mode="Probabilities", threshold=0.5):
        hide = gr.DownloadButton(visible=False)
        if index is None:
            return {}, None, None, None, "Pick a segment above to see its coding landscape.", hide
        record = catalog.records[int(index)]
        start = record["segment_start_bp"]
        end = record["segment_end_bp"]
        try:
            table, stats = catalog.fetch(index)
            frame, step = catalog.window(index, start, end, table=table, mode=mode, threshold=threshold)
            plot = make_plot(frame, mode, threshold)
        except (ValueError, TypeError, OverflowError) as exc:
            return record, start, end, None, str(exc), hide
        return record, start, end, plot, plot_note(step, stats, mode, threshold), hide

    def plot_note(step, stats, mode="Probabilities", threshold=0.5):
        origin = "local sample" if stats.get("local") else "cache" if stats["cache_hit"] else "bucket"
        if mode == "Binary labels":
            detail = (f"1 where the higher strand exceeds {threshold:g}" if step == 1 else
                      f"each step covers up to {step:,} bases and is 1 if any of them exceeds {threshold:g}; zoom in for exact labels")
        else:
            detail = ("one point per base, per strand" if step == 1 else
                      f"each column covers up to {step:,} bases; colour counts how many sit at each probability of the higher strand")
        fetched = f", {stats['bytes_read'] / 1_000_000:.2f} MB" if origin == "bucket" else ""
        return f"0-based, end-exclusive · {detail} · loaded from {origin} in {stats['seconds']:.2f} s{fetched}"

    def update_window(index, start, end, mode="Probabilities", threshold=0.5):
        if index is None:
            return None, "Look up an accession and pick a segment first."
        try:
            table, stats = catalog.fetch(index)
            frame, step = catalog.window(index, start, end, table=table, mode=mode, threshold=threshold)
            plot = make_plot(frame, mode, threshold)
        except (ValueError, TypeError, OverflowError) as exc:
            return None, str(exc)
        return plot, plot_note(step, stats, mode, threshold)

    def export(index):
        if index is None:
            raise gr.Error("Choose a segment first.")
        try:
            table = catalog.segment_table(index)
        except RemoteReadError as exc:
            raise gr.Error(str(exc)) from exc
        assembly = re.sub(r"[^A-Za-z0-9._-]", "_", table["assembly_accession"][0].as_py())
        record = re.sub(r"[^A-Za-z0-9._-]", "_", table["record_name"][0].as_py())
        metadata = catalog.records[int(index)]
        start = metadata["segment_start_bp"]
        end = metadata["segment_end_bp"]
        filename = f"{assembly}__{record}__{start}-{end}.parquet"
        target = Path(tempfile.mkdtemp(prefix="genbank-export-")) / filename
        pq.write_table(table, target, compression="zstd")
        return gr.DownloadButton(value=str(target), visible=True)

    with gr.Blocks(title="GenBank Annotation Explorer", delete_cache=(3600, 3600)) as demo:
        with gr.Tabs(selected="atlas", elem_id="atlas-navigation") as navigation:
            with gr.Tab("Genome Atlas", id="atlas", elem_id="atlas-overview"):
                tab_intro()
                atlas = build_taxonomy_tab()
            with gr.Tab("Database", id="database", elem_id="atlas-database"):
                tab_intro()
                hits = gr.State([])
                selected = gr.State(None)
                with gr.Column(elem_classes="atlas-panel"):
                    gr.HTML('<h2 class="db-section">Find an accession</h2>', apply_default_css=False, elem_classes="db-heading")
                    with gr.Row(equal_height=True, elem_classes="db-search"):
                        accession = gr.Textbox(show_label=False, container=False, scale=5,
                                               placeholder="Assembly (GCA_…) or contig accession")
                        search_button = gr.Button("Find annotations", variant="primary", scale=1, min_width=170)
                    examples = [first["assembly_accession"], first["record_name"]]
                    example_labels = None
                    if remote_mode:
                        examples += catalog.manifest.get("example_record_names", [])[:6]
                        if not full_snapshot:
                            examples += ["JBPJTW010000350.1"]
                    suggestions_path = Path(__file__).parent / "data/suggested_accessions.json"
                    if full_snapshot and suggestions_path.exists():
                        suggestions = json.loads(suggestions_path.read_text())
                        if suggestions["inventory_sha256"] == catalog.manifest.get("inventory_sha256"):
                            examples = [entry["accession"] for entry in suggestions["examples"]]
                            example_labels = [f"{entry['organism']} · {entry['segments']:,} segments" for entry in suggestions["examples"]]
                    gr.Examples(examples=[[e] for e in dict.fromkeys(examples)], inputs=accession,
                                example_labels=example_labels, label="Try", elem_id="annotation-examples")
                    status = gr.Markdown(visible=False, elem_classes="quiet-note")
                    # Hidden until a search returns rows; a click on a row loads that segment.
                    results = gr.Dataframe(value=display_table([]), interactive=False, show_label=False,
                                           visible=False, elem_id="annotation-results")

                with gr.Column(elem_classes="atlas-panel"):
                    gr.HTML('<h2 class="db-section">Coding landscape</h2>', apply_default_css=False, elem_classes="db-heading")
                    with gr.Row(equal_height=True, elem_classes="db-toolbar"):
                        start = gr.Number(label="Start", precision=0, min_width=110, scale=2)
                        end = gr.Number(label="End", precision=0, min_width=110, scale=2)
                        mode = gr.Radio(["Probabilities", "Binary labels"], value="Probabilities", label="View", min_width=320, scale=3)
                        threshold = gr.Slider(0, 1, value=0.5, step=0.01, label="Threshold", min_width=200, scale=3)
                    with gr.Row(elem_classes="db-actions"):
                        view = gr.Button("Update region", variant="primary", size="sm", min_width=140, scale=0)
                        download = gr.Button("Prepare download", size="sm", min_width=150, scale=0, elem_id="db-prepare")
                        file = gr.DownloadButton("Download per-base data (.parquet)", visible=False,
                                                 size="sm", min_width=240, scale=0)
                    plot = gr.Plot(show_label=False, elem_id="annotation-plot")
                    note = gr.Markdown("Pick a segment above to see its coding landscape.", elem_classes="quiet-note")
                    with gr.Accordion("Segment metadata and provenance", open=False, elem_classes="atlas-disclosure"):
                        metadata = gr.JSON(show_label=False)
                    with gr.Accordion("Browse the index and its sources", open=False, elem_classes="atlas-disclosure"):
                        gr.Markdown(f"**{len(catalog.records):,} indexed segments** · **{len(catalog.manifest['assemblies']):,} assemblies** · "
                                    f"{catalog.manifest['bases']:,} bases. Search covers the {scope}; results are model predictions "
                                    "and assembly coverage may be partial.\n\n"
                                    f"Showing the first {len(all_ids):,} indexed segments. Search an accession to find other indexed records. "
                                    "An indexed file does not imply complete coverage of its assembly.\n\n"
                                    "Source: [HuggingFaceBio/genbank-annotations](https://huggingface.co/buckets/HuggingFaceBio/genbank-annotations). "
                                    + (f"Annotations load on demand from {catalog.manifest.get('source_count', len(catalog.manifest['sources']))} bucket files. "
                                       f"Index updated {catalog.manifest['created_at'][:10]}." if remote_mode else "Offline sample."), elem_classes="quiet-note")
                        gr.Dataframe(value=display_table(all_ids), interactive=False, show_label=False)
                        # The manifest lists every indexed assembly, 33,722 of them. Rendered
                        # as a JSON tree that is ~200k DOM nodes, which Gradio builds on the
                        # first visit to this tab: a 3.8 s stall. The count says the same.
                        gr.JSON(value={k: len(v) if k == "assemblies" else v for k, v in catalog.manifest.items()},
                                label="Index provenance" if remote_mode else "Sample provenance")

        found = [status, results, hits, selected, file]
        shown = [metadata, start, end, plot, note, file]

        def pick_row(rows, evt: gr.SelectData):
            return rows[evt.index[0]] if rows and evt.index and evt.index[0] < len(rows) else None

        if atlas:
            # The atlas names a group; the Database tab searches accessions. The
            # jump hands over one annotated assembly from the selected group and
            # runs the ordinary search with it.
            def open_database(path):
                return gr.Tabs(selected="database"), atlas["accession_for"](path)

            atlas["button"].click(open_database, atlas["route"], [navigation, accession]) \
                .then(search, accession, found) \
                .then(select_segment, [selected, mode, threshold], shown)
        for event in (search_button.click, accession.submit):
            event(search, accession, found).then(select_segment, [selected, mode, threshold], shown)
        results.select(pick_row, hits, selected).then(select_segment, [selected, mode, threshold], shown)
        region_inputs = [selected, start, end, mode, threshold]
        view.click(update_window, region_inputs, [plot, note])
        mode.input(update_window, region_inputs, [plot, note])
        threshold.release(update_window, region_inputs, [plot, note])
        # Writing a large segment takes seconds, and the only output is hidden
        # until it is done, so the button itself has to show the work. One
        # generator covers busy, done and failed: a chained .then does not run
        # after an error, which left the button stuck on "Preparing".
        def prepare(index):
            idle = gr.Button("Prepare download", interactive=True)
            yield gr.Button("Preparing download…", interactive=False), gr.DownloadButton(visible=False)
            try:
                ready = export(index)
            except gr.Error:
                yield idle, gr.DownloadButton(visible=False)
                raise
            yield idle, ready

        download.click(prepare, selected, [download, file], show_progress="hidden")
    return demo


if __name__ == "__main__":
    build_app().queue(default_concurrency_limit=2).launch(server_name="0.0.0.0", theme=atlas_theme(), css=APP_CSS)