Repository navigation
Expand file tree
/
Copy pathtransform.py
More file actions
80 lines (71 loc) · 4.13 KB
/
Copy pathtransform.py
File metadata and controls
80 lines (71 loc) · 4.13 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
import sys, time
sys.path.append("NL-Augmenter")
import nltk
# pip install spacy torchtext cucco fastpunct sacremoses
# python -m spacy download en_core_web_sm
from nlaugmenter.transformations.butter_fingers_perturbation.transformation import ButterFingersPerturbation
from nlaugmenter.transformations.random_deletion.transformation import RandomDeletion
from nlaugmenter.transformations.synonym_substitution.transformation import SynonymSubstitution
from nlaugmenter.transformations.back_translation.transformation import BackTranslation
from nlaugmenter.transformations.change_char_case.transformation import ChangeCharCase
from nlaugmenter.transformations.whitespace_perturbation.transformation import WhitespacePerturbation
from nlaugmenter.transformations.underscore_trick.transformation import UnderscoreTrick
from nlaugmenter.transformations.style_paraphraser.transformation import StyleTransferParaphraser
from nlaugmenter.transformations.punctuation.transformation import PunctuationWithRules
def aug_generator(text_list, aug_style):
n_empty = sum(1 for text in text_list if text == "")
if n_empty > 0:
print(f'{n_empty} strings are empty in text list.')
if aug_style == "butter_fingers":
t1 = ButterFingersPerturbation(max_outputs=1)
# return [t1.generate(text_list[i], prob = 0.1)[0] for i in range(len(text_list))]
return [t1.generate(text_list[i])[0] for i in range(len(text_list))]
elif aug_style == "random_deletion":
t1 = RandomDeletion(prob=0.25)
result = []
for i in range(len(text_list)):
if len(nltk.word_tokenize(text_list[i])) > 1:
# print(f'text check: [{text_list[i]}], {len(text_list[i])}')
result.append(t1.generate(text_list[i])[0])
else:
result.append(text_list[i])
return result
elif aug_style == "synonym_substitution":
syn = SynonymSubstitution(max_outputs=1, prob = 0.2)
return [syn.generate(text_list[i])[0] for i in range(len(text_list))]
elif aug_style == "back_translation":
trans = BackTranslation()
return [trans.generate(text_list[i])[0] for i in range(len(text_list))]
elif aug_style == "change_char_case":
t1 = ChangeCharCase()
# return [t1.generate(text_list[i], prob = 0.25)[0] for i in range(len(text_list))]
return [t1.generate(text_list[i])[0] for i in range(len(text_list))]
elif aug_style == "whitespace_perturbation":
t1 = WhitespacePerturbation()
# return [t1.generate(text_list[i], prob = 0.25)[0] for i in range(len(text_list))]
return [t1.generate(text_list[i])[0] for i in range(len(text_list))]
elif aug_style == "underscore_trick":
t1 = UnderscoreTrick(prob = 0.25)
return [t1.generate(text_list[i])[0] for i in range(len(text_list))]
elif aug_style == "style_paraphraser":
t1 = StyleTransferParaphraser(style = "Basic", upper_length="same_5")
return [t1.generate(text_list[i])[0] for i in range(len(text_list))]
elif aug_style == "punctuation_perturbation":
normalizations = ['remove_extra_white_spaces', ('replace_characters', {'characters': 'was', 'replacement': 'TZ'}),
('replace_emojis', {'replacement': 'TESTO'})]
punc = PunctuationWithRules(rules=normalizations)
return [punc.generate(text_list[i])[0] for i in range(len(text_list))]
else:
raise ValueError("Augmentation style not found. Please check the available styles.")
def generate_perturbations(text_list):
augmentation_styles = ["synonym_substitution", "butter_fingers", "random_deletion", "change_char_case", "whitespace_perturbation", "underscore_trick"]
all_augmented = {}
for style in augmentation_styles:
start = time.time()
aug_list = aug_generator(text_list, style)
all_augmented[style] = aug_list
print(f"Perturbing with {style} took {time.time() - start} seconds")
return all_augmented
if __name__ == "__main__":
text_list = ["This is a test sentence. It is a good sentence.", "This is another test sentence. It is a bad sentence."]
print(generate_perturbations(text_list))