Spaces:
Running on Zero
Running on Zero
File size: 115,123 Bytes
1eb9bad 28fa57e 1eb9bad af8366e 3910264 796164b 3910264 796164b 3910264 b25e4d0 3910264 796164b 3910264 796164b 3910264 796164b b25e4d0 3910264 796164b 3910264 b25e4d0 3910264 b25e4d0 796164b b25e4d0 796164b b25e4d0 796164b b25e4d0 3910264 796164b b25e4d0 3910264 6a1cf05 cbdaba0 6a1cf05 cbdaba0 89c7c26 6a1cf05 1eb9bad 97cfc6b 1eb9bad 7d21a61 1eb9bad af8366e 3910264 6a1cf05 1eb9bad 97cfc6b 7d21a61 97cfc6b 7d21a61 97cfc6b 7d21a61 97cfc6b 1eb9bad 97cfc6b 1eb9bad 7d21a61 1eb9bad 28fa57e 1eb9bad 28fa57e 1eb9bad 28fa57e 1eb9bad 993d53c 1eb9bad 993d53c 1eb9bad 993d53c 1eb9bad af8366e 3910264 6a1cf05 1eb9bad af8366e 3910264 6a1cf05 1eb9bad 89c7c26 1eb9bad af8366e 3910264 1eb9bad 89c7c26 1eb9bad af8366e 3910264 1eb9bad af8366e 3910264 6a1cf05 1eb9bad af8366e 3910264 6a1cf05 1eb9bad af8366e 3910264 6a1cf05 1eb9bad 6a1cf05 1eb9bad 5a66063 1eb9bad 5a66063 1eb9bad 5a66063 1eb9bad 5a66063 1eb9bad 5a66063 1eb9bad 5a66063 1eb9bad 5a66063 1eb9bad 5a66063 1eb9bad 5a66063 89c7c26 1eb9bad af8366e 3910264 1eb9bad 5a66063 1eb9bad | 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 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 1033 1034 1035 1036 1037 1038 1039 1040 1041 1042 1043 1044 1045 1046 1047 1048 1049 1050 1051 1052 1053 1054 1055 1056 1057 1058 1059 1060 1061 1062 1063 1064 1065 1066 1067 1068 1069 1070 1071 1072 1073 1074 1075 1076 1077 1078 1079 1080 1081 1082 1083 1084 1085 1086 1087 1088 1089 1090 1091 1092 1093 1094 1095 1096 1097 1098 1099 1100 1101 1102 1103 1104 1105 1106 1107 1108 1109 1110 1111 1112 1113 1114 1115 1116 1117 1118 1119 1120 1121 1122 1123 1124 1125 1126 1127 1128 1129 1130 1131 1132 1133 1134 1135 1136 1137 1138 1139 1140 1141 1142 1143 1144 1145 1146 1147 1148 1149 1150 1151 1152 1153 1154 1155 1156 1157 1158 1159 1160 1161 1162 1163 1164 1165 1166 1167 1168 1169 1170 1171 1172 1173 1174 1175 1176 1177 1178 1179 1180 1181 1182 1183 1184 1185 1186 1187 1188 1189 1190 1191 1192 1193 1194 1195 1196 1197 1198 1199 1200 1201 1202 1203 1204 1205 1206 1207 1208 1209 1210 1211 1212 1213 1214 1215 1216 1217 1218 1219 1220 1221 1222 1223 1224 1225 1226 1227 1228 1229 1230 1231 1232 1233 1234 1235 1236 1237 1238 1239 1240 1241 1242 1243 1244 1245 1246 1247 1248 1249 1250 1251 1252 1253 1254 1255 1256 1257 1258 1259 1260 1261 1262 1263 1264 1265 1266 1267 1268 1269 1270 1271 1272 1273 1274 1275 1276 1277 1278 1279 1280 1281 1282 1283 1284 1285 1286 1287 1288 1289 1290 1291 1292 1293 1294 1295 1296 1297 1298 1299 1300 1301 1302 1303 1304 1305 1306 1307 1308 1309 1310 1311 1312 1313 1314 1315 1316 1317 1318 1319 1320 1321 1322 1323 1324 1325 1326 1327 1328 1329 1330 1331 1332 1333 1334 1335 1336 1337 1338 1339 1340 1341 1342 1343 1344 1345 1346 1347 1348 1349 1350 1351 1352 1353 1354 1355 1356 1357 1358 1359 1360 1361 1362 1363 1364 1365 1366 1367 1368 1369 1370 1371 1372 1373 1374 1375 1376 1377 1378 1379 1380 1381 1382 1383 1384 1385 1386 1387 1388 1389 1390 1391 1392 1393 1394 1395 1396 1397 1398 1399 1400 1401 1402 1403 1404 1405 1406 1407 1408 1409 1410 1411 1412 1413 1414 1415 1416 1417 1418 1419 1420 1421 1422 1423 1424 1425 1426 1427 1428 1429 1430 1431 1432 1433 1434 1435 1436 1437 1438 1439 1440 1441 1442 1443 1444 1445 1446 1447 1448 1449 1450 1451 1452 1453 1454 1455 1456 1457 1458 1459 1460 1461 1462 1463 1464 1465 1466 1467 1468 1469 1470 1471 1472 1473 1474 1475 1476 1477 1478 1479 1480 1481 1482 1483 1484 1485 1486 1487 1488 1489 1490 1491 1492 1493 1494 1495 1496 1497 1498 1499 1500 1501 1502 1503 1504 1505 1506 1507 1508 1509 1510 1511 1512 1513 1514 1515 1516 1517 1518 1519 1520 1521 1522 1523 1524 1525 1526 1527 1528 1529 1530 1531 1532 1533 1534 1535 1536 1537 1538 1539 1540 1541 1542 1543 1544 1545 1546 1547 1548 1549 1550 1551 1552 1553 1554 1555 1556 1557 1558 1559 1560 1561 1562 1563 1564 1565 1566 1567 1568 1569 1570 1571 1572 1573 1574 1575 1576 1577 1578 1579 1580 1581 1582 1583 1584 1585 1586 1587 1588 1589 1590 1591 1592 1593 1594 1595 1596 1597 1598 1599 1600 1601 1602 1603 1604 1605 1606 1607 1608 1609 1610 1611 1612 1613 1614 1615 1616 1617 1618 1619 1620 1621 1622 1623 1624 1625 1626 1627 1628 1629 1630 1631 1632 1633 1634 1635 1636 1637 1638 1639 1640 1641 1642 1643 1644 1645 1646 1647 1648 1649 1650 1651 1652 1653 1654 1655 1656 1657 1658 1659 1660 1661 1662 1663 1664 1665 1666 1667 1668 1669 1670 1671 1672 1673 1674 1675 1676 1677 1678 1679 1680 1681 1682 1683 1684 1685 1686 1687 1688 1689 1690 1691 1692 1693 1694 1695 1696 1697 1698 1699 1700 1701 1702 1703 1704 1705 1706 1707 1708 1709 1710 1711 1712 1713 1714 1715 1716 1717 1718 1719 1720 1721 1722 1723 1724 1725 1726 1727 1728 1729 1730 1731 1732 1733 1734 1735 1736 1737 1738 1739 1740 1741 1742 1743 1744 1745 1746 1747 1748 1749 1750 1751 1752 1753 1754 1755 1756 1757 1758 1759 1760 1761 1762 1763 1764 1765 1766 1767 1768 1769 1770 1771 1772 1773 1774 1775 1776 1777 1778 1779 1780 1781 1782 1783 1784 1785 1786 1787 1788 1789 1790 1791 1792 1793 1794 1795 1796 1797 1798 1799 1800 1801 1802 1803 1804 1805 1806 1807 1808 1809 1810 1811 1812 1813 1814 1815 1816 1817 1818 1819 1820 1821 1822 1823 1824 1825 1826 1827 1828 1829 1830 1831 1832 1833 1834 1835 1836 1837 1838 1839 1840 1841 1842 1843 1844 1845 1846 1847 1848 1849 1850 1851 1852 1853 1854 1855 1856 1857 1858 1859 1860 1861 1862 1863 1864 1865 1866 1867 1868 1869 1870 1871 1872 1873 1874 1875 1876 1877 1878 1879 1880 1881 1882 1883 1884 1885 1886 1887 1888 1889 1890 1891 1892 1893 1894 1895 1896 1897 1898 1899 1900 1901 1902 1903 1904 1905 1906 1907 1908 1909 1910 1911 1912 1913 1914 1915 1916 1917 1918 1919 1920 1921 1922 1923 1924 1925 1926 1927 1928 1929 1930 1931 1932 1933 1934 1935 1936 1937 1938 1939 1940 1941 1942 1943 1944 1945 1946 1947 1948 1949 1950 1951 1952 1953 1954 1955 1956 1957 1958 1959 1960 1961 1962 1963 1964 1965 1966 1967 1968 1969 1970 1971 1972 1973 1974 1975 1976 1977 1978 1979 1980 1981 1982 1983 1984 1985 1986 1987 1988 1989 1990 1991 1992 1993 1994 1995 1996 1997 1998 1999 2000 2001 2002 2003 2004 2005 2006 2007 2008 2009 2010 2011 2012 2013 2014 2015 2016 2017 2018 2019 2020 2021 2022 2023 2024 2025 2026 2027 2028 2029 2030 2031 2032 2033 2034 2035 2036 2037 2038 2039 2040 2041 2042 2043 2044 2045 2046 2047 2048 2049 2050 2051 2052 2053 2054 2055 2056 2057 2058 2059 2060 2061 2062 2063 2064 2065 2066 2067 2068 2069 2070 2071 2072 2073 2074 2075 2076 2077 2078 2079 2080 2081 2082 2083 2084 2085 2086 2087 2088 2089 2090 2091 2092 2093 2094 2095 2096 2097 2098 2099 2100 2101 2102 2103 2104 2105 2106 2107 2108 2109 2110 2111 2112 2113 2114 2115 2116 2117 2118 2119 2120 2121 2122 2123 2124 2125 2126 2127 2128 2129 2130 2131 2132 2133 2134 2135 2136 2137 2138 2139 2140 2141 2142 2143 2144 2145 2146 2147 2148 2149 2150 2151 2152 2153 2154 2155 2156 2157 2158 2159 2160 2161 2162 2163 2164 2165 2166 2167 2168 2169 2170 2171 2172 2173 2174 2175 2176 2177 2178 2179 2180 2181 2182 2183 2184 2185 2186 2187 2188 2189 2190 2191 2192 2193 2194 2195 2196 2197 2198 2199 2200 2201 2202 2203 2204 2205 2206 2207 2208 2209 2210 2211 2212 2213 2214 2215 2216 2217 2218 2219 2220 2221 2222 2223 2224 2225 2226 2227 2228 2229 2230 2231 2232 2233 2234 2235 2236 2237 2238 2239 2240 2241 2242 2243 2244 2245 2246 2247 2248 2249 2250 2251 2252 2253 2254 2255 2256 2257 2258 2259 2260 2261 2262 2263 2264 2265 2266 2267 2268 2269 2270 2271 2272 2273 2274 2275 2276 2277 2278 2279 2280 2281 2282 2283 2284 2285 2286 2287 2288 2289 2290 2291 2292 2293 2294 2295 2296 2297 2298 2299 2300 2301 2302 2303 2304 2305 2306 2307 2308 2309 2310 2311 2312 2313 2314 2315 2316 2317 2318 2319 2320 2321 2322 2323 2324 2325 2326 2327 2328 2329 2330 2331 2332 2333 2334 2335 2336 2337 2338 2339 2340 2341 2342 2343 2344 2345 2346 2347 2348 2349 2350 2351 2352 2353 2354 2355 2356 2357 2358 2359 2360 2361 2362 2363 2364 2365 2366 2367 2368 2369 2370 2371 2372 2373 2374 2375 2376 2377 2378 2379 2380 2381 2382 2383 2384 2385 2386 2387 2388 2389 2390 2391 2392 2393 2394 2395 2396 2397 2398 2399 2400 2401 2402 2403 2404 2405 2406 2407 2408 2409 2410 2411 2412 2413 2414 2415 2416 2417 2418 2419 2420 2421 2422 2423 2424 2425 2426 2427 2428 2429 2430 2431 2432 2433 2434 2435 2436 2437 2438 2439 2440 2441 2442 2443 2444 2445 2446 2447 2448 2449 2450 2451 2452 2453 2454 2455 2456 2457 2458 2459 2460 2461 2462 2463 2464 2465 2466 2467 2468 2469 2470 2471 2472 2473 2474 2475 2476 2477 2478 2479 2480 2481 2482 2483 2484 2485 2486 2487 2488 2489 2490 2491 2492 2493 2494 2495 2496 2497 2498 2499 2500 2501 2502 2503 2504 2505 2506 2507 2508 2509 2510 2511 2512 2513 2514 2515 2516 2517 2518 2519 2520 2521 2522 2523 2524 2525 2526 2527 2528 2529 2530 2531 2532 2533 2534 2535 2536 2537 2538 2539 2540 2541 2542 2543 2544 2545 2546 2547 2548 2549 2550 2551 2552 2553 2554 2555 2556 2557 2558 2559 2560 2561 2562 2563 2564 2565 2566 2567 2568 2569 2570 2571 2572 2573 2574 2575 2576 2577 2578 2579 2580 2581 2582 2583 2584 2585 2586 2587 2588 2589 2590 | #!/usr/bin/env python3
"""
Inference script for MHG text normalization.
Apply trained model to new, unannotated MHG texts.
Supports sliding-window stride inference for long texts that exceed
the model's maximum sequence length. When stride is enabled, long
inputs are split into overlapping windows, each normalized independently,
and predictions are merged preferring window centers (where context is
richest).
"""
import os
import yaml
import torch
import argparse
from pathlib import Path
from typing import List, Optional, Tuple
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
import re
import difflib
from tqdm import tqdm
# Editorial punctuation characters to strip before normalization and restore after.
# These appear in edited/printed text but were NOT present in the training data.
# MHG abbreviation markers (hyphens '-', apostrophes '\u2019') are NOT included
# here because they are meaningful for the normalization task.
_PUNCT_CHARS = frozenset(".,;:!?()[]«»‹›\u201c\u201d\u201e\"\u201a\u2018\u2019\u2014\u2013<>")
def _attach_en_proclitic(text: str) -> str:
"""Attach free-standing negation proclitic 'en' to the following word.
In Middle High German, the negation proclitic 'en' appears as a separate
token before the word it negates. This post-processing step joins them
directly without a separator, e.g. 'en guot' → 'enguot'.
Capitalization of the proclitic is preserved: 'En guot' → 'Enguot'.
Only attaches when 'en' appears as a standalone, case-insensitive word
(bounded by whitespace or start/end of string) and is followed by another
word. Punctuation after 'en' prevents attachment.
Args:
text: Input text to process
Returns:
Text with 'en' proclitics attached to the following word
"""
# Match standalone "en" (case-insensitive) as a whole word,
# followed by whitespace and a non-whitespace character.
pattern = re.compile(r'(?<!\S)(en)(?!\S)\s+(\S)', re.IGNORECASE)
def _replace(m):
en_word = m.group(1)
next_char = m.group(2)
return f"{en_word}{next_char}"
return pattern.sub(_replace, text)
def _lenition_t_after_n_before_vowel(text: str) -> str:
"""MHG Lenition: /t/ → /d/ after /n/ and before a vowel.
In Middle High German, the voiceless stop /t/ becomes voiced /d/ when
it appears between /n/ and a vowel (within the same word). This covers
all MHG vowels including those with diacritics (ä, ö, ü, ê, î, etc.)
and base vowels followed by combining diacritical marks.
Capitalization is preserved: 'Ntô' → 'Ndô', 'nTô' → 'ndô'.
Args:
text: Input text to process
Returns:
Text with lenited t → d where applicable
"""
# Vowel set: MHG base vowels plus all common diacritic variants
vowel_chars = 'aeiouyAEIOUY'
vowel_chars += 'äöüÄÖÜ' # umlaut
vowel_chars += 'âêîôûÂÊÎÔÛ' # circumflex
vowel_chars += 'āēīōūĀĒĪŌŪ' # macron
vowel_chars += 'åæœøÅÆŒØ' # other
# Match 'n' + 't' followed by a vowel (optionally with combining diacritics)
pattern = re.compile(
r'([nN])([tT])(?=[' + re.escape(vowel_chars) + r'][\u0300-\u036f]*)'
)
def _replace(m):
n_char = m.group(1)
t_char = m.group(2)
return n_char + ('d' if t_char == 't' else 'D')
return pattern.sub(_replace, text)
def _lenition_t_after_l_before_vowel(text: str) -> str:
"""MHG Lenition: /t/ → /d/ after /l/ and before a vowel.
In Middle High German, the voiceless stop /t/ becomes voiced /d/ when
it appears between /l/ and a vowel (within the same word). This covers
all MHG vowels including those with diacritics (ä, ö, ü, ê, î, etc.)
and base vowels followed by combining diacritical marks.
Capitalization is preserved: 'Ltô' → 'Ldô', 'lTô' → 'ldô'.
Args:
text: Input text to process
Returns:
Text with lenited t → d where applicable
"""
# Vowel set: MHG base vowels plus all common diacritic variants
vowel_chars = 'aeiouyAEIOUY'
vowel_chars += 'äöüÄÖÜ' # umlaut
vowel_chars += 'âêîôûÂÊÎÔÛ' # circumflex
vowel_chars += 'āēīōūĀĒĪŌŪ' # macron
vowel_chars += 'åæœøÅÆŒØ' # other
# Match 'l' + 't' followed by a vowel (optionally with combining diacritics)
pattern = re.compile(
r'([lL])([tT])(?=[' + re.escape(vowel_chars) + r'][\u0300-\u036f]*)'
)
def _replace(m):
l_char = m.group(1)
t_char = m.group(2)
return l_char + ('d' if t_char == 't' else 'D')
return pattern.sub(_replace, text)
def _niet_to_niht(text: str) -> str:
"""Change freestanding 'niet'/'niut' to 'niht'.
In some MHG texts, the negation word 'niet' (or variant 'niut') appears
as a freestanding token but the preferred normalized form is 'niht'.
This post-processing step replaces standalone occurrences of 'niet' and
'niut' (case-insensitive) with 'niht', preserving the original
capitalization pattern.
Handles words immediately followed by editorial punctuation (from
_PUNCT_CHARS) by capturing and re-attaching the trailing punctuation.
Examples:
'niet' → 'niht'
'niut' → 'niht'
'Niet' → 'Niht'
'Niut' → 'Niht'
'NIET' → 'NIHT'
'NIUT' → 'NIHT'
'niet.' → 'niht.'
'niet,' → 'niht,'
'niet.)' → 'niht.)'
Args:
text: Input text to process
Returns:
Text with freestanding 'niet'/'niut' changed to 'niht'
"""
# Build character class of all editorial punctuation characters so we can
# explicitly match trailing punctuation and re-attach it after replacement.
# This is necessary because post-processing runs AFTER punctuation restoration,
# so words like "niet." have punctuation attached.
punct_class = re.escape(''.join(_PUNCT_CHARS))
# Match standalone niet/niut (bounded by non-word on left) followed by
# zero or more punctuation characters and then a non-word boundary.
pattern = re.compile(
r'(?<!\S)(ni[eu]t)([' + punct_class + r']*)(?!\S)', re.IGNORECASE
)
def _replace(m):
word = m.group(1)
trailing = m.group(2)
if word in ('niet', 'niut'):
result = 'niht'
elif word in ('Niet', 'Niut'):
result = 'Niht'
elif word in ('NIET', 'NIUT'):
result = 'NIHT'
else:
# Mixed case — preserve first letter case
result = (word[0].upper() if word[0].isupper() else word[0].lower()) + 'iht'
return result + trailing
return pattern.sub(_replace, text)
def _common_apocopes(text: str) -> str:
"""Apply common MHG apocope (final vowel/consonant loss) rules.
In Middle High German, certain words lose their final syllables in
common usage. This post-processing step handles the most frequent
apocopes:
Examples:
'wile' → 'wil'
'vore' → 'vor'
Args:
text: Input text to process
Returns:
Text with common apocopes applied
"""
# Build character class of all editorial punctuation so we can
# capture trailing punctuation and re-attach it after replacement.
punct_class = re.escape(''.join(_PUNCT_CHARS))
replacements = {
'vile': 'vil',
'vore': 'vor',
'wile': 'wil',
'wole': 'wol',
}
# Build a single regex that matches any of the keys as whole words
# (bounded by whitespace or start/end), followed by optional punctuation.
pattern_str = r'(?<!\S)(' + '|'.join(re.escape(k) for k in replacements.keys()) + r')([' + punct_class + r']*)(?!\S)'
pattern = re.compile(pattern_str, re.IGNORECASE)
def _replace(m):
word = m.group(1)
trailing = m.group(2)
result = replacements.get(word.lower(), word)
# Preserve capitalization pattern from the original word
if word[0].isupper() and result[0].islower():
result = result[0].upper() + result[1:]
return result + trailing
return pattern.sub(_replace, text)
def _capitalize_first_alpha(word: str) -> str:
"""Capitalize the first alphabetic character in a word.
Handles words that may start with non-alphabetic characters
(e.g., MHG abbreviation markers like hyphens or apostrophes,
or leading punctuation that was not stripped).
Args:
word: Word to capitalize
Returns:
Word with first alphabetic character uppercased
"""
for i, ch in enumerate(word):
if ch.isalpha():
return word[:i] + ch.upper() + word[i+1:]
return word # No alphabetic characters found; return unchanged
def _select_device() -> torch.device:
"""Select the best available device: CUDA > ROCm > DirectML > MPS > CPU."""
if torch.cuda.is_available():
return torch.device('cuda')
if hasattr(torch.backends, 'mps') and torch.backends.mps.is_available():
return torch.device('mps')
try:
import torch_directml
device = torch_directml.device()
return torch.device(device)
except ImportError:
pass
return torch.device('cpu')
class MHGNormalizer:
"""Normalizer for Middle High German text."""
def __init__(self, model_path: str, config_path: Optional[str] = None, defer_gpu: bool = False):
"""
Initialize normalizer with trained model.
Args:
model_path: Path to trained model checkpoint
config_path: Path to configuration file (optional)
"""
self.device = _select_device()
print(f"Using device: {self.device}")
# Load configuration if provided
if config_path:
with open(config_path, 'r', encoding='utf-8') as f:
self.config = yaml.safe_load(f)
else:
# Default configuration
self.config = {
'inference': {
'max_length': 512,
'num_beams': 5,
'length_penalty': 0.6,
'do_sample': False,
'top_p': 0.9,
'temperature': 1.0,
'stride': True,
'stride_window_tokens': 400,
'stride_step_tokens': 300,
}
}
# Ensure stride config exists (for configs created before stride was added)
inference_config = self.config.setdefault('inference', {})
if 'stride' not in inference_config:
inference_config['stride'] = True
if 'stride_window_tokens' not in inference_config:
inference_config['stride_window_tokens'] = 400
if 'stride_step_tokens' not in inference_config:
inference_config['stride_step_tokens'] = 300
# Ensure preserve_punctuation and preserve_capitalization defaults
if 'preserve_punctuation' not in inference_config:
inference_config['preserve_punctuation'] = True
if 'preserve_capitalization' not in inference_config:
inference_config['preserve_capitalization'] = True
# Ensure post_processing defaults
post_config = self.config.setdefault('post_processing', {})
if 'attach_en_proclitic' not in post_config:
post_config['attach_en_proclitic'] = False
if 'lenition_t_after_n' not in post_config:
post_config['lenition_t_after_n'] = False
if 'lenition_t_after_l' not in post_config:
post_config['lenition_t_after_l'] = False
if 'niet_to_niht' not in post_config:
post_config['niet_to_niht'] = False
if 'common_apocopes' not in post_config:
post_config['common_apocopes'] = False
# Load model and tokenizer
print(f"\nLoading model from: {model_path}")
self.tokenizer = AutoTokenizer.from_pretrained(model_path)
# Detect dtype from model config (FP16 models have torch_dtype=float16)
import json as _json
from pathlib import Path as _Path
_cfg = _Path(model_path) / "config.json"
if _cfg.exists():
_cfg_data = _json.loads(_cfg.read_text())
_model_dtype = torch.float16 if _cfg_data.get("torch_dtype") == "float16" else None
else:
_model_dtype = None
# CPU lacks efficient FP16 support — load as FP32 there. On CUDA/ZeroGPU
# load FP16 weights as bfloat16, which is numerically stable for T5/mT5
# (FP16 tends to overflow/NaN in T5) and fast on modern GPUs.
# When defer_gpu=True (ZeroGPU), the model is preloaded to CPU at
# startup and moved to GPU later inside @spaces.GPU — load as BF16
# so the GPU transfer is fast and inference runs in BF16.
if _model_dtype == torch.float16:
if str(self.device) == "cpu" and not defer_gpu:
print(" CPU detected — loading FP16 model as FP32 (CPU has no FP16 acceleration)")
_load_dtype = torch.float32
else:
_where = "GPU" if str(self.device) != "cpu" else "CPU (GPU deferred)"
print(f" Loading FP16 model as bfloat16 ({_where})")
_load_dtype = torch.bfloat16
else:
_load_dtype = _model_dtype
# Do NOT pass tie_word_embeddings=False — T5/mT5 always forces
# tie_word_embeddings=True in T5Config.__post_init__. Passing False
# prevents HuggingFace from tying encoder/decoder embed_tokens to
# shared.weight when loading checkpoints, causing broken inference.
self.model = AutoModelForSeq2SeqLM.from_pretrained(
model_path,
torch_dtype=_load_dtype,
low_cpu_mem_usage=True,
)
self.model.to(self.device)
self.model.eval()
print("Model loaded successfully!\n")
def _ensure_on_device(self):
"""Move the model to the best available device if it has changed.
On ZeroGPU, __init__ runs at startup when no GPU is visible, so the
model loads to CPU. This is called inside @spaces.GPU (where CUDA is
available) to transfer the model to the GPU. Idempotent: if the model
is already on the target device, nothing happens.
"""
target = _select_device()
if target != self.device:
print(f"Moving model from {self.device} to {target} …")
self.device = target
self.model.to(self.device)
@staticmethod
def _extract_punctuation(text: str):
"""Extract editorial punctuation from text for preservation during normalization.
Strips leading and trailing editorial punctuation characters from each word,
returning the stripped text and a punctuation map for later restoration.
Words that consist entirely of punctuation characters are absorbed into
the trailing punctuation of the previous word (or the leading punctuation
of the next word if they appear at the start of the text).
Args:
text: Input text that may contain editorial punctuation
Returns:
tuple: (stripped_text, punct_map) where:
- stripped_text: text with editorial punctuation removed from word edges
- punct_map: list of (leading_punct, trailing_punct) tuples, one per
word in stripped_text
"""
words = text.split()
punct_map = []
stripped_words = []
pending_leading = ''
for word in words:
# Extract leading punctuation
leading = ''
i = 0
while i < len(word) and word[i] in _PUNCT_CHARS:
leading += word[i]
i += 1
# Extract trailing punctuation
trailing = ''
j = len(word) - 1
while j >= i and word[j] in _PUNCT_CHARS:
trailing = word[j] + trailing
j -= 1
core = word[i:j+1]
if core:
# Prepend any pending leading punctuation (from standalone punct at start)
leading = pending_leading + leading
pending_leading = ''
punct_map.append((leading, trailing))
stripped_words.append(core)
else:
# Word is entirely punctuation (e.g., "..." or "—")
# Absorb into adjacent word's punctuation. Since ' '.join() provides
# only one space between words, we include an explicit space to
# preserve the whitespace around standalone punctuation marks.
if punct_map:
# Append to previous word's trailing punctuation with leading space
prev_leading, prev_trailing = punct_map[-1]
punct_map[-1] = (prev_leading, prev_trailing + ' ' + leading + trailing)
else:
# At start of text — save as leading punctuation for next word
pending_leading += leading + trailing
# If text ended with standalone punctuation, attach to last word's trailing
if pending_leading and punct_map:
prev_leading, prev_trailing = punct_map[-1]
punct_map[-1] = (prev_leading, prev_trailing + pending_leading)
stripped_text = ' '.join(stripped_words)
return stripped_text, punct_map
@staticmethod
def _restore_punctuation(text: str, punct_map: list) -> str:
"""Restore editorial punctuation to normalized text.
Uses position-based mapping: word i gets punct_map[i]. When the
model produces a different number of words than the original,
trailing punctuation from unmatched positions is appended to the
last word.
For across-lines mode where word counts may change significantly,
use _restore_punctuation_aligned() instead, which uses character-level
alignment to correctly map punctuation even when words split or merge.
Args:
text: Normalized text (without editorial punctuation)
punct_map: List of (leading_punct, trailing_punct) tuples from
_extract_punctuation
Returns:
Text with editorial punctuation restored
"""
words = text.split()
result = []
for i, word in enumerate(words):
if i < len(punct_map):
leading, trailing = punct_map[i]
result.append(leading + word + trailing)
else:
# More words in output than in punct_map; no punctuation info
result.append(word)
# If model produced fewer words than original, collect trailing punctuation
# from unmatched positions and append to the last output word.
# This preserves sentence-final punctuation (e.g., trailing periods).
if len(words) > 0 and len(words) < len(punct_map):
extra_trailing = ''.join(t for _, t in punct_map[len(words):])
if extra_trailing:
result[-1] = result[-1] + extra_trailing
return ' '.join(result)
@staticmethod
def _restore_punctuation_aligned(original_text: str, normalized_text: str,
punct_map: list) -> str:
"""Restore editorial punctuation using character-level alignment.
When the model changes word counts (e.g., "inalten" → "in alten"),
position-based _restore_punctuation() fails because word indices shift.
This method uses character-level sequence alignment to correctly map
each normalized word to its corresponding original word, then applies
the original word's punctuation.
For example, if "inalten" (with no punctuation) splits into "in alten",
character alignment maps "in" to the start of "inalten" and "alten" to
the end, so both correctly inherit the punctuation of "inalten".
Args:
original_text: Original text with punctuation already stripped
normalized_text: Normalized text (without editorial punctuation)
punct_map: List of (leading_punct, trailing_punct) tuples,
one per word in original_text
Returns:
Text with editorial punctuation restored
"""
normalized_words = normalized_text.split()
if not normalized_words:
return normalized_text
original_words = original_text.split()
if not original_words:
return normalized_text
# If word counts match, use simple position-based restoration
if len(normalized_words) == len(original_words):
return MHGNormalizer._restore_punctuation(normalized_text, punct_map)
# Use character-level alignment to map each normalized word to its
# corresponding original word index. This correctly handles word
# splits like "inalten" → "in alten" because character alignment
# maps "in" to the "in" prefix of "inalten".
norm_to_orig_word = MHGNormalizer._align_normalized_to_original_charlevel(
original_text, normalized_text
)
# Apply punctuation, handling word splits:
# When multiple normalized words map to the same original word,
# only the first gets leading punct and only the last gets trailing punct.
result = []
for i, word in enumerate(normalized_words):
orig_idx = norm_to_orig_word[i]
if orig_idx < len(punct_map):
leading, trailing = punct_map[orig_idx]
else:
leading, trailing = '', ''
# Check if this is the first/last normalized word mapping
# to this original word (handles word splits)
is_first = (i == 0 or norm_to_orig_word[i - 1] != orig_idx)
is_last = (i == len(normalized_words) - 1 or
norm_to_orig_word[i + 1] != orig_idx)
applied_leading = leading if is_first else ''
applied_trailing = trailing if is_last else ''
result.append(applied_leading + word + applied_trailing)
# Collect trailing punctuation from unmatched original words
# (original words that were deleted by the model)
matched_orig = set(norm_to_orig_word)
extra_trailing = ''
for orig_idx in range(len(punct_map)):
if orig_idx not in matched_orig:
_, trailing = punct_map[orig_idx]
extra_trailing += trailing
if extra_trailing and result:
result[-1] = result[-1] + extra_trailing
return ' '.join(result)
@staticmethod
def _extract_capitalization(text: str):
"""Extract capitalization pattern from text for preservation during normalization.
Records which words start with an uppercase letter, then lowercases
the text so the model receives consistent lowercase input.
Words whose first alphabetic character is uppercase are marked in the
cap_map. Non-alphabetic leading characters (e.g., punctuation that
wasn't stripped, MHG abbreviation markers) are skipped when finding
the first alphabetic character.
This method should be called AFTER _extract_punctuation so that
the cap_map indices align with the punct_map indices (both indexed
by words in the punctuation-stripped text).
Args:
text: Input text (typically with punctuation already stripped)
Returns:
tuple: (lowercased_text, cap_map) where:
- lowercased_text: the text converted to lowercase
- cap_map: list of bools, one per word; True if the word's
first alphabetic character was uppercase
"""
words = text.split()
cap_map = []
for word in words:
has_cap = False
for ch in word:
if ch.isalpha():
has_cap = ch.isupper()
break
cap_map.append(has_cap)
lowered_text = text.lower()
return lowered_text, cap_map
@staticmethod
def _restore_capitalization(text: str, cap_map: list) -> str:
"""Restore capitalization to normalized text using position-based mapping.
For each word at position i, if cap_map[i] is True, the first
alphabetic character of that word is uppercased. This is the
simple position-based approach used when word counts match.
For across-lines mode where word counts may change significantly,
use _restore_capitalization_aligned() instead.
Args:
text: Normalized text (lowercase)
cap_map: List of booleans from _extract_capitalization
Returns:
Text with capitalization restored
"""
words = text.split()
result = []
for i, word in enumerate(words):
if i < len(cap_map) and cap_map[i]:
result.append(_capitalize_first_alpha(word))
else:
result.append(word)
# If model produced fewer words than original, trailing cap_map
# entries are simply ignored (no capitalization to apply).
return ' '.join(result)
@staticmethod
def _restore_capitalization_aligned(original_text: str, normalized_text: str,
cap_map: list) -> str:
"""Restore capitalization using character-level alignment.
When the model changes word counts (e.g., "inalten" → "in alten"),
position-based _restore_capitalization() fails because word indices
shift. This method uses character-level sequence alignment to correctly
map each normalized word to its corresponding original word, then
applies capitalization based on the original word's case.
For word splits (multiple normalized words mapping to one original
word), only the first normalized word is capitalized, matching the
original word's pattern.
Args:
original_text: Original text (lowercase, punctuation-stripped;
same text that was sent to the model)
normalized_text: Normalized text (lowercase, without punctuation)
cap_map: List of booleans from _extract_capitalization,
one per word in original_text
Returns:
Text with capitalization restored
"""
normalized_words = normalized_text.split()
if not normalized_words:
return normalized_text
original_words = original_text.split()
if not original_words:
return normalized_text
# If word counts match, use simple position-based restoration
if len(normalized_words) == len(original_words):
return MHGNormalizer._restore_capitalization(normalized_text, cap_map)
# Use character-level alignment to map normalized words to original words
norm_to_orig_word = MHGNormalizer._align_normalized_to_original_charlevel(
original_text, normalized_text
)
# Apply capitalization, handling word splits:
# When multiple normalized words map to the same original word,
# only the first gets capitalized (matching the original word's
# pattern where only the first letter is uppercase).
result = []
for i, word in enumerate(normalized_words):
orig_idx = norm_to_orig_word[i]
if orig_idx < len(cap_map) and cap_map[orig_idx]:
is_first = (i == 0 or norm_to_orig_word[i - 1] != orig_idx)
if is_first:
result.append(_capitalize_first_alpha(word))
else:
result.append(word)
else:
result.append(word)
return ' '.join(result)
@staticmethod
def _normalize_word_for_alignment(word: str) -> str:
"""Normalize a word for alignment matching.
Converts to lowercase and simplifies common MHG spelling variations
to improve SequenceMatcher's ability to find correct alignments.
Args:
word: Word to normalize
Returns:
Normalized word form for matching
"""
# Lowercase and long-s to short-s
result = word.lower().replace('ſ', 's')
# Remove diacritics for matching
for char, replacement in [
('â', 'a'), ('î', 'i'), ('û', 'u'), ('ô', 'o'),
('ä', 'a'), ('ë', 'e'), ('ï', 'i'), ('ö', 'o'), ('ü', 'u'),
('é', 'e'), ('è', 'e'), ('ê', 'e'), ('à', 'a'),
('â', 'a'), ('î', 'i'), ('ô', 'o'), ('û', 'u'),
]:
result = result.replace(char, replacement)
return result
@staticmethod
def _align_normalized_to_original_charlevel(
original_text: str,
normalized_text: str
) -> List[int]:
"""Align normalized words to original words using word-level sequence matching.
Uses difflib.SequenceMatcher at the word level with normalized word forms
to find the best alignment. This handles:
- Word splits (one original → multiple normalized)
- Word changes (different spelling)
- Repeated words (matched by sequence context, not just identity)
Args:
original_text: Original text (space-joined words, punctuation-stripped)
normalized_text: Normalized text (space-joined words)
Returns:
List of original word indices, one per normalized word.
Each normalized word is mapped to the index of the original word
it aligns to. The mapping is monotonically non-decreasing.
"""
normalized_words = normalized_text.split()
if not normalized_words:
return []
original_words = original_text.split()
if not original_words:
return [0] * len(normalized_words)
total_orig = len(original_words)
total_norm = len(normalized_words)
# If word counts match, use simple position-based mapping
if total_norm == total_orig:
return list(range(total_norm))
# Use word-level SequenceMatcher with normalized word forms for better matching
# This helps match "dise" to "diſe" and "vâhten" to "vochten" as equal blocks
orig_normalized = [MHGNormalizer._normalize_word_for_alignment(w) for w in original_words]
norm_normalized = [MHGNormalizer._normalize_word_for_alignment(w) for w in normalized_words]
matcher = difflib.SequenceMatcher(None, orig_normalized, norm_normalized,
autojunk=False)
# Build the mapping from normalized word index to original word index
norm_to_orig_word = [0] * total_norm
for tag, i1, i2, j1, j2 in matcher.get_opcodes():
if tag == 'equal':
# Direct match: normalized words j1:j2 match original words i1:i2
for k in range(j2 - j1):
norm_to_orig_word[j1 + k] = i1 + k
elif tag == 'replace':
# Words changed but position preserved (e.g., "vō" → "von").
# Map proportionally within the block using rounded division so
# that word splits (e.g., 1 orig → 2 norm) correctly cluster on
# the same original word rather than drifting to neighbours.
orig_len = i2 - i1
norm_len = j2 - j1
for k in range(norm_len):
if orig_len > 0:
# Rounded proportional index: maps k ∈ [0, norm_len) to
# orig_idx ∈ [i1, i1+orig_len). When norm_len > orig_len
# (word splits), multiple k values round to the same
# orig_idx, keeping split fragments together.
orig_idx = i1 + min(round(k * orig_len / norm_len), orig_len - 1)
else:
orig_idx = i1
norm_to_orig_word[j1 + k] = min(orig_idx, total_orig - 1)
elif tag == 'insert':
# Extra words in normalized (e.g., from word splits).
# Map to the next original word position (i1), since inserts in
# normalization are typically split fragments that belong to the
# following original word, not the preceding one.
for j in range(j1, j2):
if i1 < total_orig:
norm_to_orig_word[j] = i1
else:
norm_to_orig_word[j] = total_orig - 1
# 'delete': original words removed, no normalized words to map
# Ensure monotonicity: each normalized word should map to an original
# word at or after the previous normalized word's mapping
prev_orig = 0
for i in range(total_norm):
if norm_to_orig_word[i] < prev_orig:
norm_to_orig_word[i] = prev_orig
prev_orig = norm_to_orig_word[i]
return norm_to_orig_word
@staticmethod
def _spread_repeated_mappings(
norm_to_orig_word: List[int],
original_words: List[str],
normalized_words: List[str]
) -> List[int]:
"""Spread repeated word mappings across their occurrences in the original.
When the character-level alignment maps multiple identical normalized words
to the same original position, this post-processing step spreads them across
the actual occurrences of that word in the original text.
For example, if "dise" appears 6 times in the original (at indices 0, 4, 8, 12, 16, 20)
and 6 normalized "dise" words all map to orig[0], this spreads them to
[0, 4, 8, 12, 16, 20].
Args:
norm_to_orig_word: Initial alignment mapping
original_words: List of original words
normalized_words: List of normalized words
Returns:
Adjusted alignment mapping with repeated words spread across occurrences
"""
if not norm_to_orig_word:
return norm_to_orig_word
# Build a map of normalized word -> list of original indices where it appears
orig_word_occurrences = {}
for i, word in enumerate(original_words):
word_lower = word.lower()
if word_lower not in orig_word_occurrences:
orig_word_occurrences[word_lower] = []
orig_word_occurrences[word_lower].append(i)
# Track usage of each original word occurrence
orig_usage = {} # orig_idx -> count of times used
# Track which normalized word forms we've seen and how many times
norm_word_count = {} # normalized_word -> count seen so far
result = []
for i, orig_idx in enumerate(norm_to_orig_word):
norm_word_lower = normalized_words[i].lower()
# Check if this normalized word form has multiple occurrences in original
occurrences = orig_word_occurrences.get(norm_word_lower, [])
if len(occurrences) > 1:
# This word appears multiple times in original
# Use the next available occurrence
count_so_far = norm_word_count.get(norm_word_lower, 0)
norm_word_count[norm_word_lower] = count_so_far + 1
if count_so_far < len(occurrences):
# Map to the next occurrence
new_orig_idx = occurrences[count_so_far]
result.append(new_orig_idx)
orig_usage[new_orig_idx] = orig_usage.get(new_orig_idx, 0) + 1
else:
# More occurrences in normalized than original; use alignment
result.append(orig_idx)
orig_usage[orig_idx] = orig_usage.get(orig_idx, 0) + 1
else:
# Unique word or single occurrence; use alignment
result.append(orig_idx)
orig_usage[orig_idx] = orig_usage.get(orig_idx, 0) + 1
return result
@staticmethod
def _assign_lines_by_word_alignment(
original_lines: List[str],
normalized_text: str,
) -> List[int]:
"""Assign each word in normalized_text to a line index using word-level alignment.
Uses _align_normalized_to_original_charlevel (word-level difflib with
normalized word forms) to map each normalized word to its corresponding
original word index, then looks up which line that original word
belongs to.
This approach is robust to MHG character differences (ſ→s, diacritics,
etc.) because the alignment normalizes word forms before matching.
It also handles word splits correctly — multiple normalized words
mapping to the same original word stay on the same line.
"""
normalized_words = normalized_text.split()
if not normalized_words:
return []
num_lines = len(original_lines)
full_orig = ' '.join(original_lines)
# Get word-level alignment: each normalized word → original word index
norm_to_orig_word = MHGNormalizer._align_normalized_to_original_charlevel(
full_orig, normalized_text
)
# Build original word index → line mapping
word_to_line = []
for line_idx, line_text in enumerate(original_lines):
for _ in line_text.split():
word_to_line.append(line_idx)
# Map each normalized word to its line via the original word index
norm_word_to_line = []
prev_line = 0
for orig_word_idx in norm_to_orig_word:
if orig_word_idx < len(word_to_line):
line_idx = word_to_line[orig_word_idx]
else:
line_idx = prev_line
norm_word_to_line.append(line_idx)
prev_line = line_idx
# Post-processing: fix cross-line split fragments.
# When a word at a line boundary (e.g., "en" from "Engülden" → "en gültin")
# was incorrectly assigned to the previous line, detect it by checking if
# the normalized word is a prefix of the first word of the next original line.
total_norm = len(normalized_words)
if len(original_lines) > 1:
for norm_idx in range(total_norm - 1):
curr_line = norm_word_to_line[norm_idx]
next_line = norm_word_to_line[norm_idx + 1]
# Only care about descending line assignments (word drifted to previous line)
if curr_line < next_line and curr_line < num_lines - 1:
# This word is on an earlier line than the next word.
# Check if it's a split fragment: does the next original line's
# first word start with this normalized word (after normalizing)?
next_orig_line = curr_line + 1
next_orig_words = original_lines[next_orig_line].split()
if next_orig_words:
first_next_word = next_orig_words[0]
norm_current = MHGNormalizer._normalize_word_for_alignment(normalized_words[norm_idx])
norm_first_next = MHGNormalizer._normalize_word_for_alignment(first_next_word)
# If the current word is a prefix of the first word of the next original line,
# it's likely a split fragment that belongs on the next line
if (norm_current and
norm_first_next.startswith(norm_current) and
len(norm_current) <= 4): # short fragments are split candidates
norm_word_to_line[norm_idx] = next_line
return norm_word_to_line
@staticmethod
def _get_normalized_to_original_char_mapping(
original_text: str,
normalized_text: str
) -> dict:
"""Map each normalized word start position to its corresponding original character position.
Uses character-level SequenceMatcher to align normalized text to original text,
then computes the starting character position of each normalized word and maps
it to the corresponding position in the original text.
This mapping is used to determine which line each normalized word belongs to
when splitting the normalized output back into lines. Unlike word-level alignment,
character-level alignment correctly handles cases where the model merges multiple
short lines into one.
Args:
original_text: Original text (space-joined words, may contain newlines)
normalized_text: Normalized text (space-joined words)
Returns:
Dictionary mapping normalized word index to original character position
"""
normalized_words = normalized_text.split()
if not normalized_words:
return {}
# Use character-level SequenceMatcher to align normalized text to original text
matcher = difflib.SequenceMatcher(None, original_text, normalized_text,
autojunk=False)
# Build character position mapping from normalized to original
norm_to_orig_char_pos = {}
for tag, i1, i2, j1, j2 in matcher.get_opcodes():
if tag == 'equal':
for k in range(j2 - j1):
norm_to_orig_char_pos[j1 + k] = i1 + k
elif tag == 'replace':
orig_len = i2 - i1
norm_len = j2 - j1
if orig_len > 0 and norm_len > 0:
for k in range(norm_len):
norm_to_orig_char_pos[j1 + k] = i1 + min(
k * orig_len // norm_len, orig_len - 1)
elif norm_len > 0:
for k in range(norm_len):
norm_to_orig_char_pos[j1 + k] = i1 if i1 < len(original_text) else len(original_text) - 1
elif tag == 'insert':
for j in range(j1, j2):
if i1 > 0:
norm_to_orig_char_pos[j] = i1 - 1
elif i1 < len(original_text):
norm_to_orig_char_pos[j] = i1
# Compute start position of each normalized word
word_start_positions = []
pos = 0
for word in normalized_words:
word_start_positions.append(pos)
pos += len(word) + 1 # +1 for space
# Map each normalized word to its original character position
word_to_orig_char = {}
for i, start_pos in enumerate(word_start_positions):
if start_pos in norm_to_orig_char_pos:
word_to_orig_char[i] = norm_to_orig_char_pos[start_pos]
else:
# Fallback: find nearest mapped position
for p in range(start_pos, -1, -1):
if p in norm_to_orig_char_pos:
word_to_orig_char[i] = norm_to_orig_char_pos[p]
break
else:
word_to_orig_char[i] = 0
return word_to_orig_char
@staticmethod
def _align_normalized_to_original(
original_words: List[str],
normalized_words: List[str]
) -> List[int]:
"""Align normalized words to original words using word-level sequence matching.
.. deprecated::
Use :meth:`_align_normalized_to_original_charlevel` instead, which
uses character-level alignment and correctly handles word splits at
line boundaries. This word-level method can incorrectly match a
split word fragment (e.g., "in" from "inallen" → "in allen") to a
standalone identical word on a different line.
Uses difflib.SequenceMatcher to map each normalized word to its
corresponding original word index. This handles word splits and merges
within a single line reasonably well, but can produce incorrect line
assignments when a split fragment matches a word on a different line.
Args:
original_words: Words from the joined original text
normalized_words: Words from the model's normalized output
Returns:
List of original word indices, one per normalized word.
Each normalized word is mapped to the index of the original word
it aligns to. For insertions (extra words from splits), the
preceding original word's index is used.
"""
if not normalized_words:
return []
if not original_words:
return [0] * len(normalized_words)
# Use SequenceMatcher to align normalized words to original words.
# autojunk=False ensures common words (like "und", "der") are still
# used for anchoring the alignment, which improves quality.
matcher = difflib.SequenceMatcher(None, original_words, normalized_words,
autojunk=False)
# Map each normalized word index to its corresponding original word index
norm_to_orig = [0] * len(normalized_words)
for tag, i1, i2, j1, j2 in matcher.get_opcodes():
if tag == 'equal':
# Direct match: normalized word j corresponds to original word i
for k in range(j2 - j1):
norm_to_orig[j1 + k] = i1 + k
elif tag == 'replace':
# Words changed but position preserved (e.g., "vō" → "von").
# Map each normalized word to the corresponding original word
# by proportional position within the replace block.
orig_len = i2 - i1
norm_len = j2 - j1
for k in range(norm_len):
orig_idx = i1 + min(k * orig_len // norm_len, orig_len - 1) if norm_len > 0 else i1
norm_to_orig[j1 + k] = orig_idx
elif tag == 'insert':
# Extra words in normalized text (e.g., from word splits like
# "inalten" → "in" "alten"). Assign to the original word just
# before the insertion point.
for j in range(j1, j2):
if i1 > 0:
norm_to_orig[j] = i1 - 1
elif i1 < len(original_words):
norm_to_orig[j] = i1
# 'delete': words removed from original, no normalized words to assign
return norm_to_orig
def _get_token_count(self, text: str) -> int:
"""
Get the number of tokens for a text including the 'normalize: ' prefix.
Args:
text: Input text (without prefix)
Returns:
Total token count including prefix and special tokens
"""
prefix = "normalize: "
input_text = prefix + text
return len(self.tokenizer.encode(input_text))
def _create_word_windows(
self,
words: List[str],
max_window_tokens: int,
stride_step_tokens: int
) -> List[Tuple[int, int]]:
"""
Create overlapping word windows that fit within the token budget.
Each window is defined by (start_word_idx, end_word_idx) where
end_word_idx is exclusive. Windows overlap so that every word
appears in the interior of at least one window, ensuring good
context for the normalization of each word.
Args:
words: List of words from the input text
max_window_tokens: Maximum number of tokens per window
(including the 'normalize: ' prefix)
stride_step_tokens: Number of tokens to advance between
consecutive windows
Returns:
List of (start_idx, end_idx) tuples defining word windows
"""
if not words:
return []
prefix = "normalize: "
prefix_tokens = len(self.tokenizer.encode(prefix, add_special_tokens=False))
effective_max = max_window_tokens - prefix_tokens
if effective_max <= 0:
# Window too small for any content, fall back to single-word windows
return [(i, i + 1) for i in range(len(words))]
# Pre-compute token count for each word (with leading space for joining)
word_token_counts = []
for w in words:
tc = len(self.tokenizer.encode(" " + w, add_special_tokens=False))
word_token_counts.append(tc)
# Cumulative token counts for efficient range queries
cum_tokens = [0]
for tc in word_token_counts:
cum_tokens.append(cum_tokens[-1] + tc)
windows = []
start_idx = 0
while start_idx < len(words):
# Find the end of this window: the largest end_idx where
# the token count from start_idx to end_idx fits in effective_max
end_idx = start_idx
while end_idx < len(words) and cum_tokens[end_idx + 1] - cum_tokens[start_idx] <= effective_max:
end_idx += 1
# Ensure at least one word per window (even if it exceeds token limit)
if end_idx == start_idx:
end_idx = start_idx + 1
windows.append((start_idx, end_idx))
# If we've covered all words, stop
if end_idx >= len(words):
break
# Advance start by stride_step_tokens worth of words
new_start = start_idx
while new_start < end_idx and cum_tokens[new_start + 1] - cum_tokens[start_idx] < stride_step_tokens:
new_start += 1
# Ensure we make progress (at least one word advance)
if new_start <= start_idx:
new_start = start_idx + 1
start_idx = new_start
return windows
def _merge_window_predictions(
self,
total_words: int,
windows: List[Tuple[int, int]],
window_predictions: List[List[str]],
original_words: List[str]
) -> List[str]:
"""
Merge predictions from overlapping windows, preferring window centers.
For each word position, the prediction from the window where that
word is closest to the center is used. This ensures that words in
overlap regions get the best context-informed prediction, since
the center of a window has the most surrounding context.
When the model produces a different number of words than expected
(due to word splits or merges), sequence alignment is used instead
of proportional mapping to prevent words from disappearing.
Args:
total_words: Total number of words in the original text
windows: List of (start_idx, end_idx) tuples defining word windows
window_predictions: List of word lists, one per window
original_words: Original input words (used as fallback)
Returns:
List of merged prediction words
"""
# Each position stores a list of words (to handle splits where
# multiple predicted words map to one original position).
merged = [None] * total_words
best_distances = [float('inf')] * total_words
for (start_idx, end_idx), pred_words in zip(windows, window_predictions):
# Center of the window
center = (start_idx + end_idx - 1) / 2.0
expected_len = end_idx - start_idx
if len(pred_words) == expected_len:
# Perfect alignment: direct word-by-word mapping
for i in range(expected_len):
word_pos = start_idx + i
if word_pos < total_words:
distance = abs(word_pos - center)
if distance < best_distances[word_pos]:
best_distances[word_pos] = distance
merged[word_pos] = [pred_words[i]]
else:
# Length mismatch: the model produced a different number of words
# than expected (word splits or merges). Use character-level
# alignment to properly map each predicted word to its original
# position, preventing word disappearance and correctly handling
# word splits (e.g., "inalten" → "in alten").
window_words = original_words[start_idx:end_idx]
window_text = ' '.join(window_words)
pred_text = ' '.join(pred_words)
alignment = self._align_normalized_to_original_charlevel(
window_text, pred_text
)
# Group predicted words by their aligned original position.
# Multiple predicted words mapping to the same position
# (e.g., from a split like "inalten" → "in" "alten") are
# collected together so no words are lost.
pos_to_preds = {}
for j, orig_idx in enumerate(alignment):
abs_pos = start_idx + orig_idx
if abs_pos not in pos_to_preds:
pos_to_preds[abs_pos] = []
pos_to_preds[abs_pos].append(pred_words[j])
for abs_pos, preds in pos_to_preds.items():
if abs_pos < total_words:
distance = abs(abs_pos - center)
if distance < best_distances[abs_pos]:
best_distances[abs_pos] = distance
merged[abs_pos] = preds
# Fill any unfilled positions with original words (model deletions
# within a window — preserve the original un-normalized word rather
# than dropping it entirely).
for i in range(total_words):
if merged[i] is None:
merged[i] = [original_words[i]]
# Flatten lists of words into a single word list.
# Splits produce multiple words at one position; merges produce one.
result = []
for words_at_pos in merged:
result.extend(words_at_pos)
return result
def normalize_long_text(
self,
text: str,
max_length: Optional[int] = None,
stride_window_tokens: Optional[int] = None,
stride_step_tokens: Optional[int] = None,
num_beams: Optional[int] = None,
length_penalty: Optional[float] = None,
do_sample: Optional[bool] = None,
top_p: Optional[float] = None,
temperature: Optional[float] = None,
preserve_punctuation: Optional[bool] = None,
preserve_capitalization: Optional[bool] = None,
progress_callback: Optional[callable] = None
) -> str:
"""
Normalize a potentially long text using sliding window with stride.
Splits the input text into overlapping windows of words, normalizes
each window independently, then merges the predictions. For each word
position, the prediction from the window where that word is closest
to the center is preferred, as center predictions have the most
surrounding context.
If the text fits within max_length tokens, delegates to normalize_text()
for single-pass normalization (no windowing needed).
Args:
text: Input text to normalize (original MHG)
max_length: Maximum generation length (tokens)
stride_window_tokens: Max input tokens per window (including prefix).
Default: 400 (from config)
stride_step_tokens: Token advance between windows. Controls overlap;
overlap = window - step. Default: 300 (from config)
num_beams: Number of beams for beam search
length_penalty: Length penalty for beam search
do_sample: Whether to use sampling
top_p: Nucleus sampling threshold
temperature: Temperature for sampling
preserve_punctuation: If True, strip editorial punctuation before
normalization and restore it afterwards.
Default: None (read from config, falls back to True)
preserve_capitalization: If True, record capitalization before
normalization (lowercasing the input) and
restore it afterwards.
Default: None (read from config, falls back to True)
Returns:
Normalized text
"""
# Resolve parameters
max_length = max_length or self.config['inference']['max_length']
stride_window_tokens = stride_window_tokens or self.config['inference'].get('stride_window_tokens', 400)
stride_step_tokens = stride_step_tokens or self.config['inference'].get('stride_step_tokens', 300)
if preserve_punctuation is None:
preserve_punctuation = self.config['inference'].get('preserve_punctuation', True)
if preserve_capitalization is None:
preserve_capitalization = self.config['inference'].get('preserve_capitalization', True)
# Extract editorial punctuation before normalization
punct_map = None
if preserve_punctuation:
text, punct_map = self._extract_punctuation(text)
# Extract capitalization before normalization (after punctuation extraction
# so cap_map indices align with punct_map indices)
cap_map = None
if preserve_capitalization:
text, cap_map = self._extract_capitalization(text)
# Check if stride is actually needed
token_count = self._get_token_count(text)
if token_count <= max_length:
# Text fits in a single window, use simple normalization
result = self.normalize_text(
text, max_length=max_length, stride=False,
num_beams=num_beams, length_penalty=length_penalty,
do_sample=do_sample, top_p=top_p, temperature=temperature,
preserve_punctuation=False, # already stripped
preserve_capitalization=False # already lowercased
)
if cap_map is not None:
result = self._restore_capitalization(result, cap_map)
if punct_map is not None:
result = self._restore_punctuation(result, punct_map)
return result
# Split text into words
words = text.split()
if not words:
return ""
if len(words) == 1:
# Single word, just normalize directly (even if token count exceeds limit)
result = self.normalize_text(
text, max_length=max_length, stride=False,
num_beams=num_beams, length_penalty=length_penalty,
do_sample=do_sample, top_p=top_p, temperature=temperature,
preserve_punctuation=False, # already stripped
preserve_capitalization=False # already lowercased
)
if cap_map is not None:
result = self._restore_capitalization(result, cap_map)
if punct_map is not None:
result = self._restore_punctuation(result, punct_map)
return result
# Create overlapping windows
windows = self._create_word_windows(words, stride_window_tokens, stride_step_tokens)
num_windows = len(windows)
# if num_windows > 1:
# print(f" Stride: {token_count} tokens > {max_length} max, "
# f"using {num_windows} windows "
# f"(window={stride_window_tokens}tok, step={stride_step_tokens}tok, "
# f"overlap={stride_window_tokens - stride_step_tokens}tok)")
# Normalize each window
window_predictions = []
windows_iter = windows
if num_windows > 1:
windows_iter = tqdm(windows, desc=" Progress",
total=num_windows, ncols=100,
bar_format='{desc}: {percentage:3.0f}%|{bar}| {elapsed}')
for wi, (start_idx, end_idx) in enumerate(windows_iter):
window_text = ' '.join(words[start_idx:end_idx])
prediction = self.normalize_text(
window_text, max_length=max_length, stride=False,
num_beams=num_beams, length_penalty=length_penalty,
do_sample=do_sample, top_p=top_p, temperature=temperature,
preserve_punctuation=False, # already stripped
preserve_capitalization=False # already lowercased
)
# Report progress scaled to a low range (5-75%) so the UI bar
# reflects the fact that significant post-processing work remains
# after this method returns. With 5 beams, each window is
# expensive, but the caller-side work (character-level alignment
# in _align_normalized_to_original_charlevel, line reassignment,
# cap/punct restoration in normalize_multiline_text) can easily
# take as long as or longer than the window loop itself.
#
# Scale: window 0 → 5%, last window → 75%.
if progress_callback:
pct = 5 + int(wi * 70 / max(num_windows - 1, 1))
progress_callback(pct, 100)
pred_words = prediction.split()
window_predictions.append(pred_words)
# Merge predictions from overlapping windows
merged_words = self._merge_window_predictions(
len(words), windows, window_predictions, words
)
# Signal that generation + merge are done. The remaining 25% of
# progress will be consumed by caller-side post-processing
# (alignment, capitalization/punctuation restoration, line splitting)
# before app.py sets progress to 100.
if progress_callback:
progress_callback(78, 100)
result = ' '.join(merged_words)
if cap_map is not None:
# Use character-level alignment for capitalization restoration
# because word counts may change due to splits/merges.
result = self._restore_capitalization_aligned(
text, result, cap_map
)
if punct_map is not None:
# Use character-level alignment for punctuation restoration
# instead of position-based mapping, because word counts may
# change due to splits/merges in the merge step.
result = self._restore_punctuation_aligned(
text, result, punct_map
)
return result
def normalize_text(
self,
text: str,
max_length: Optional[int] = None,
num_beams: Optional[int] = None,
length_penalty: Optional[float] = None,
do_sample: Optional[bool] = None,
top_p: Optional[float] = None,
temperature: Optional[float] = None,
stride: Optional[bool] = None,
stride_window_tokens: Optional[int] = None,
stride_step_tokens: Optional[int] = None,
preserve_punctuation: Optional[bool] = None,
preserve_capitalization: Optional[bool] = None,
attach_en_proclitic: Optional[bool] = None,
lenition_t_after_n: Optional[bool] = None,
lenition_t_after_l: Optional[bool] = None,
niet_to_niht: Optional[bool] = None,
common_apocopes: Optional[bool] = None,
progress_callback: Optional[callable] = None
) -> str:
"""
Normalize a single text string.
When stride=True (default) and the input text exceeds the model's
maximum sequence length, automatically uses sliding-window inference
via normalize_long_text() to process the full text without truncation.
Args:
text: Input text to normalize (original MHG)
max_length: Maximum generation length
num_beams: Number of beams for beam search
length_penalty: Length penalty for beam search
do_sample: Whether to use sampling (vs beam search)
top_p: Nucleus sampling threshold (0.0 to 1.0)
temperature: Temperature for sampling (0.1 to 2.0)
stride: Whether to use sliding window for long texts.
Default: True (from config). Set to False to always truncate.
stride_window_tokens: Max input tokens per window (for stride mode)
stride_step_tokens: Token advance between windows (for stride mode)
preserve_punctuation: If True, strip editorial punctuation before
normalization and restore it afterwards.
Default: None (read from config, falls back to True)
preserve_capitalization: If True, record capitalization before
normalization (lowercasing the input) and
restore it afterwards.
Default: None (read from config, falls back to True)
attach_en_proclitic: If True, attach free-standing negation proclitic
"en" to the following word with a hyphen
(e.g. "en guot" → "en-guot"). Applied as the
final post-processing step.
Default: None (read from config, falls back to False)
Returns:
Normalized text
"""
# Resolve preserve_punctuation
if preserve_punctuation is None:
preserve_punctuation = self.config['inference'].get('preserve_punctuation', True)
# Resolve preserve_capitalization
if preserve_capitalization is None:
preserve_capitalization = self.config['inference'].get('preserve_capitalization', True)
# Resolve post-processing options
if attach_en_proclitic is None:
attach_en_proclitic = self.config.get('post_processing', {}).get('attach_en_proclitic', False)
if lenition_t_after_n is None:
lenition_t_after_n = self.config.get('post_processing', {}).get('lenition_t_after_n', False)
if lenition_t_after_l is None:
lenition_t_after_l = self.config.get('post_processing', {}).get('lenition_t_after_l', False)
if niet_to_niht is None:
niet_to_niht = self.config.get('post_processing', {}).get('niet_to_niht', False)
if common_apocopes is None:
common_apocopes = self.config.get('post_processing', {}).get('common_apocopes', False)
# Extract editorial punctuation before normalization
punct_map = None
if preserve_punctuation:
text, punct_map = self._extract_punctuation(text)
# Extract capitalization before normalization (after punctuation extraction
# so cap_map indices align with punct_map indices)
cap_map = None
if preserve_capitalization:
text, cap_map = self._extract_capitalization(text)
# Resolve stride setting
if stride is None:
stride = self.config['inference'].get('stride', True)
# If stride is enabled, check if we need it
if stride:
max_length_resolved = max_length or self.config['inference']['max_length']
token_count = self._get_token_count(text)
if token_count > max_length_resolved:
result = self.normalize_long_text(
text, max_length=max_length_resolved,
stride_window_tokens=stride_window_tokens,
stride_step_tokens=stride_step_tokens,
num_beams=num_beams, length_penalty=length_penalty,
do_sample=do_sample, top_p=top_p, temperature=temperature,
preserve_punctuation=False, # already stripped above
preserve_capitalization=False, # already lowercased above
progress_callback=progress_callback
)
if cap_map is not None:
result = self._restore_capitalization(result, cap_map)
if punct_map is not None:
result = self._restore_punctuation(result, punct_map)
# Post-processing: common apocopes (before en-attachment so "en wile" → "en wil" → "enwil")
if common_apocopes:
result = _common_apocopes(result)
# Post-processing: attach negation proclitic
if attach_en_proclitic:
result = _attach_en_proclitic(result)
# Post-processing: lenition of t after n before vowel
if lenition_t_after_n:
result = _lenition_t_after_n_before_vowel(result)
# Post-processing: lenition of t after l before vowel
if lenition_t_after_l:
result = _lenition_t_after_l_before_vowel(result)
# Post-processing: niet → niht
if niet_to_niht:
result = _niet_to_niht(result)
return result
# Use config values if not provided
max_length = max_length or self.config['inference']['max_length']
num_beams = num_beams or self.config['inference']['num_beams']
length_penalty = length_penalty or self.config['inference']['length_penalty']
do_sample = do_sample if do_sample is not None else self.config['inference']['do_sample']
top_p = top_p if top_p is not None else self.config['inference']['top_p']
temperature = temperature if temperature is not None else self.config['inference'].get('temperature', 1.0)
# Add prefix if using T5-style model
prefix = "normalize: "
input_text = prefix + text
# Tokenize
inputs = self.tokenizer(
input_text,
return_tensors='pt',
padding=True,
truncation=True,
max_length=max_length
)
# Move to device
inputs = {k: v.to(self.device) for k, v in inputs.items()}
# Generate
with torch.no_grad():
# Prepare generation arguments by unpacking inputs directly
generate_kwargs = {
**inputs, # This includes 'input_ids' and 'attention_mask'
'max_length': max_length,
'num_beams': num_beams,
'length_penalty': length_penalty,
'do_sample': do_sample,
'early_stopping': True
}
# Only add sampling-specific parameters if sampling is enabled
if do_sample:
generate_kwargs['top_p'] = top_p
generate_kwargs['temperature'] = temperature
outputs = self.model.generate(**generate_kwargs)
# Decode
normalized = self.tokenizer.decode(outputs[0], skip_special_tokens=True)
# Restore capitalization (before punctuation, since cap_map indices
# align with punctuation-stripped words)
if cap_map is not None:
normalized = self._restore_capitalization(normalized, cap_map)
# Restore editorial punctuation
if punct_map is not None:
normalized = self._restore_punctuation(normalized, punct_map)
# Post-processing: common apocopes (before en-attachment so "en wile" → "en wil" → "enwil")
if common_apocopes:
normalized = _common_apocopes(normalized)
# Post-processing: attach negation proclitic (after apocopes)
if attach_en_proclitic:
normalized = _attach_en_proclitic(normalized)
# Post-processing: lenition of t after n before vowel
if lenition_t_after_n:
normalized = _lenition_t_after_n_before_vowel(normalized)
# Post-processing: lenition of t after l before vowel
if lenition_t_after_l:
normalized = _lenition_t_after_l_before_vowel(normalized)
# Post-processing: niet → niht
if niet_to_niht:
normalized = _niet_to_niht(normalized)
return normalized
def normalize_texts(
self,
texts: List[str],
batch_size: int = 8,
show_progress: bool = True,
**kwargs
) -> List[str]:
"""
Normalize multiple texts efficiently using batching.
Args:
texts: List of input texts
batch_size: Batch size for processing
show_progress: Whether to show progress bar
**kwargs: Additional arguments for normalize_text
(including stride, stride_window_tokens, stride_step_tokens)
Returns:
List of normalized texts
"""
normalized_texts = []
# Calculate total number of batches
total_batches = (len(texts) + batch_size - 1) // batch_size
# Create progress bar if requested
iterator = range(0, len(texts), batch_size)
if show_progress:
iterator = tqdm(iterator, desc="Normalizing lines",
total=total_batches,
ncols=100,
bar_format='{desc}: {percentage:3.0f}%|{bar}| {elapsed}')
for i in iterator:
batch = texts[i:i + batch_size]
batch_normalized = [
self.normalize_text(text, **kwargs)
for text in batch
]
normalized_texts.extend(batch_normalized)
return normalized_texts
def normalize_file(
self,
input_path: str,
output_path: str,
batch_size: int = 8,
across_lines: bool = True,
preserve_punctuation: Optional[bool] = None,
preserve_capitalization: Optional[bool] = None,
**kwargs
):
"""
Normalize texts from a file.
When across_lines=True (default), all lines are joined into a single
continuous text and normalized together, utilizing the model's full
capacity and sliding-window stride across line boundaries. The output
is split back into lines based on original word counts per line.
When across_lines=False, each line is normalized independently
(original behavior).
Args:
input_path: Path to input file (one text per line)
output_path: Path to output file
batch_size: Batch size for processing (only used when across_lines=False)
across_lines: If True, join all lines and normalize as one continuous
text, then split back by word counts. Default: True.
preserve_punctuation: If True, strip editorial punctuation before
normalization and restore it afterwards.
Default: None (read from config, falls back to True)
preserve_capitalization: If True, record capitalization before
normalization (lowercasing the input) and
restore it afterwards.
Default: None (read from config, falls back to True)
**kwargs: Additional arguments for normalize_text
(including stride, stride_window_tokens, stride_step_tokens)
"""
input_path = Path(input_path)
output_path = Path(output_path)
# Resolve preserve_punctuation
if preserve_punctuation is None:
preserve_punctuation = self.config['inference'].get('preserve_punctuation', True)
# Resolve preserve_capitalization
if preserve_capitalization is None:
preserve_capitalization = self.config['inference'].get('preserve_capitalization', True)
# Read input file — do NOT lowercase here; capitalization extraction
# handles lowercasing when preserve_capitalization is True, otherwise
# we lowercase explicitly below.
# Preserve empty lines: track their positions so they can be reinserted
# in the output at the same locations.
print(f"\nNormalizing: {input_path}")
with open(input_path, 'r', encoding='utf-8') as f:
all_lines = [line.rstrip('\n').rstrip('\r') for line in f]
# Build a map of which line indices are empty, and extract non-empty
# lines for normalization. Empty lines will be reinserted as-is.
empty_line_indices = set()
texts = []
for i, line in enumerate(all_lines):
if line.strip() == '':
empty_line_indices.add(i)
else:
texts.append(line.strip())
# Map from non-empty line index to original line index, so we can
# place normalized output back into the correct positions.
non_empty_to_original = []
for i, line in enumerate(all_lines):
if line.strip() != '':
non_empty_to_original.append(i)
# If preserve_capitalization is disabled, lowercase all input now
# (the model expects lowercase input). When preserve_capitalization
# is enabled, _extract_capitalization will lowercase the text and
# record the original capitalization pattern for later restoration.
if not preserve_capitalization:
texts = [line.lower() for line in texts]
total_texts = len(texts)
#print(f"Loaded {total_texts} lines")
if across_lines:
# Per-line pre-processing: extract punctuation and capitalization
# before joining into a single text for normalization.
# Order matters: punctuation extraction first (on original-case text),
# then capitalization extraction (on punctuation-stripped text).
# This ensures cap_map and punct_map are both indexed by the same
# word list (punctuation-stripped words), making per-line restoration
# straightforward.
line_punct_maps = None
line_cap_maps = None
stripped_lines = None # punctuation-stripped lines (original case)
lowered_lines = None # lowercase, punctuation-stripped lines
if preserve_punctuation:
line_punct_maps = []
stripped_lines = []
for line in texts:
stripped_line, punct_map = self._extract_punctuation(line)
stripped_lines.append(stripped_line)
line_punct_maps.append(punct_map)
if preserve_capitalization:
line_cap_maps = []
lowered_lines = []
# Capitalization extraction operates on the punctuation-stripped
# text (if punctuation was extracted) or the original text.
cap_source_lines = stripped_lines if stripped_lines is not None else texts
for line in cap_source_lines:
lowered_line, cap_map = self._extract_capitalization(line)
lowered_lines.append(lowered_line)
line_cap_maps.append(cap_map)
# Determine the text to send to the model:
# - If both preserved: lowered_lines (lowercase, no punctuation)
# - If only cap preserved: lowered_lines (lowercase, with punctuation)
# - If only punct preserved: stripped_lines (original case, no punctuation)
# - If neither preserved: texts (lowercased earlier)
if lowered_lines is not None:
word_counts = [len(line.split()) for line in lowered_lines]
full_text = ' '.join(lowered_lines)
elif stripped_lines is not None:
word_counts = [len(line.split()) for line in stripped_lines]
full_text = ' '.join(stripped_lines)
else:
word_counts = [len(line.split()) for line in texts]
full_text = ' '.join(texts)
normalized_full = self.normalize_text(
full_text, preserve_punctuation=False,
preserve_capitalization=False, **kwargs
)
# Split normalized output back into lines using occurrence-based tracking
# with word split detection.
#
# This correctly handles:
# - Word splits (e.g., "enwere" → "en wære"): detected via alignment
# - Repeated words (e.g., "dise" x10): occurrence tracking maps k-th occurrence
# in normalized to k-th occurrence in original
# - Different word counts: proportional mapping fallback
num_lines = len(texts)
normalized_words = normalized_full.split()
total_norm = len(normalized_words)
# Build word_to_line mapping from original word indices to line indices
word_to_line = []
for line_idx, count in enumerate(word_counts):
for _ in range(count):
word_to_line.append(line_idx)
# Calculate total original words
total_orig = sum(word_counts)
# Build original word list for occurrence tracking
original_words = full_text.split()
# Helper function to normalize word for occurrence tracking
def normalize_word_for_tracking(word: str) -> str:
"""Normalize word for occurrence tracking (case-insensitive, ſ→s, simplified)."""
# Basic normalization: lowercase and long-s to short-s
base = word.lower().replace('ſ', 's')
# Simplify common MHG spelling variations for matching
# Remove diacritics for matching purposes
simplified = base
for char, replacement in [
('â', 'a'), ('î', 'i'), ('û', 'u'), ('ô', 'o'),
('ä', 'a'), ('ë', 'e'), ('ï', 'i'), ('ö', 'o'), ('ü', 'u'),
('é', 'e'), ('è', 'e'), ('ê', 'e'), ('à', 'a'),
('â', 'a'), ('î', 'i'), ('ô', 'o'), ('û', 'u'),
]:
simplified = simplified.replace(char, replacement)
return simplified
# Step 1: Build occurrence tracking for original words
# Maps normalized_word -> list of original indices where it appears
orig_occurrences = {}
for i, word in enumerate(original_words):
norm_word = normalize_word_for_tracking(word)
if norm_word not in orig_occurrences:
orig_occurrences[norm_word] = []
orig_occurrences[norm_word].append(i)
# Get character-level alignment first (used in Steps 2 and 3)
norm_to_orig_word = self._align_normalized_to_original_charlevel(
full_text, normalized_full
)
# Step 2: Map each normalized word to original
# Use alignment as primary method. For word splits (multiple normalized
# words mapping to the same original), the split detection in Step 3
# handles keeping them on the same line.
#
# The occurrence-based approach was removed because it caused incorrect
# mappings: when a word like "daz" appears multiple times in the original,
# counting global occurrences led to mapping normalized words to wrong
# positions (e.g., line 10's "daz" mapped to line 4's "das").
#
# The character-level alignment correctly handles:
# - Word splits (e.g., "inalten" → "in alten")
# - Word merges
# - Position-based mapping for non-split words
norm_to_orig_occurrence = norm_to_orig_word.copy()
# Step 3: Use alignment to detect word splits
# A true word split is when consecutive normalized words map to the same original
# AND the original word is longer (contains the split parts)
# We filter out false splits caused by repeated words like "dise"
# (norm_to_orig_word already computed above for fallback)
# Detect split groups: consecutive normalized words that map to the same original
# These should stay on the same line
# BUT only if it's a true word split (original word contains the parts)
split_group = {} # norm_idx -> group_id
group_id = 0
i = 0
while i < total_norm:
if i == total_norm - 1:
# Last word, no split possible
split_group[i] = group_id
break
# Check if current and next word map to same original (potential word split)
if norm_to_orig_word[i] == norm_to_orig_word[i + 1]:
current_orig = norm_to_orig_word[i]
orig_word = original_words[current_orig]
# Collect consecutive words mapping to this original
split_words = []
j = i
while j < total_norm and norm_to_orig_word[j] == current_orig:
split_words.append(normalized_words[j])
j += 1
# Check if this is a true word split:
# The combined normalized words should be similar to the original
# (allowing for minor spelling differences)
combined = ''.join(split_words)
is_true_split = (
len(combined) >= len(orig_word) * 0.5 and # Combined is at least half the original length
len(combined) <= len(orig_word) * 1.5 and # Combined is at most 1.5x original length
len(split_words) > 1 # At least 2 parts
)
if is_true_split:
# True word split - keep parts together
for k in range(i, j):
split_group[k] = group_id
group_id += 1
i = j
else:
# False split (repeated words) - assign each word its own group
for k in range(i, j):
split_group[k] = group_id
group_id += 1
i = j
else:
split_group[i] = group_id
group_id += 1
i += 1
# Assign words to lines using word-level alignment
norm_to_line = self._assign_lines_by_word_alignment(
lowered_lines if lowered_lines is not None else
(stripped_lines if stripped_lines is not None else texts),
normalized_full
)
line_word_lists = [[] for _ in range(num_lines)]
for norm_idx in range(total_norm):
target_line = norm_to_line[norm_idx]
line_word_lists[target_line].append(normalized_words[norm_idx])
# Ensure no empty lines by redistributing words if necessary
# This can happen when the model produces very different structure
empty_lines = [i for i in range(num_lines) if not line_word_lists[i]]
if empty_lines:
# Find lines with multiple words and move some to empty lines
for empty_idx in empty_lines:
# Find the nearest line with words
best_source = None
best_dist = float('inf')
for src_idx, words in enumerate(line_word_lists):
if len(words) > 1:
dist = abs(src_idx - empty_idx)
if dist < best_dist:
best_dist = dist
best_source = src_idx
if best_source is not None:
# Move the last word from source to empty line
word = line_word_lists[best_source].pop()
line_word_lists[empty_idx].append(word)
normalized = []
for line_idx in range(num_lines):
line_text = ' '.join(line_word_lists[line_idx])
# Restore capitalization first (normalized output is lowercase)
if preserve_capitalization and line_cap_maps is not None and line_idx < len(line_cap_maps):
cap_ref = lowered_lines[line_idx] if lowered_lines is not None else None
if cap_ref is not None:
line_text = self._restore_capitalization_aligned(
cap_ref, line_text, line_cap_maps[line_idx]
)
# Then restore punctuation (on capitalized text)
if preserve_punctuation and line_punct_maps is not None and line_idx < len(line_punct_maps):
punct_ref = stripped_lines[line_idx] if stripped_lines is not None else None
if punct_ref is not None:
line_text = self._restore_punctuation_aligned(
punct_ref, line_text, line_punct_maps[line_idx]
)
normalized.append(line_text)
# Ensure we have the right number of output lines
while len(normalized) < total_texts:
normalized.append('')
else:
# Original behavior: normalize each line independently
# Forward preserve_punctuation and preserve_capitalization explicitly
# so normalize_text doesn't fall back to the config default
kwargs['preserve_punctuation'] = preserve_punctuation
kwargs['preserve_capitalization'] = preserve_capitalization
print()
normalized = self.normalize_texts(
texts,
batch_size=batch_size,
show_progress=True,
**kwargs
)
# Reinsert empty lines at their original positions.
# normalized[] contains results for non-empty lines only; we need
# to merge them back into the full line sequence, preserving empty
# lines at the same positions as the input.
total_original_lines = len(all_lines)
if empty_line_indices:
full_output = [''] * total_original_lines
for norm_idx, orig_idx in enumerate(non_empty_to_original):
if norm_idx < len(normalized) and orig_idx < total_original_lines:
full_output[orig_idx] = normalized[norm_idx]
# Any remaining normalized lines beyond the mapping go at the end
for norm_idx in range(len(non_empty_to_original), len(normalized)):
full_output.append(normalized[norm_idx])
output_lines = full_output
else:
output_lines = normalized
# Write output
print(f"\nWriting to: {output_path}")
output_path.parent.mkdir(parents=True, exist_ok=True)
with open(output_path, 'w', encoding='utf-8') as f:
for text in output_lines:
f.write(text + '\n')
print("Done!")
def normalize_multiline_text(
self,
input_text: str,
preserve_punctuation: Optional[bool] = None,
preserve_capitalization: Optional[bool] = None,
attach_en_proclitic: Optional[bool] = None,
lenition_t_after_n: Optional[bool] = None,
lenition_t_after_l: Optional[bool] = None,
niet_to_niht: Optional[bool] = None,
common_apocopes: Optional[bool] = None,
progress_callback: Optional[callable] = None,
**kwargs
) -> str:
"""
Normalize multi-line text preserving line structure with cross-line context.
Replicates the across_lines=True logic from normalize_file() but operates
on a string rather than a file.
Returns normalized text with the same number of lines as input.
"""
if self.config is not None:
if preserve_punctuation is None:
preserve_punctuation = self.config.get('preserve_editorial_punctuation', False)
if preserve_capitalization is None:
preserve_capitalization = self.config.get('preserve_capitalization', False)
if attach_en_proclitic is None:
attach_en_proclitic = self.config.get('post_processing', {}).get('attach_en_proclitic', False)
if lenition_t_after_n is None:
lenition_t_after_n = self.config.get('post_processing', {}).get('lenition_t_after_n', False)
if lenition_t_after_l is None:
lenition_t_after_l = self.config.get('post_processing', {}).get('lenition_t_after_l', False)
if niet_to_niht is None:
niet_to_niht = self.config.get('post_processing', {}).get('niet_to_niht', False)
if common_apocopes is None:
common_apocopes = self.config.get('post_processing', {}).get('common_apocopes', False)
else:
preserve_punctuation = preserve_punctuation if preserve_punctuation is not None else False
preserve_capitalization = preserve_capitalization if preserve_capitalization is not None else False
attach_en_proclitic = attach_en_proclitic if attach_en_proclitic is not None else False
lenition_t_after_n = lenition_t_after_n if lenition_t_after_n is not None else False
lenition_t_after_l = lenition_t_after_l if lenition_t_after_l is not None else False
niet_to_niht = niet_to_niht if niet_to_niht is not None else False
common_apocopes = common_apocopes if common_apocopes is not None else False
all_lines = [line.rstrip('\n').rstrip('\r') for line in input_text.split('\n')]
# Build empty-line map and extract non-empty lines
empty_line_indices = set()
texts = []
for i, line in enumerate(all_lines):
if line.strip() == '':
empty_line_indices.add(i)
else:
texts.append(line.strip())
non_empty_to_original = []
for i, line in enumerate(all_lines):
if line.strip() != '':
non_empty_to_original.append(i)
if not preserve_capitalization:
texts = [line.lower() for line in texts]
total_texts = len(texts)
if total_texts == 0:
return '\n'.join(all_lines)
# across_lines=True logic
line_punct_maps = None
line_cap_maps = None
stripped_lines = None
lowered_lines = None
if preserve_punctuation:
line_punct_maps = []
stripped_lines = []
for line in texts:
stripped_line, punct_map = self._extract_punctuation(line)
stripped_lines.append(stripped_line)
line_punct_maps.append(punct_map)
if preserve_capitalization:
line_cap_maps = []
lowered_lines = []
cap_source_lines = stripped_lines if stripped_lines is not None else texts
for line in cap_source_lines:
lowered_line, cap_map = self._extract_capitalization(line)
lowered_lines.append(lowered_line)
line_cap_maps.append(cap_map)
if lowered_lines is not None:
word_counts = [len(line.split()) for line in lowered_lines]
full_text = ' '.join(lowered_lines)
elif stripped_lines is not None:
word_counts = [len(line.split()) for line in stripped_lines]
full_text = ' '.join(stripped_lines)
else:
word_counts = [len(line.split()) for line in texts]
full_text = ' '.join(texts)
normalized_full = self.normalize_text(
full_text, preserve_punctuation=False,
preserve_capitalization=False,
progress_callback=progress_callback,
common_apocopes=False, # applied below on final result
**kwargs
)
num_lines = len(texts)
normalized_words = normalized_full.split()
total_norm = len(normalized_words)
word_to_line = []
for line_idx, count in enumerate(word_counts):
for _ in range(count):
word_to_line.append(line_idx)
total_orig = sum(word_counts)
original_words = full_text.split()
def normalize_word_for_tracking(word: str) -> str:
base = word.lower().replace('ſ', 's')
simplified = base
for char, replacement in [
('â', 'a'), ('î', 'i'), ('û', 'u'), ('ô', 'o'),
('ä', 'a'), ('ë', 'e'), ('ï', 'i'), ('ö', 'o'), ('ü', 'u'),
('é', 'e'), ('è', 'e'), ('ê', 'e'), ('à', 'a'),
]:
simplified = simplified.replace(char, replacement)
return simplified
orig_occurrences = {}
for i, word in enumerate(original_words):
norm_word = normalize_word_for_tracking(word)
if norm_word not in orig_occurrences:
orig_occurrences[norm_word] = []
orig_occurrences[norm_word].append(i)
# Progress: alignment is the slowest step on CPU (O(n*m) difflib)
if progress_callback:
progress_callback(78, 100)
# Progress: about to start character-level alignment (slowest CPU step)
if progress_callback:
progress_callback(79, 100)
norm_to_orig_word = self._align_normalized_to_original_charlevel(
full_text, normalized_full
)
# Progress: alignment done
if progress_callback:
progress_callback(84, 100)
norm_to_orig_occurrence = norm_to_orig_word.copy()
split_group = {}
group_id = 0
i = 0
while i < total_norm:
if i == total_norm - 1:
split_group[i] = group_id
group_id += 1
break
# Only group consecutive normalized words that map to the *same*
# original word (true word splits like "inalten" → "in alten").
# Do NOT group words mapping to adjacent originals — that collapses
# normal line boundaries and piles words onto the first line.
orig_i = norm_to_orig_word[i]
orig_next = norm_to_orig_word[i + 1]
if orig_i == orig_next:
# Collect consecutive words mapping to this original
split_words = []
j = i
while j < total_norm and norm_to_orig_word[j] == orig_i:
split_words.append(normalized_words[j])
j += 1
# Validate: true splits have combined length similar to original
combined = ''.join(split_words)
orig_word = original_words[orig_i] if orig_i < len(original_words) else ''
is_true_split = (
len(split_words) > 1 and
len(combined) >= len(orig_word) * 0.5 and
len(combined) <= len(orig_word) * 1.5
)
if is_true_split:
for k in range(i, j):
split_group[k] = group_id
group_id += 1
i = j
else:
# False split (repeated words) — each gets its own group
split_group[i] = group_id
group_id += 1
i += 1
else:
split_group[i] = group_id
group_id += 1
i += 1
# Progress: split grouping done, starting line assignment
if progress_callback:
progress_callback(85, 100)
norm_to_line = self._assign_lines_by_word_alignment(
lowered_lines if lowered_lines is not None else
(stripped_lines if stripped_lines is not None else texts),
normalized_full
)
# Progress: line assignment done
if progress_callback:
progress_callback(88, 100)
line_word_lists = [[] for _ in range(num_lines)]
for norm_idx in range(total_norm):
target_line = norm_to_line[norm_idx]
line_word_lists[target_line].append(normalized_words[norm_idx])
# Progress: word distribution done
if progress_callback:
progress_callback(90, 100)
empty_lines = [i for i in range(num_lines) if not line_word_lists[i]]
if empty_lines:
for empty_idx in empty_lines:
best_source = None
best_dist = float('inf')
for src_idx, words in enumerate(line_word_lists):
if len(words) > 1:
dist = abs(src_idx - empty_idx)
if dist < best_dist:
best_dist = dist
best_source = src_idx
if best_source is not None:
word = line_word_lists[best_source].pop()
line_word_lists[empty_idx].append(word)
# Progress: starting capitalization/punctuation restoration loop
if progress_callback:
progress_callback(91, 100)
normalized = []
for line_idx in range(num_lines):
line_text = ' '.join(line_word_lists[line_idx])
if preserve_capitalization and line_cap_maps is not None and line_idx < len(line_cap_maps):
cap_ref = lowered_lines[line_idx] if lowered_lines is not None else None
if cap_ref is not None:
line_text = self._restore_capitalization_aligned(
cap_ref, line_text, line_cap_maps[line_idx]
)
if preserve_punctuation and line_punct_maps is not None and line_idx < len(line_punct_maps):
punct_ref = stripped_lines[line_idx] if stripped_lines is not None else None
if punct_ref is not None:
line_text = self._restore_punctuation_aligned(
punct_ref, line_text, line_punct_maps[line_idx]
)
normalized.append(line_text)
# Progress: restoration loop done
if progress_callback:
progress_callback(93, 100)
while len(normalized) < total_texts:
normalized.append('')
# Reinsert empty lines
total_original_lines = len(all_lines)
if empty_line_indices:
full_output = [''] * total_original_lines
for norm_idx, orig_idx in enumerate(non_empty_to_original):
if norm_idx < len(normalized) and orig_idx < total_original_lines:
full_output[orig_idx] = normalized[norm_idx]
for norm_idx in range(len(non_empty_to_original), len(normalized)):
full_output.append(normalized[norm_idx])
output_lines = full_output
else:
output_lines = normalized
# Progress: empty line reinsertion done
if progress_callback:
progress_callback(95, 100)
result = '\n'.join(output_lines)
# Progress: join done, starting post-processing regex
if progress_callback:
progress_callback(96, 100)
# Post-processing: common apocopes (before en-attachment so "en wile" → "en wil" → "enwil")
if common_apocopes:
result = _common_apocopes(result)
# Post-processing: attach negation proclitic (after apocopes)
if attach_en_proclitic:
result = _attach_en_proclitic(result)
# Post-processing: lenition of t after n before vowel
if lenition_t_after_n:
result = _lenition_t_after_n_before_vowel(result)
# Post-processing: lenition of t after l before vowel
if lenition_t_after_l:
result = _lenition_t_after_l_before_vowel(result)
# Post-processing: niet → niht
if niet_to_niht:
result = _niet_to_niht(result)
# Progress: all done
if progress_callback:
progress_callback(100, 100)
return result
def normalize_and_compare(
self,
texts: List[str],
reference_texts: Optional[List[str]] = None,
num_samples: int = 5,
**kwargs
):
"""
Normalize texts and optionally compare with references.
Args:
texts: List of input texts
reference_texts: Optional list of reference normalized texts
num_samples: Number of samples to display
**kwargs: Additional arguments for normalize_text
"""
print(f"\n{'='*60}")
print(f"Normalizing {len(texts)} texts")
print(f"{'='*60}\n")
normalized = self.normalize_texts(texts, **kwargs)
# Print samples
num_samples = min(num_samples, len(texts))
for i in range(num_samples):
print(f"Sample {i + 1}:")
print(f" Original: {texts[i]}")
print(f" Normalized: {normalized[i]}")
if reference_texts:
print(f" Reference: {reference_texts[i]}")
match = normalized[i] == reference_texts[i]
print(f" Match: {match}")
print()
return normalized
def main():
parser = argparse.ArgumentParser(description='Normalize Middle High German text')
# Model arguments
parser.add_argument('--model', type=str,
default='/home/jonas/normaere/checkpoints/best_model',
help='Path to trained model')
parser.add_argument('--config', type=str, default='/home/jonas/normaere/config.yaml',
help='Path to configuration file')
# Input arguments
parser.add_argument('--text', type=str, help='Single text to normalize')
parser.add_argument('--input_file', type=str, help='Input file with texts (one per line)')
parser.add_argument('--output_file', type=str, help='Output file for normalized texts')
# Generation arguments
parser.add_argument('--max_length', type=int, default=512,
help='Maximum generation length')
parser.add_argument('--num_beams', type=int, default=5,
help='Number of beams')
parser.add_argument('--length_penalty', type=float, default=0.6,
help='Length penalty')
parser.add_argument('--do_sample', action='store_true',
help='Use sampling instead of beam search')
parser.add_argument('--top_p', type=float, default=None,
help='Nucleus sampling threshold (0.0 to 1.0)')
parser.add_argument('--temperature', type=float, default=None,
help='Temperature for sampling (0.1 to 2.0, lower = more deterministic, higher = more random)')
parser.add_argument('--batch_size', type=int, default=8,
help='Batch size for file processing')
# Stride arguments
parser.add_argument('--stride', action='store_true', default=True,
help='Use sliding window stride for long texts (default: True)')
parser.add_argument('--no-stride', action='store_false', dest='stride',
help='Disable stride; truncate long texts instead')
parser.add_argument('--stride_window_tokens', type=int, default=None,
help='Max input tokens per stride window (default: 400)')
parser.add_argument('--stride_step_tokens', type=int, default=None,
help='Token advance between windows; overlap = window - step (default: 300)')
# Across-lines arguments
parser.add_argument('--across-lines', action='store_true', default=True,
help='Join all lines into one continuous text for normalization, '
'utilizing full model capacity and stride across line boundaries '
'(default: True)')
parser.add_argument('--no-across-lines', action='store_false', dest='across_lines',
help='Normalize each line independently (original behavior)')
# Punctuation preservation arguments
parser.add_argument('--preserve-punctuation', action='store_true', default=True,
help='Strip editorial punctuation before normalization and '
'restore it afterwards (default: True)')
parser.add_argument('--no-preserve-punctuation', action='store_false', dest='preserve_punctuation',
help='Do not strip/restore punctuation; pass text as-is to the model')
# Capitalization preservation arguments
parser.add_argument('--preserve-capitalization', action='store_true', default=True,
help='Record capitalization before normalization (lowercasing input) '
'and restore it afterwards (default: True)')
parser.add_argument('--no-preserve-capitalization', action='store_false', dest='preserve_capitalization',
help='Do not preserve capitalization; input is lowercased without restoration')
args = parser.parse_args()
# Initialize normalizer
normalizer = MHGNormalizer(args.model, args.config)
# Build stride kwargs
stride_kwargs = {
'stride': args.stride,
'stride_window_tokens': args.stride_window_tokens,
'stride_step_tokens': args.stride_step_tokens,
}
if args.text:
# Normalize single text
print("\nOriginal text:")
print(args.text)
print("\nNormalized text:")
normalized = normalizer.normalize_text(
args.text,
max_length=args.max_length,
num_beams=args.num_beams,
length_penalty=args.length_penalty,
do_sample=args.do_sample,
top_p=args.top_p,
temperature=args.temperature,
preserve_punctuation=args.preserve_punctuation,
preserve_capitalization=args.preserve_capitalization,
**stride_kwargs
)
print(normalized)
print()
elif args.input_file:
# Normalize file
if not args.output_file:
# Create output path based on input
input_path = Path(args.input_file)
args.output_file = input_path.parent / f"{input_path.stem}_normalized{input_path.suffix}"
normalizer.normalize_file(
args.input_file,
args.output_file,
batch_size=args.batch_size,
across_lines=args.across_lines,
preserve_punctuation=args.preserve_punctuation,
preserve_capitalization=args.preserve_capitalization,
max_length=args.max_length,
num_beams=args.num_beams,
length_penalty=args.length_penalty,
do_sample=args.do_sample,
top_p=args.top_p,
temperature=args.temperature,
**stride_kwargs
)
else:
# Interactive mode
print("\nInteractive mode. Enter MHG text to normalize (Ctrl+D to exit):")
print("="*60 + "\n")
try:
texts = []
while True:
text = input("> ")
if text.strip():
texts.append(text)
except EOFError:
pass
if texts:
print(f"\nNormalizing {len(texts)} texts...\n")
normalized = normalizer.normalize_texts(
texts,
preserve_punctuation=args.preserve_punctuation,
preserve_capitalization=args.preserve_capitalization,
**stride_kwargs
)
for i, (orig, norm) in enumerate(zip(texts, normalized)):
print(f"\nText {i + 1}:")
print(f" Original: {orig}")
print(f" Normalized: {norm}")
if __name__ == '__main__':
main()
|