lvwerra's picture
lvwerra HF Staff
Describe the dataset at the top of both tabs (#15)
19eb242
Raw History Blame Contribute Delete
19.1 kB
"""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)