Buckets:
| import numpy as np | |
| import os | |
| import torch | |
| from torch.utils.data import Dataset | |
| from datasets import load_dataset | |
| from .nucleotide_tokenizers import FixedSizeNucleotidesKmersTokenizer, compute_tokens_to_ids_v2 | |
| cache_dir = os.environ.get("POPE_HF_CACHE_DIR") | |
| num_proc = 1 | |
| class HRGDataset(Dataset): | |
| def __init__(self, split: str, max_length: int, k_for_kmers: int = 6, chunk_size: int = 6100, overlap: int = 50): | |
| self.k_for_kmers = k_for_kmers | |
| self.chunk_size = chunk_size | |
| self.overlap = overlap | |
| self.split = split | |
| self.hf_dataset = load_dataset("InstaDeepAI/human_reference_genome", split=split, | |
| num_proc=num_proc, cache_dir=cache_dir, trust_remote_code=True | |
| ) | |
| tokens_to_ids, standard_tokens = compute_tokens_to_ids_v2(k_mers=k_for_kmers) | |
| self.raw_seq_length = len(self.hf_dataset[0]['sequence']) | |
| self.tokenizer = FixedSizeNucleotidesKmersTokenizer( | |
| k_mers=k_for_kmers, fixed_length=max_length, | |
| tokens_to_ids=tokens_to_ids, | |
| prepend_cls_token=True | |
| ) | |
| self.hrg_vocab_size = self.tokenizer.vocabulary_size | |
| self.hrg_pad_token_id = self.tokenizer.pad_token_id | |
| self.max_length = max_length | |
| def __len__(self): | |
| return len(self.hf_dataset) | |
| def convert_from_long_idx(self, idx): | |
| return idx // self.raw_seq_length, idx % self.raw_seq_length | |
| def get_overlapping_chunks(self, idx): | |
| # split into overlapping chunks of size `chunk_size` with `overlap` elements shared | |
| start_idx = idx * self.chunk_size - self.overlap | |
| end_idx = start_idx + self.chunk_size | |
| start_seq_j, start_k = self.convert_from_long_idx(start_idx) | |
| end_seq_j, end_k = self.convert_from_long_idx(end_idx) | |
| if start_seq_j == end_seq_j: | |
| chunk = self.hf_dataset[start_seq_j]['sequence'][start_k:end_k] | |
| return chunk | |
| start_seq = self.hf_dataset[start_seq_j]['sequence'] | |
| end_seq = self.hf_dataset[end_seq_j]['sequence'] | |
| chunk = start_seq[start_k:] + end_seq[:end_k] | |
| return chunk | |
| def __getitem__(self, idx): | |
| chunk = self.get_overlapping_chunks(idx) | |
| start_idx = 0 | |
| # data augmentation: random start point from the first 100 nucleotides | |
| if self.split == 'train': | |
| start_idx = np.random.randint(0, 100, size=1)[0] | |
| # Tokenize the overlapping chunk. | |
| tokens, token_ids = self.tokenizer.tokenize(chunk[start_idx:]) | |
| # Crop if longer than max_length | |
| if len(token_ids) > self.max_length: | |
| final_token_ids = token_ids[:self.max_length] | |
| # Pad if shorter than max_length | |
| else: | |
| padding_to_add = self.max_length - len(token_ids) | |
| final_token_ids = token_ids + [self.tokenizer.pad_token_id] * padding_to_add | |
| return {'input_ids': torch.tensor(final_token_ids, dtype=torch.int64)} | |
Xet Storage Details
- Size:
- 2.98 kB
- Xet hash:
- 644bc0c1af3b5a6e8f6980992555d27dcb63b9873a8b2cf7e4e162ad50fa16cd
·
Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.