diff --git a/opustools_pkg/opustools/formatting.py b/opustools_pkg/opustools/formatting.py index a8c182b..3a0e801 100644 --- a/opustools_pkg/opustools/formatting.py +++ b/opustools_pkg/opustools/formatting.py @@ -328,7 +328,8 @@ def moses(sentences, ids): format_fs = {'normal': (normal_src, normal_trg), 'tmx': (tmx_src, tmx_trg), 'moses': (moses, moses), - 'links': (None, None)} + 'links': (None, None), + 'yield_tuple': (moses, moses)} return format_fs[wmode] diff --git a/opustools_pkg/opustools/opus_read.py b/opustools_pkg/opustools/opus_read.py index fa924ad..fcb69a7 100644 --- a/opustools_pkg/opustools/opus_read.py +++ b/opustools_pkg/opustools/opus_read.py @@ -15,7 +15,7 @@ def skip_regex_type(n, N): - "Select function to skip document names" + """Select function to skip document names""" def get_re(doc_name): return not re.search(n, doc_name) @@ -92,7 +92,7 @@ def __init__( download_dir -- Directory where files will be downloaded (default .) preserve_inline_tags -- Preserve inline tags within sentences n -- Get only documents that match the regex - N -- Skip all doucments that match the regex + N -- Skip all documents that match the regex chunk_size -- Number of sentence pairs in chunks to be processed (default 1000000) doc_level -- Print full documents @@ -126,7 +126,7 @@ def __init__( source_annotations, target_annotations = \ target_annotations.copy(), source_annotations.copy() - lang_filters = [src_cld2, src_langid, trg_cld2, trg_langid] + self.lang_filters = lang_filters = [src_cld2, src_langid, trg_cld2, trg_langid] default_alignment = os.path.join( root_directory, directory, release, 'xml', @@ -244,60 +244,21 @@ def doc_level_link_list(self, link_list, src_parser, trg_parser): new_link_list.append(('', tid)) return new_link_list - def printPairs(self): - logger.debug("printPairs called!") - resultfile = None - mosessrc = None - mosestrg = None - id_file = None - - if self.write_ids: - id_file = file_open(self.write_ids, 'w', encoding='utf-8') - - if self.write: - if self.write_mode == 'moses' and len(self.write) == 2: - mosessrc = file_open(self.write[0], mode='w', encoding='utf-8') - mosestrg = file_open(self.write[1], mode='w', encoding='utf-8') - else: - resultfile = file_open(self.write[0], mode='w', encoding='utf-8') - - if self.preprocess == 'moses': - # If preprocessing is moses, download - if not self.write or len(self.write) != 2: - # Write to current path and return - if self.write and len(self.write) != 2: - resultfile.close() - logger.warning('"moses" preprocessing requires two output ' - 'file names. Using default names.') - moses_names = self.of_handler.open_moses_files( - outpath=self.of_handler.download_dir) - logger.info('Moses files written to %s', ', '.join(moses_names)) - return - with tempfile.TemporaryDirectory() as tmpdir: - # Write to specified files - logger.info('Extracting data...') - moses_names = self.of_handler.open_moses_files(outpath=tmpdir) - with file_open(os.path.join(tmpdir, moses_names[0])) as in1, \ - file_open(os.path.join(tmpdir, moses_names[1])) as in2: - if self.switch_langs: - in1, in2 = in2, in1 - for fin, fout in [(in1, mosessrc), (in2, mosestrg)]: - for line in fin: - fout.write(line) - mosessrc.close() - mosestrg.close() - logger.info('Moses files written to %s', ', '.join(self.write)) - return + def _iter_pairs(self, wmode, format_pair, on_doc_start=None, on_doc_end=None): + """Yield (src_result, trg_result, link_attr, src_doc_name, trg_doc_name) + for each alignment link across all documents. - self.add_file_header(resultfile) + Handles chunked link collection, sentence file opening/parsing, + skip_doc, doc_level, and maximum enforcement. + on_doc_start(src_doc_name, trg_doc_name) is called before processing + each document chunk. on_doc_end() is called after processing each + document chunk (even if no pairs are yielded). + """ src_parser = None trg_parser = None - total = 0 - stop = False cur_pos = 0 - prev_src_doc_name = None prev_trg_doc_name = None src_doc_size = -1 @@ -323,67 +284,121 @@ def printPairs(self): if self.skip_doc(src_doc_name): continue - if (self.write_mode != 'links' or - (self.write_mode == 'links' and self.check_lang)): + if wmode != 'links' or (wmode == 'links' and self.check_lang): try: src_doc = self.of_handler.open_sentence_file(src_doc_name, 'src') trg_doc = self.of_handler.open_sentence_file(trg_doc_name, 'trg') except KeyError as e: print('\n'+e.args[0]+'\nContinuing from next sentence file pair.', file=sys.stderr) continue - try: src_parser = SentenceParser( src_doc, preprocessing=self.preprocess, anno_attrs=self.src_annot, - preserve=self.preserve, delimiter=self.annot_delimiter, doc_level=self.doc_level, len_name=self.len_name) + preserve=self.preserve, delimiter=self.annot_delimiter, + doc_level=self.doc_level, len_name=self.len_name) src_doc_size = src_parser.store_sentences(src_set, src_doc_size, self.verbose) trg_parser = SentenceParser( trg_doc, preprocessing=self.preprocess, anno_attrs=self.trg_annot, - preserve=self.preserve, delimiter=self.annot_delimiter, doc_level=self.doc_level, len_name=self.len_name) + preserve=self.preserve, delimiter=self.annot_delimiter, + doc_level=self.doc_level, len_name=self.len_name) trg_doc_size = trg_parser.store_sentences(trg_set, trg_doc_size, self.verbose) except SentenceParserError as e: print('\n'+e.message+'\nContinuing from next sentence file pair.', file=sys.stderr) continue - self.add_doc_names( - src_doc_name, trg_doc_name, resultfile, mosessrc, mosestrg) + if on_doc_start: + on_doc_start(src_doc_name, trg_doc_name) - if self.doc_level and self.write_mode != 'links': + if self.doc_level and wmode != 'links': link_list = self.doc_level_link_list(link_list, src_parser, trg_parser) len_link_list = len(link_list) - for i, link_a in enumerate(link_list): if self.verbose: if i % 1000 == 0 or i + 1 == len_link_list: progress = str(round((i+1)/len_link_list*100, 2)) print("\x1b[2KWriting chunk ... {}%".format(progress), end="\r", file=sys.stderr) - src_result, trg_result = self.format_pair( + src_result, trg_result = format_pair( link_a, src_parser, trg_parser, self.fromto) if src_result == -1: continue - link_attr = attrs_list[i] if i < len(attrs_list) else None - - self.out_put_pair( - src_result, trg_result, resultfile, mosessrc, mosestrg, - link_attr, id_file, src_doc_name, trg_doc_name) + yield (src_result, trg_result, + attrs_list[i] if i < len(attrs_list) else None, + src_doc_name, trg_doc_name) total += 1 if total == self.maximum: - stop = True break - self.add_doc_ending(resultfile) + if on_doc_end: + on_doc_end() if self.verbose and self.write: print("\033[F\033[F\033[F", end="", file=sys.stderr) - if stop: + if total == self.maximum: break + def printPairs(self): + logger.debug("printPairs called!") + resultfile = None + mosessrc = None + mosestrg = None + id_file = None + + if self.write_ids: + id_file = file_open(self.write_ids, 'w', encoding='utf-8') + + if self.write: + if self.write_mode == 'moses' and len(self.write) == 2: + mosessrc = file_open(self.write[0], mode='w', encoding='utf-8') + mosestrg = file_open(self.write[1], mode='w', encoding='utf-8') + else: + resultfile = file_open(self.write[0], mode='w', encoding='utf-8') + + if self.preprocess == 'moses': + if not self.write or len(self.write) != 2: + if self.write and len(self.write) != 2: + resultfile.close() + logger.warning('"moses" preprocessing requires two output ' + 'file names. Using default names.') + moses_names = self.of_handler.open_moses_files( + outpath=self.of_handler.download_dir) + logger.info('Moses files written to %s', ', '.join(moses_names)) + return + with tempfile.TemporaryDirectory() as tmpdir: + logger.info('Extracting data...') + moses_names = self.of_handler.open_moses_files(outpath=tmpdir) + with file_open(os.path.join(tmpdir, moses_names[0])) as in1, \ + file_open(os.path.join(tmpdir, moses_names[1])) as in2: + if self.switch_langs: + in1, in2 = in2, in1 + for fin, fout in [(in1, mosessrc), (in2, mosestrg)]: + for line in fin: + fout.write(line) + mosessrc.close() + mosestrg.close() + logger.info('Moses files written to %s', ', '.join(self.write)) + return + + self.add_file_header(resultfile) + + def doc_start(src, trg): + self.add_doc_names(src, trg, resultfile, mosessrc, mosestrg) + + def doc_end(): + self.add_doc_ending(resultfile) + + for src_result, trg_result, link_attr, src_doc_name, trg_doc_name \ + in self._iter_pairs(self.write_mode, self.format_pair, + on_doc_start=doc_start, on_doc_end=doc_end): + self.out_put_pair( + src_result, trg_result, resultfile, mosessrc, mosestrg, + link_attr, id_file, src_doc_name, trg_doc_name) + if self.verbose and self.write: print("\n\n", file=sys.stderr) @@ -402,3 +417,39 @@ def printPairs(self): id_file.close() self.of_handler.close_zipfiles() + + def yieldPairs(self): + """Yield (source_sentence, target_sentence) tuples from the corpus. + + Applies the same filters as printPairs (maximum, skip_doc, src/tgt_range, + language filters, etc.) but yields plain-text sentence pairs instead of + writing to files or stdout. Useful for programmatic access. + """ + if self.preprocess == 'moses': + with tempfile.TemporaryDirectory() as tmpdir: + moses_names = self.of_handler.open_moses_files(outpath=tmpdir) + with file_open(os.path.join(tmpdir, moses_names[0])) as src_f, \ + file_open(os.path.join(tmpdir, moses_names[1])) as trg_f: + if self.switch_langs: + src_f, trg_f = trg_f, src_f + for src_line, trg_line in zip(src_f, trg_f): + yield src_line.rstrip('\n'), trg_line.rstrip('\n') + return + + form_sent_langs = self.fromto.copy() + if self.switch_langs: + form_sent_langs = [self.fromto[1], self.fromto[0]] + format_sentences = sentence_format_type('yield_tuple', form_sent_langs) + check_filters, _ = check_lang_conf_type(self.lang_filters) + format_pair = pair_format_type( + 'yield_tuple', self.switch_langs, check_filters, False, + format_sentences) + + for src_result, trg_result, _, _, _ \ + in self._iter_pairs('yield_tuple', format_pair): + src = src_result.rstrip('\n').replace('\n', ' ') + tgt = trg_result.rstrip('\n').replace('\n', ' ') + yield src, tgt + + self.alignmentParser.bp.close_document() + self.of_handler.close_zipfiles() diff --git a/opustools_pkg/tests/__init__.py b/opustools_pkg/tests/__init__.py index 44e5380..0211d30 100644 --- a/opustools_pkg/tests/__init__.py +++ b/opustools_pkg/tests/__init__.py @@ -1,7 +1,7 @@ from .test_block_parser import TestBlockParser from .test_sentence_parser import TestSentenceParser from .test_alignment_parser import TestAlignmentParser -from .test_opus_read import TestOpusRead, add_to_root_dir +from .test_opus_read import TestOpusRead, TestOpusReadYieldPairsLocal, add_to_root_dir from .test_opus_cat import TestOpusCat from .test_opus_get import TestOpusGet from .test_opus_langid import TestOpusLangid diff --git a/opustools_pkg/tests/test_opus_read.py b/opustools_pkg/tests/test_opus_read.py index 9f172c2..78d55f6 100644 --- a/opustools_pkg/tests/test_opus_read.py +++ b/opustools_pkg/tests/test_opus_read.py @@ -2250,5 +2250,145 @@ def test_moses_doc_level(self): with open(os.path.join(self.tempdir1, 'doc_level.en-fr.txt')) as doc_out: self.assertEqual(doc_out.readlines(), result) + def test_yield_pairs_xml_maximum(self): + pairs = list(OpusRead( + directory='TEST', source='en', target='fr', + root_directory=self.root_directory, maximum=3).yieldPairs()) + self.assertEqual(len(pairs), 3) + + def test_yield_pairs_xml_all_pairs(self): + pairs = list(OpusRead( + directory='TEST', source='en', target='fr', + root_directory=self.root_directory).yieldPairs()) + self.assertEqual(len(pairs), 7) + self.assertEqual( + pairs[0], + ('e1.1 e1.2 e1.3 e1.4 e1.5', 'f1.1 f1.2 f1.3 f1.4 f1.5')) + self.assertEqual( + pairs[1], + ('e2.1 e2.2 e2.3 e2.4 e2.5', 'f2.1 f2.2 f2.3 f2.4 f2.5')) + self.assertEqual( + pairs[2], + ('e4.1 e4.2 e4.3 e4.4 e4.5', 'f4.1 f4.2 f4.3 f4.4 f4.5')) + self.assertEqual( + pairs[3], + ('e2.1 e2.2 e2.3 e2.4 e2.5', + 'f2.1 f2.2 f2.3 f2.4 f2.5 f3.1 f3.2 f3.3 f3.4 f3.5')) + self.assertEqual( + pairs[4], + ('e4.1 e4.2 e4.3 e4.4 e4.5 e5.1 e5.2 e5.3 e5.4 e5.5', + 'f4.1 f4.2 f4.3 f4.4 f4.5')) + self.assertEqual( + pairs[5], + ('e1.1 e1.2 e1.3 e1.4 e1.5 e2.1 e2.2 e2.3 e2.4 e2.5', + 'f2.1 f2.2 f2.3 f2.4 f2.5')) + self.assertEqual( + pairs[6], + ('e3.1 e3.2 e3.3 e3.4 e3.5', 'f3.1 f3.2 f3.3 f3.4 f3.5')) + + def test_yield_pairs_raw_basic(self): + pairs = list(OpusRead( + directory='TEST', source='en', target='fr', + root_directory=self.root_directory, maximum=1, + preprocess='raw').yieldPairs()) + self.assertEqual(len(pairs), 1) + self.assertEqual( + pairs[0], + ('e1.1 e1.2 e1.3 e1.4 e1.5', 'f1.1 f1.2 f1.3 f1.4 f1.5')) + + def test_yield_pairs_all_maximum_minus_one(self): + pairs = list(OpusRead( + directory='TEST', source='en', target='fr', + root_directory=self.root_directory, maximum=-1).yieldPairs()) + self.assertEqual(len(pairs), 7) + + def test_yield_pairs_switch_langs(self): + pairs = list(OpusRead( + directory='TEST', source='fr', target='en', + root_directory=self.root_directory, maximum=1).yieldPairs()) + self.assertEqual(len(pairs), 1) + self.assertEqual( + pairs[0], + ('f1.1 f1.2 f1.3 f1.4 f1.5', 'e1.1 e1.2 e1.3 e1.4 e1.5')) + + +class TestOpusReadYieldPairsLocal(unittest.TestCase): + """Tests for yieldPairs using locally generated data (no OPUS API).""" + + @classmethod + def setUpClass(self): + self.tempdir = tempfile.mkdtemp() + + os.makedirs(os.path.join(self.tempdir, 'TEST', 'latest', 'xml', 'en')) + os.makedirs(os.path.join(self.tempdir, 'TEST', 'latest', 'xml', 'fr')) + + en_xml = ''' +

+ Hello + world + + Good + morning +

''' + fr_xml = ''' +

+ Bonjour + le + monde + + Bon + matin +

''' + + for lang, content in [('en', en_xml), ('fr', fr_xml)]: + with open(os.path.join( + self.tempdir, 'TEST', 'latest', 'xml', lang, 'doc1.xml'), 'w') as f: + f.write(content) + + with zipfile.ZipFile(os.path.join( + self.tempdir, 'TEST', 'latest', 'xml', 'en.zip'), 'w') as zf: + zf.write(os.path.join(self.tempdir, 'TEST', 'latest', 'xml', 'en', 'doc1.xml'), + arcname='TEST/xml/en/doc1.xml') + with zipfile.ZipFile(os.path.join( + self.tempdir, 'TEST', 'latest', 'xml', 'fr.zip'), 'w') as zf: + zf.write(os.path.join(self.tempdir, 'TEST', 'latest', 'xml', 'fr', 'doc1.xml'), + arcname='TEST/xml/fr/doc1.xml') + + align_xml = ''' + + + + + + + ''' + + align_path = os.path.join(self.tempdir, 'TEST', 'latest', 'xml', 'en-fr.xml') + with open(align_path, 'w') as f: + f.write(align_xml) + with gzip.open(align_path + '.gz', 'wb') as f: + with open(align_path, 'rb') as b: + f.write(b.read()) + + @classmethod + def tearDownClass(self): + shutil.rmtree(self.tempdir) + + def test_yield_pairs_basic(self): + pairs = list(OpusRead( + directory='TEST', source='en', target='fr', + root_directory=self.tempdir).yieldPairs()) + self.assertEqual(len(pairs), 2) + self.assertEqual(pairs[0], ('Hello world', 'Bonjour le monde')) + self.assertEqual(pairs[1], ('Good morning', 'Bon matin')) + + def test_yield_pairs_maximum_1(self): + pairs = list(OpusRead( + directory='TEST', source='en', target='fr', + root_directory=self.tempdir, maximum=1).yieldPairs()) + self.assertEqual(len(pairs), 1) + self.assertEqual(pairs[0], ('Hello world', 'Bonjour le monde')) + + if __name__ == '__main__': unittest.main()