cgeorgiaw's picture
cgeorgiaw HF Staff
Remove Wet Lab placeholder and use two equal-width navigation tabs
f733efd verified
Raw History Blame
15 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 pyarrow.parquet as pq
import plotly.graph_objects as go
from taxonomy import build_taxonomy_tab
from style import APP_CSS, atlas_theme, section_header
from catalog import Catalog
from remote_catalog import RemoteCatalog, RemoteReadError
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 search(accession):
began = time.perf_counter()
ids, total = catalog.find(accession)
elapsed = time.perf_counter() - began
choices = [(f"{catalog.records[i]['record_name']} · {catalog.records[i]['assembly_accession']} · "
f"[{catalog.records[i]['segment_start_bp']:,}, {catalog.records[i]['segment_end_bp']:,})", str(i)) for i in ids]
if not str(accession or "").strip():
message = "Enter an assembly or contig accession. Try an example below."
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 message, catalog.table(ids), gr.Dropdown(choices=choices, value=str(ids[0]) if ids else None), None
def make_plot(frame, mode, threshold):
binary = mode == "Binary labels"
column = "Predicted CDS" if binary else "P(CDS)"
figure = go.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):
if index is None:
return {}, None, None, None, "Choose a matching segment.", None
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), None
return record, start, end, plot, plot_note(step, stats, mode, threshold), None
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":
resolution = (f"**1 = max(P_positive, P_negative) > {threshold:g}; 0 = otherwise.** "
"This is one CDS/background label per base, combining both strands. "
+ ("The stepped track preserves every base label in this region." if step == 1 else
f"**Binned overview:** each interval spans up to {step:,} bases and is 1 if **any** base exceeds the threshold. "
"This does not mean every base in that interval is CDS. Narrow the region for exact labels."))
else:
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."
return ("Coordinates are **0-based, end-exclusive**. "
+ resolution
+ " Download the segment for the original per-base probabilities.\n\n"
+ f"Loaded from **{origin}** in **{stats['seconds']:.2f} s**"
+ (f" · {stats['bytes_read'] / 1_000_000:.2f} MB fetched." if origin == "bucket" else "."))
def update_window(index, start, end, mode="Probabilities", threshold=0.5):
if index is None:
return None, "Look up an accession and choose 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 str(target)
with gr.Blocks(title="GenBank Annotation Explorer", delete_cache=(3600, 3600)) as demo:
with gr.Tabs(selected="atlas", elem_id="atlas-navigation"):
with gr.Tab("Genome Atlas", id="atlas", elem_id="atlas-overview"):
build_taxonomy_tab()
with gr.Tab("Database", id="database", elem_id="atlas-database"):
gr.HTML('<header class="database-heading"><p class="database-eyebrow">THE ANNOTATION COLLECTION</p>'
'<h1>Explore the database</h1><p>Find an accession, explore its coding landscape, and download the annotations.</p></header>',
apply_default_css=False)
with gr.Column(elem_classes="atlas-panel"):
gr.HTML(section_header("01", "Find an accession", "Start with an assembly or contig to explore its predicted coding regions."), apply_default_css=False)
gr.Markdown(f"**{len(catalog.records):,} indexed segments** · **{len(catalog.manifest['assemblies']):,} assemblies** · "
f"{catalog.manifest['bases']:,} bases\n\n"
f"Search covers the {scope}. Results are model predictions and assembly coverage may be partial.",
elem_classes=["quiet-note", "scope-note"])
with gr.Row():
accession = gr.Textbox(label="Accession ID", placeholder="Assembly (GCA_…) or contig accession", scale=5)
search_button = gr.Button("Find annotations", variant="primary", scale=1, elem_classes="action-button")
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="Explore an assembly with many segments" if example_labels else "Try an accession",
elem_id="annotation-examples")
status = gr.Markdown("Enter an accession or select an example. IDs are case-insensitive; version suffixes are optional.", elem_classes="quiet-note")
results = gr.Dataframe(value=catalog.table([]), interactive=False, label="Matching segments", elem_id="annotation-results")
segment = gr.Dropdown(choices=[], label="Segment to explore", interactive=True)
with gr.Column(elem_classes="atlas-panel"):
gr.HTML(section_header("02", "Explore the coding landscape", "View CDS probabilities or apply a threshold to see one label per base."), apply_default_css=False)
with gr.Row():
start = gr.Number(label="Start (0-based, inclusive)", precision=0)
end = gr.Number(label="End (exclusive)", precision=0)
view = gr.Button("Update region", elem_classes="action-button")
with gr.Row():
mode = gr.Radio(["Probabilities", "Binary labels"], value="Probabilities", label="Viewer mode")
threshold = gr.Slider(0, 1, value=0.5, step=0.01, label="CDS threshold (max strand P > threshold)")
plot = gr.Plot(label="CDS tracks", elem_id="annotation-plot")
note = gr.Markdown("Choose a segment above to bring its coding landscape into view.", elem_classes="quiet-note")
with gr.Accordion("Segment metadata and provenance", open=False, elem_classes="atlas-disclosure"):
metadata = gr.JSON(label="Source metadata")
with gr.Column(elem_classes="atlas-panel"):
gr.HTML(section_header("03", "Take the annotations with you", "Download the original segment, with its accession and record name in the filename."), apply_default_css=False)
with gr.Row():
download = gr.Button("Prepare segment download", variant="primary", scale=1, elem_classes="action-button")
file = gr.File(label="Original segment annotations (Parquet)", interactive=False, scale=3)
gr.Markdown("Full per-base probabilities are preserved in the download, including when the viewer shows a binned overview.", elem_classes="quiet-note")
with gr.Accordion("Browse the index and its sources", open=False, elem_classes="atlas-disclosure"):
gr.Markdown(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=catalog.table(all_ids), interactive=False)
gr.JSON(value=catalog.manifest, label="Index provenance" if remote_mode else "Sample provenance")
gr.HTML('<footer class="workspace-footer"><span>Hugging Face Bio · Genome Atlas</span>'
'<a href="https://huggingface.co/buckets/HuggingFaceBio/genbank-annotations" target="_blank" rel="noopener noreferrer">Explore the annotation collection ↗</a></footer>', apply_default_css=False)
outputs = [status, results, segment, file]
for event in (search_button.click, accession.submit):
event(search, accession, outputs).then(select_segment, [segment, mode, threshold], [metadata, start, end, plot, note, file])
segment.input(select_segment, [segment, mode, threshold], [metadata, start, end, plot, note, file])
region_inputs = [segment, 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])
download.click(export, segment, file)
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)