File size: 2,304 Bytes
d5a62f1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""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)}