JavRedstone's picture
download
raw
2.98 kB
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.