thewall/Simulation
06
1import os2import numpy as np3from enum import IntEnum4import datasets5 6 7logger = datasets.logging.get_logger(__name__)8 9 10_CITATION = """\11@article{iwano2022generative,12 title={Generative aptamer discovery using RaptGen},13 author={Iwano, Natsuki and Adachi, Tatsuo and Aoki, Kazuteru and Nakamura, Yoshikazu and Hamada, Michiaki},14 journal={Nature Computational Science},15 pages={1--9},16 year={2022},17 publisher={Nature Publishing Group}18}19"""20 21_DESCRIPTION = """\22https://github.com/hmdlab/raptgen/blob/master/raptgen/data.py23"""24 25 26class SNV(IntEnum):27 Mutation = 028 Insertion = 129 Deletion = 230 31 32class SequenceGenerator():33 def __init__(self, num_motifs=1, motif_length=10, motifs=None,34 target_length=20, fix_random_region_length=True, error_rate=0.0, generate_motifs=True, middle_insert_range=(2, 6),35 seed=0, add_primer=True, forward_primer="AAAAA", reverse_primer="GGGGG", one_side_proba=0.5, paired=False):36 np.random.seed(seed)37 38 if generate_motifs:39 self.motifs = ["".join(np.random.choice(40 list("ATGC"), motif_length)) for _ in range(num_motifs)]41 else:42 self.motifs = motifs43 44 self.error_indices = 1 + \45 np.argsort(np.random.random(size=motif_length-1))[:3]46 self.mut_idx, self.ins_idx, self.del_idx = self.error_indices47 48 logger.info(f"error rate is {error_rate*100:.1f}%")49 for idx, motif in enumerate(self.motifs):50 seq = [ch for ch in motif]51 mut = self.mutate(seq[self.mut_idx])52 if error_rate != 0:53 seq[self.mut_idx] = f"[{seq[self.mut_idx]}>{mut}]"54 seq[self.ins_idx] = f"[+]{seq[self.ins_idx]}"55 seq[self.del_idx] = f"{seq[self.del_idx].lower()}"56 seq = "".join(seq)57 logger.info(f"motif {idx} is {seq}")58 59 self.num_motifs = num_motifs60 self.error_rate = error_rate61 self.target_length = target_length62 self.forward_primer = forward_primer63 self.reverse_primer = reverse_primer64 self.add_primer = add_primer65 66 self.one_side_proba = one_side_proba67 self.middle_insert_range = middle_insert_range68 self.paired = paired69 70 def mutate(self, char):71 return "TGCA"["ATGC".index(char)]72 73 def sample_motif(self, n):74 motif_indices = np.random.randint(self.num_motifs, size=n)75 has_errors = np.random.random(size=n) < self.error_rate76 # mutation, insertion, deletion77 error_types = np.random.choice(SNV, size=n)78 sequences = []79 valid_masks = []80 for motif_index, has_error, error_type in zip(motif_indices, has_errors, error_types):81 motif = self.motifs[motif_index]82 seq = [ch for ch in motif]83 mask = [1]*len(motif)84 if has_error:85 if error_type == SNV.Mutation:86 seq[self.mut_idx] = self.mutate(seq[self.mut_idx])87 mask[self.mut_idx] = 088 elif error_type == SNV.Insertion:89 seq[self.ins_idx] = np.random.choice(90 list("ATGC")) + seq[self.ins_idx]91 mask.insert(self.ins_idx, 0)92 elif error_type == SNV.Deletion:93 seq[self.del_idx] = ""94 del mask[self.del_idx]95 else:96 raise NotImplementedError97 seq = "".join(seq)98 sequences.append(seq)99 valid_masks.append(mask)100 return sequences, valid_masks, motif_indices.tolist()101 102 def sample(self, n=1, with_indices=True):103 motifs, valid_masks, motif_indices = self.sample_motif(n)104 sequences = []105 motif_masks = []106 paired_indices = []107 for seq, mask in zip(motifs, valid_masks):108 if self.paired:109 seq, mask, idx = self.insert_in_the_middle(110 seq, mask, nrange=self.middle_insert_range, one_side_proba=self.one_side_proba)111 paired_indices += [idx]112 random_region = "".join(np.random.choice(113 list("ATGC"), size=self.target_length-len(seq)))114 l = np.random.randint(len(random_region))115 if self.add_primer:116 sequences.append(117 self.forward_primer + random_region[:l] + seq + random_region[l:] + self.reverse_primer)118 motif_masks.append([0]*(len(self.forward_primer)+l)+mask+[0]*(len(random_region)-l+len(self.reverse_primer)))119 else:120 sequences.append(random_region[:l] + seq + random_region[l:])121 motif_masks.append([0]*l+mask+[0]*(len(random_region)-l))122 123 if self.paired and with_indices:124 return sequences, motif_masks, motif_indices, paired_indices125 elif with_indices:126 return sequences, motif_masks, motif_indices127 return sequences, motif_masks128 129 def insert_in_the_middle(self, sequence, mask, nrange=(2, 6), one_side_proba=0.5):130 n = np.random.randint(*nrange)131 if np.random.random() < one_side_proba:132 if np.random.choice(["l", "r"]) == "l":133 l_motif = sequence[:len(sequence)//2]134 r_motif = ""135 idx = 1136 else:137 l_motif = ""138 r_motif = sequence[len(sequence)//2:]139 idx = 2140 else:141 l_motif = sequence[:len(sequence)//2]142 r_motif = sequence[len(sequence)//2:]143 idx = 0144 seq = l_motif + "".join(np.random.choice(list("ATGC"), size=n)) + r_motif145 new_mask = mask[:len(l_motif)]+([0]*n)+mask[len(sequence)-len(r_motif):]146 return seq, new_mask, idx147 148 149DATA_FILES = {"multiple-666": {"train": "https://huggingface.co/datasets/thewall/Simulation/resolve/main/data/multiple-666-train.parquet",150 "test": "https://huggingface.co/datasets/thewall/Simulation/resolve/main/data/multiple-666-test.parquet"},151 "paired-666": {"train": "https://huggingface.co/datasets/thewall/Simulation/resolve/main/data/paired-666-train.parquet",152 "test": "https://huggingface.co/datasets/thewall/Simulation/resolve/main/data/paired-666-test.parquet"},153 "paired-42": {"train": "https://huggingface.co/datasets/thewall/Simulation/resolve/main/data/paired-42-train.parquet",154 "test": "https://huggingface.co/datasets/thewall/Simulation/resolve/main/data/paired-42-test.parquet"},155 }156 157class SimulationConfig(datasets.BuilderConfig):158 def __init__(self, n_seq, num_motifs=1, motif_length=10, error_rate=0.0, seed=0, add_primer=False, paired=False, **kwargs):159 super(SimulationConfig, self).__init__(**kwargs)160 self.n_seq = n_seq161 self.num_motifs = num_motifs162 self.motif_length = motif_length163 self.error_rate = error_rate164 self.seed = seed165 self.add_primer = add_primer166 self.paired = paired167 # if "paired" in kwargs['name']:168 # self.paired = True169 # else:170 # self.paired = False171 172 173class Simulation(datasets.GeneratorBasedBuilder):174 175 BUILDER_CONFIGS = [176 SimulationConfig(name="multiple", num_motifs=10, error_rate=0.1, n_seq=10000, seed=0),177 SimulationConfig(name="paired", n_seq=5000, seed=0, paired=True),178 SimulationConfig(name="multiple-666", num_motifs=10, error_rate=0.1, n_seq=10000, seed=0),179 SimulationConfig(name="paired-666", n_seq=5000, seed=0, paired=True),180 SimulationConfig(name="paired-42", n_seq=10000, seed=0, paired=True),181 ]182 183 DEFAULT_CONFIG_NAME = "multiple-666"184 185 def _info(self):186 return datasets.DatasetInfo(187 description=_DESCRIPTION,188 features=datasets.Features(189 {190 "id": datasets.Value("int32"),191 "seq": datasets.Value("string"),192 "motif": datasets.Value("string"),193 "motif_ids": datasets.Value("int32"),194 "motif_mask": datasets.Sequence(feature=datasets.Value("int32")),195 }196 ),197 homepage="https://github.com/hmdlab/raptgen/blob/master/raptgen/data.py",198 citation=_CITATION,199 )200 201 def _split_generators(self, dl_manager):202 if self.config.name in DATA_FILES:203 train_data_file = dl_manager.download(DATA_FILES[self.config.name]['train'])204 test_data_file = dl_manager.download(DATA_FILES[self.config.name]['test'])205 dataset = datasets.load_dataset("parquet", data_files={"train": train_data_file,206 "test": test_data_file})207 train_iterator = self._iterator(dataset['train'])208 test_iterator = self._iterator(dataset['test'])209 return [210 datasets.SplitGenerator(name=datasets.Split.TRAIN, gen_kwargs={"iterator_fn": train_iterator}),211 datasets.SplitGenerator(name=datasets.Split.TEST, gen_kwargs={"iterator_fn": test_iterator}),212 ]213 else:214 kwargs = {"num_motifs": self.config.num_motifs,215 "motif_length": self.config.motif_length,216 "error_rate": self.config.error_rate,217 "seed": self.config.seed,218 "add_primer": self.config.add_primer,219 "sample_num": self.config.n_seq,220 "paired": self.config.paired221 }222 iterator = self._sample(**kwargs)223 224 return [225 datasets.SplitGenerator(name=datasets.Split.TRAIN, gen_kwargs={"iterator_fn": iterator}),226 ]227 228 def _sample(self, num_motifs, motif_length, error_rate, seed, add_primer, sample_num, paired):229 simulator = SequenceGenerator(num_motifs=num_motifs, motif_length=motif_length,230 error_rate=error_rate, seed=seed,231 add_primer=add_primer, paired=paired)232 data = simulator.sample(sample_num)233 motifs = simulator.motifs234 for key, (seq, mask, motif_ids, label) in enumerate(zip(data[0], data[1], data[2], data[-1])):235 yield key, {"id": key,236 "seq": seq,237 "motif": motifs[motif_ids],238 "motif_ids": label,239 "motif_mask": mask,240 }241 242 def _iterator(self, dataset):243 for row in dataset:244 yield row['id'], row245 246 def _generate_examples(self, iterator_fn):247 yield from iterator_fn248 249 250if __name__=="__main__":251 from datasets import load_dataset252 # splited_data = dataset.train_test_split(train_size=0.9, seed=666)253 # splited_data['train'].to_parquet("paired-666-train.parquet")254 # splited_data['test'].to_parquet("paired-666-test.parquet")255 256 # dataset = load_dataset(path = "thewall/simulation", name="multiple", split="all")257 # splited_data = dataset.train_test_split(train_size=0.9, seed=666)258 # splited_data['train'].to_parquet("multiple-666-train.parquet")259 # splited_data['test'].to_parquet("multiple-666-test.parquet")260 261 dataset = load_dataset("simulation.py", name="paired-666", split="test")