CoolFace
Datasetpublic

thewall/Simulation

sourceHugging Faceopenrailupdated 3y agoView on Hugging Face
0likes6downloads
simulation.py261 linesDownload Raw Back to root
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")