"""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()