diff --git a/tests/test_augment_api.py b/tests/test_augment_api.py index e2d76589..a46bbd9f 100644 --- a/tests/test_augment_api.py +++ b/tests/test_augment_api.py @@ -93,6 +93,32 @@ def test_deletion_augmenter(): assert augmented_s in augmented_text_list +def test_augment_text_with_ids_keeps_each_original_once(monkeypatch): + from textattack.augmentation import Augmenter + from textattack.transformations import WordDeletion + + augmenter = Augmenter(WordDeletion()) + monkeypatch.setattr( + augmenter, + "augment", + lambda text: [f"{text} augmented-1", f"{text} augmented-2"], + ) + + texts, ids = augmenter.augment_text_with_ids( + ["first", "second"], [1, 2], show_progress=False + ) + + assert texts == [ + "first", + "first augmented-1", + "first augmented-2", + "second", + "second augmented-1", + "second augmented-2", + ] + assert ids == [1, 1, 1, 2, 2, 2] + + def test_high_yield_scales_with_transformations_per_example(): # Regression test: the retry-bound fix for issue #800 (stop a # low-diversity transformation from silently returning fewer than diff --git a/textattack/augmentation/augmenter.py b/textattack/augmentation/augmenter.py index 74a86e1c..97f42694 100644 --- a/textattack/augmentation/augmenter.py +++ b/textattack/augmentation/augmenter.py @@ -239,9 +239,8 @@ def augment_text_with_ids(self, text_list, id_list, show_progress=True): all_text_list.append(text) all_id_list.append(_id) augmented_texts = self.augment(text) - all_text_list.extend - all_text_list.extend([text] + augmented_texts) - all_id_list.extend([_id] * (1 + len(augmented_texts))) + all_text_list.extend(augmented_texts) + all_id_list.extend([_id] * len(augmented_texts)) return all_text_list, all_id_list def __repr__(self):