| 46 | |
| 47 | |
| 48 | class Encoder(object): |
| 49 | def __init__(self, args): |
| 50 | self.args = args |
| 51 | |
| 52 | def initializer(self): |
| 53 | # Use Encoder class as a container for global data |
| 54 | Encoder.tokenizer = build_tokenizer(self.args) |
| 55 | if self.args.split_sentences: |
| 56 | if not nltk_available: |
| 57 | print("NLTK is not available to split sentences.") |
| 58 | exit() |
| 59 | if os.environ.get("NLTK_DATA"): |
| 60 | library = os.path.join(os.environ.get("NLTK_DATA"), "tokenizers", "punkt", f"{self.args.lang}.pickle") |
| 61 | url = f"file:{library}" |
| 62 | else: |
| 63 | library = os.path.join("tokenizers", "punkt", f"{self.args.lang}.pickle") |
| 64 | url = f"nltk:{library}" |
| 65 | splitter = nltk.load(url) |
| 66 | if self.args.keep_newlines: |
| 67 | # this prevents punkt from eating newlines after sentences |
| 68 | Encoder.splitter = nltk.tokenize.punkt.PunktSentenceTokenizer( |
| 69 | train_text = splitter._params, |
| 70 | lang_vars = CustomLanguageVars()) |
| 71 | else: |
| 72 | Encoder.splitter = splitter |
| 73 | |
| 74 | else: |
| 75 | Encoder.splitter = IdentitySplitter() |
| 76 | |
| 77 | def split(self, json_line): |
| 78 | data = json.loads(json_line) |
| 79 | output = {} |
| 80 | for key in self.args.json_keys: |
| 81 | text = data[key] |
| 82 | max_len = 1000000 |
| 83 | tokens_list = [Encoder.splitter.tokenize(text[i:i+max_len]) for i in range(0, len(text), max_len)] |
| 84 | output[key] = [tokens for partial in tokens_list for tokens in partial] |
| 85 | return json.dumps(output), len(json_line) |
| 86 | |
| 87 | def encode(self, json_line): |
| 88 | data = json.loads(json_line) |
| 89 | ids = {} |
| 90 | lens = {} |
| 91 | for key in self.args.json_keys: |
| 92 | text = data[key] |
| 93 | if isinstance(text, list): |
| 94 | sentences = text |
| 95 | else: |
| 96 | sentences = [text] |
| 97 | doc_ids = [] |
| 98 | sentence_lens = [] |
| 99 | for sentence in sentences: |
| 100 | sentence_ids = Encoder.tokenizer.tokenize(sentence) |
| 101 | if len(sentence_ids) > 0: |
| 102 | doc_ids.extend(sentence_ids) |
| 103 | sentence_lens.append(len(sentence_ids)) |
| 104 | if len(doc_ids) > 0 and self.args.append_eod: |
| 105 | doc_ids.append(Encoder.tokenizer.eod) |
no outgoing calls