carbon-a-database-explorer / prepare_sample.py
cgeorgiaw's picture
cgeorgiaw HF Staff
Deploy initial accession explorer with a bounded GenBank annotation sample
97d5f73 verified
Raw History Blame
3.89 kB
"""Extract a bounded, reproducible sample from one published bucket object."""
import argparse
from datetime import datetime, timezone
import hashlib
import json
from pathlib import Path
import tempfile
import pyarrow as pa
import pyarrow.parquet as pq
BUCKET = "HuggingFaceBio/genbank-annotations"
DEFAULT_SOURCE = "annotations/fungi/s0000/s0000.chunk-00510.parquet"
ROOT = Path(__file__).resolve().parent
def extract(source, output, source_path, max_rows=24, max_bp=2_000_000, xet_hash=None):
if max_rows < 1 or max_bp < 1:
raise ValueError("Sample limits must be positive.")
pf = pq.ParquetFile(source)
selected, bases = [], 0
# Read lengths before loading the much larger per-base columns.
lengths = pf.read(columns=["segment_bp_length"]).column(0).to_pylist()
for i, length in enumerate(lengths):
if 0 < length <= max_bp - bases:
selected.append(i)
bases += length
if len(selected) >= max_rows:
break
if not selected:
raise ValueError("No complete segments fit the sample budget.")
pieces, offset = [], 0
for group in range(pf.num_row_groups):
count = pf.metadata.row_group(group).num_rows
indices = [i - offset for i in selected if offset <= i < offset + count]
if indices:
pieces.append(pf.read_row_group(group).take(pa.array(indices)))
offset += count
sample = pa.concat_tables(pieces)
output = Path(output)
output.mkdir(parents=True, exist_ok=True)
target = output / "sample.parquet"
# One group per segment lets the app fetch just the selected segment.
pq.write_table(sample, target, compression="zstd", row_group_size=1)
metadata = sample.select([c for c in sample.column_names if not c.startswith("pred_prob_") and c != "sequence"]).to_pylist()
manifest = {
"bucket_id": BUCKET, "source_path": source_path, "source_xet_hash": xet_hash,
"created_at": datetime.now(timezone.utc).isoformat(),
"selection": "First complete segments fitting the row and base budgets, in source order",
"max_rows": max_rows, "max_bp": max_bp, "rows": sample.num_rows,
"bases": bases, "source_rows": pf.metadata.num_rows, "source_row_indices": selected,
"assemblies": sorted({r["assembly_accession"] for r in metadata}),
"sha256": hashlib.sha256(target.read_bytes()).hexdigest(),
}
(output / "manifest.json").write_text(json.dumps(manifest, indent=2) + "\n")
print(json.dumps(manifest, indent=2))
return manifest
def main():
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--source", default=DEFAULT_SOURCE)
parser.add_argument("--local-file", type=Path)
parser.add_argument("--output", type=Path, default=ROOT / "data")
parser.add_argument("--max-rows", type=int, default=24)
parser.add_argument("--max-bp", type=int, default=2_000_000)
args = parser.parse_args()
if args.local_file:
extract(args.local_file, args.output, args.source, args.max_rows, args.max_bp)
else:
from huggingface_hub import HfApi
api = HfApi()
entries = api.list_bucket_tree(BUCKET, args.source, recursive=False)
entry = next((e for e in entries if e.path == args.source), None)
if entry is None:
raise ValueError(f"Bucket object not found: {args.source}")
if entry.size > 100_000_000:
raise ValueError("Source exceeds the 100 MB download budget. Choose a smaller object or use --local-file.")
with tempfile.TemporaryDirectory() as temp:
local = Path(temp) / "source.parquet"
api.download_bucket_files(BUCKET, [(entry, local)], raise_on_missing_files=True)
extract(local, args.output, args.source, args.max_rows, args.max_bp, entry.xet_hash)
if __name__ == "__main__":
main()