Buckets:
| """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)} | |
Xet Storage Details
- Size:
- 2.3 kB
- Xet hash:
- 8a59076d26fff9822f5bf3fc69b4957b31e883a05f422a362afada21f461d8f2
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.