Carbon-A-1.2B / load_data.py
cgeorgiaw's picture
cgeorgiaw HF Staff
Package original Carbon-A checkpoint with packed-data inference example
d5a62f1 verified
Raw History Blame Contribute Delete
2.3 kB
"""Minimal loader for the released 16,384-token GENERanno annotation HDF5s."""
from bisect import bisect_right
from pathlib import Path
import h5py
import numpy as np
import torch
from torch.utils.data import Dataset
class PackedAnnotationDataset(Dataset):
"""Read local downloaded shards; one row = 98,304 bp with two label tracks.
No augmentation, retokenization, or context truncation is applied here.
Supply special IDs from the bundled tokenizer to reproduce training masking.
"""
def __init__(self, paths, ignored_token_ids=(0, 1, 2, 3, 4)):
self.paths = sorted(Path(p) for p in paths)
if not self.paths:
raise ValueError("Provide at least one downloaded .h5 file")
self.ignored_token_ids = np.asarray(ignored_token_ids, dtype=np.int64)
self.ends = []
total = 0
for path in self.paths:
with h5py.File(path, "r") as f:
n, tokens = f['input_ids'].shape
if tokens != 16384 or any(f[k].shape != (n, tokens * 6) for k in ('label_plus', 'label_minus')):
raise ValueError(f"Expected aligned 16k annotation rows in {path}")
total += n
self.ends.append(total)
def __len__(self):
return self.ends[-1]
def __getitem__(self, index):
if not 0 <= index < len(self):
raise IndexError(index)
file_index = bisect_right(self.ends, index)
row = index - (self.ends[file_index - 1] if file_index else 0)
# Open per read: no shared HDF5 handle between DataLoader workers.
with h5py.File(self.paths[file_index], "r") as f:
ids = np.asarray(f['input_ids'][row], dtype=np.int64)
plus = (f['label_plus'][row] != 0).astype(np.int64)
minus = (f['label_minus'][row] != 0).astype(np.int64)
labels = np.concatenate([plus, minus])
ignored_bases = np.repeat(np.isin(ids, self.ignored_token_ids), 6)
labels[np.concatenate([ignored_bases, ignored_bases])] = -100
return {'input_ids': torch.from_numpy(ids),
# Fixed-length pretokenized rows; matches the original collator.
'attention_mask': torch.ones(len(ids), dtype=torch.long),
'labels': torch.from_numpy(labels)}