File size: 3,889 Bytes
97d5f73
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
"""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()