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