(self, ds, max_seq_len=512, mask_lm_prob=.15, max_preds_per_seq=None, short_seq_prob=.01,
dataset_size=None, presplit_sentences=False, weighted=True, **kwargs)
| 851 | """ |
| 852 | |
| 853 | def __init__(self, ds, max_seq_len=512, mask_lm_prob=.15, max_preds_per_seq=None, short_seq_prob=.01, |
| 854 | dataset_size=None, presplit_sentences=False, weighted=True, **kwargs): |
| 855 | self.ds = ds |
| 856 | self.ds_len = len(self.ds) |
| 857 | self.tokenizer = self.ds.GetTokenizer() |
| 858 | self.vocab_words = list(self.tokenizer.text_token_vocab.values()) |
| 859 | self.ds.SetTokenizer(None) |
| 860 | self.max_seq_len = max_seq_len |
| 861 | self.mask_lm_prob = mask_lm_prob |
| 862 | if max_preds_per_seq is None: |
| 863 | max_preds_per_seq = math.ceil(max_seq_len * mask_lm_prob / 10) * 10 |
| 864 | self.max_preds_per_seq = max_preds_per_seq |
| 865 | self.short_seq_prob = short_seq_prob |
| 866 | self.dataset_size = dataset_size |
| 867 | if self.dataset_size is None: |
| 868 | self.dataset_size = self.ds_len * (self.ds_len - 1) |
| 869 | self.presplit_sentences = presplit_sentences |
| 870 | if not self.presplit_sentences: |
| 871 | nltk.download('punkt', download_dir="./nltk") |
| 872 | self.weighted = weighted |
| 873 | self.get_weighting() |
| 874 | |
| 875 | def get_weighting(self): |
| 876 | if self.weighted: |
nothing calls this directly
no test coverage detected