"""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 = ""
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}
from %{x:,}"))
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}"))
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('
Find an accession
', 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('Coding landscape
', 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)