TangibleAI/mathtext
1
1import spacy # noqa2import time3from editdistance import eval as edit_dist4from transformers import pipeline5 6# import os7# os.environ['KMP_DUPLICATE_LIB_OK']='True'8# import spacy9 10# Change this according to what words should be corrected to11SPELL_CORRECT_MIN_CHAR_DIFF = 212 13TOKENS2INT_ERROR_INT = 3220214 15ONES = [16 "zero", "one", "two", "three", "four", "five", "six", "seven", "eight",17 "nine", "ten", "eleven", "twelve", "thirteen", "fourteen", "fifteen",18 "sixteen", "seventeen", "eighteen", "nineteen",19]20 21CHAR_MAPPING = {22 "-": " ",23 "_": " ",24 "and": " ",25}26# CHAR_MAPPING.update((str(i), word) for i, word in enumerate([" " + s + " " for s in ONES]))27TOKEN_MAPPING = {28 "and": " ",29 "oh": "0",30}31 32 33def find_char_diff(a, b):34 # Finds the character difference between two str objects by counting the occurences of every character. Not edit distance.35 char_counts_a = {}36 char_counts_b = {}37 for char in a:38 if char in char_counts_a.keys():39 char_counts_a[char] += 140 else:41 char_counts_a[char] = 142 for char in b:43 if char in char_counts_b.keys():44 char_counts_b[char] += 145 else:46 char_counts_b[char] = 147 char_diff = 048 for i in char_counts_a:49 if i in char_counts_b.keys():50 char_diff += abs(char_counts_a[i] - char_counts_b[i])51 else:52 char_diff += char_counts_a[i]53 return char_diff54 55 56def tokenize(text):57 text = text.lower()58 # print(text)59 text = replace_tokens(''.join(i for i in replace_chars(text)).split())60 # print(text)61 text = [i for i in text if i != ' ']62 # print(text)63 output = []64 for word in text:65 # print(word)66 output.append(convert_word_to_int(word))67 output = [i for i in output if i != ' ']68 # print(output)69 return output70 71 72def detokenize(tokens):73 return ' '.join(tokens)74 75 76def replace_tokens(tokens, token_mapping=TOKEN_MAPPING):77 return [token_mapping.get(tok, tok) for tok in tokens]78 79 80def replace_chars(text, char_mapping=CHAR_MAPPING):81 return [char_mapping.get(c, c) for c in text]82 83 84def convert_word_to_int(in_word, numwords={}):85 # Converts a single word/str into a single int86 tens = ["", "", "twenty", "thirty", "forty", "fifty", "sixty", "seventy", "eighty", "ninety"]87 scales = ["hundred", "thousand", "million", "billion", "trillion"]88 if not numwords:89 for idx, word in enumerate(ONES):90 numwords[word] = idx91 for idx, word in enumerate(tens):92 numwords[word] = idx * 1093 for idx, word in enumerate(scales):94 numwords[word] = 10 ** (idx * 3 or 2)95 if in_word in numwords:96 # print(in_word)97 # print(numwords[in_word])98 return numwords[in_word]99 try:100 int(in_word)101 return int(in_word)102 except ValueError:103 pass104 # Spell correction using find_char_diff105 char_diffs = [find_char_diff(in_word, i) for i in ONES + tens + scales]106 min_char_diff = min(char_diffs)107 if min_char_diff <= SPELL_CORRECT_MIN_CHAR_DIFF:108 return char_diffs.index(min_char_diff)109 110 111def tokens2int(tokens):112 # Takes a list of tokens and returns a int representation of them113 types = []114 for i in tokens:115 if i <= 9:116 types.append(1)117 118 elif i <= 90:119 types.append(2)120 121 else:122 types.append(3)123 # print(tokens)124 if len(tokens) <= 3:125 current = 0126 for i, number in enumerate(tokens):127 if i != 0 and types[i] < types[i - 1] and current != tokens[i - 1] and types[i - 1] != 3:128 current += tokens[i] + tokens[i - 1]129 elif current <= tokens[i] and current != 0:130 current *= tokens[i]131 elif 3 not in types and 1 not in types:132 current = int(''.join(str(i) for i in tokens))133 break134 elif '111' in ''.join(str(i) for i in types) and 2 not in types and 3 not in types:135 current = int(''.join(str(i) for i in tokens))136 break137 else:138 current += number139 140 elif 3 not in types and 2 not in types:141 current = int(''.join(str(i) for i in tokens))142 143 else:144 """145 double_list = []146 current_double = []147 double_type_list = []148 for i in tokens:149 if len(current_double) < 2:150 current_double.append(i)151 else:152 double_list.append(current_double)153 current_double = []154 current_double = []155 for i in types:156 if len(current_double) < 2:157 current_double.append(i)158 else:159 double_type_list.append(current_double)160 current_double = []161 print(double_type_list)162 print(double_list)163 current = 0164 for i, type_double in enumerate(double_type_list):165 if len(type_double) == 1:166 current += double_list[i][0]167 elif type_double[0] == type_double[1]:168 current += int(str(double_list[i][0]) + str(double_list[i][1]))169 elif type_double[0] > type_double[1]:170 current += sum(double_list[i])171 elif type_double[0] < type_double[1]:172 current += double_list[i][0] * double_list[i][1]173 # print(current)174 """175 count = 0176 current = 0177 for i, token in enumerate(tokens):178 count += 1179 if count == 2:180 if types[i - 1] == types[i]:181 current += int(str(token) + str(tokens[i - 1]))182 elif types[i - 1] > types[i]:183 current += tokens[i - 1] + token184 else:185 current += tokens[i - 1] * token186 count = 0187 elif i == len(tokens) - 1:188 current += token189 190 return current191 192 193def text2int(text):194 # Wraps all of the functions up into one195 return tokens2int(tokenize(text))196 197###############################################198# Vish editdistance approach doesn't halt199 200 201def lev_dist(a, b):202 '''203 This function will calculate the levenshtein distance between two input204 strings a and b205 206 params:207 a (String) : The first string you want to compare208 b (String) : The second string you want to compare209 210 returns:211 This function will return the distance between string a and b.212 213 example:214 a = 'stamp'215 b = 'stomp'216 lev_dist(a,b)217 >> 1.0218 '''219 if not isinstance(a, str) and isinstance(b, str):220 raise ValueError(f"lev_dist() requires 2 strings not lev_dist({repr(a)}, {repr(b)}")221 if a == b:222 return 0223 224 def min_dist(s1, s2):225 226 print(f"{a[s1]}s1{b[s2]}s2 ", end='')227 if s1 >= len(a) or s2 >= len(b):228 return len(a) - s1 + len(b) - s2229 230 # no change required231 if a[s1] == b[s2]:232 return min_dist(s1 + 1, s2 + 1)233 234 return 1 + min(235 min_dist(s1, s2 + 1), # insert character236 min_dist(s1 + 1, s2), # delete character237 min_dist(s1 + 1, s2 + 1), # replace character238 )239 240 dist = min_dist(0, 0)241 print(f"\n lev_dist({a}, {b}) => {dist}")242 return dist243 244 245def correct_number_text(text):246 """ Convert an English str containing number words with possibly incorrect spellings into an int247 248 >>> robust_text2int("too")249 2250 >>> robust_text2int("fore")251 4252 >>> robust_text2int("1 2 tree")253 123254 """255 words = {256 "zero": 0,257 "one": 1,258 "two": 2,259 "three": 3,260 "four": 4,261 "five": 5,262 "six": 6,263 "seven": 7,264 "eight": 8,265 "nine": 9,266 "ten": 10,267 "eleven": 11,268 "twelve": 12,269 "thirteen": 13,270 "fourteen": 14,271 "fifteen": 15,272 "sixteen": 16,273 "seventeen": 17,274 "eighteen": 18,275 "nineteen": 19,276 "score": 20,277 "twenty": 20,278 "thirty": 30,279 "forty": 40,280 "fifty": 50,281 "sixty": 60,282 "seventy": 70,283 "eighty": 80,284 "ninety": 90,285 "hundred": 100,286 "thousand": 1000,287 "million": 1000000,288 "billion": 1000000000,289 }290 291 text = text.lower()292 text_words = text.replace("-", " ").split()293 corrected_words = []294 for text_word in text_words:295 if text_word not in words:296 print(f"{text_word} not in words")297 if not isinstance(text_word, str):298 return TOKENS2INT_ERROR_INT299 t0 = time.time()300 min_dist = len(text_word)301 correct_spelling = None302 for word in words:303 dist = edit_dist(word, text_word)304 if dist < min_dist:305 correct_spelling = word306 min_dist = dist307 corrected_words.append(correct_spelling)308 t1 = time.time()309 print(f"{text_word} dt:{t1-t0}")310 else:311 corrected_words.append(text_word)312 313 corrected_text = " ".join(corrected_words)314 315 print(corrected_text)316 return corrected_text317 318 # From hereon, we can use text2int319 # TODO320 321 322sentiment = pipeline(task="sentiment-analysis", model="distilbert-base-uncased-finetuned-sst-2-english")323 324 325def get_sentiment(text):326 return sentiment(text)327 328 329def robust_text2int(text):330 """ Correct spelling of number words in text before using text2int """331 try:332 return tokens2int(tokenize(correct_number_text(text)))333 except Exception as e:334 print(e)335 return TOKENS2INT_ERROR_INT336 