Skip to content

Commit 55207ec

Browse files
authored
refactor: code review cleanup and bug fixes (#350)
- Fix import ordering in processor.py - Fix 亿级 bug: 100001000 => 一亿零一千 (was 一亿一千) - Share rule instances in Chinese ITN verbalizer (build time -15%) - Don't create Transliteration when transliterate=False in Japanese TN - Share Whitelist instance in Chinese TN normalizer - Support full-width comma in Measure strip_comma
1 parent aa078d0 commit 55207ec

6 files changed

Lines changed: 46 additions & 28 deletions

File tree

itn/chinese/inverse_normalizer.py

Lines changed: 30 additions & 20 deletions
Original file line numberDiff line numberDiff line change
@@ -52,36 +52,46 @@ def __init__(
5252

5353
def build_tagger_and_verbalizer(self):
5454
cardinal = Cardinal(self.convert_number, self.enable_0_to_9, self.enable_million)
55+
char = Char()
56+
date = Date()
57+
fraction = Fraction(cardinal=cardinal)
58+
train_number = TrainNumber()
59+
math = Math(cardinal=cardinal)
60+
measure = Measure(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal)
61+
money = Money(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal)
62+
time = Time()
63+
license_plate = LicensePlate()
64+
whitelist = Whitelist()
5565

5666
tagger = (
57-
add_weight(Date().tagger, 1.02)
58-
| add_weight(Whitelist().tagger, 1.01)
59-
| add_weight(Fraction(cardinal=cardinal).tagger, 1.05)
60-
| add_weight(Measure(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal).tagger, 1.05)
61-
| add_weight(Money(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal).tagger, 1.04)
62-
| add_weight(Time().tagger, 1.05)
67+
add_weight(date.tagger, 1.02)
68+
| add_weight(whitelist.tagger, 1.01)
69+
| add_weight(fraction.tagger, 1.05)
70+
| add_weight(measure.tagger, 1.05)
71+
| add_weight(money.tagger, 1.04)
72+
| add_weight(time.tagger, 1.05)
6373
| add_weight(cardinal.tagger, 1.06)
64-
| add_weight(Math(cardinal=cardinal).tagger, 1.10)
65-
| add_weight(LicensePlate().tagger, 1.0)
66-
| add_weight(TrainNumber().tagger, 1.0)
67-
| add_weight(Char().tagger, 100)
74+
| add_weight(math.tagger, 1.10)
75+
| add_weight(license_plate.tagger, 1.0)
76+
| add_weight(train_number.tagger, 1.0)
77+
| add_weight(char.tagger, 100)
6878
).optimize()
6979

7080
tagger = tagger.star
7181
self.tagger = tagger @ self.build_rule(delete(" "), "", "[EOS]")
7282

7383
verbalizer = (
7484
cardinal.verbalizer
75-
| Char().verbalizer
76-
| Date().verbalizer
77-
| Fraction().verbalizer
78-
| TrainNumber().verbalizer
79-
| Math().verbalizer
80-
| Measure(enable_0_to_9=self.enable_0_to_9).verbalizer
81-
| Money(enable_0_to_9=self.enable_0_to_9).verbalizer
82-
| Time().verbalizer
83-
| LicensePlate().verbalizer
84-
| Whitelist().verbalizer
85+
| char.verbalizer
86+
| date.verbalizer
87+
| fraction.verbalizer
88+
| train_number.verbalizer
89+
| math.verbalizer
90+
| measure.verbalizer
91+
| money.verbalizer
92+
| time.verbalizer
93+
| license_plate.verbalizer
94+
| whitelist.verbalizer
8595
).optimize()
8696
postprocessor = PostProcessor(remove_interjections=self.remove_interjections).processor
8797

tn/chinese/normalizer.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -58,7 +58,7 @@ def build_tagger_and_verbalizer(self):
5858
processor = PreProcessor(traditional_to_simple=self.traditional_to_simple).processor
5959
cardinal = Cardinal()
6060
date = Date()
61-
whitelist = Whitelist()
61+
whitelist = Whitelist(remove_erhua=self.remove_erhua)
6262
sport = Sport(cardinal=cardinal)
6363
fraction = Fraction(cardinal=cardinal)
6464
measure = Measure(cardinal=cardinal)
@@ -92,7 +92,7 @@ def build_tagger_and_verbalizer(self):
9292
| money.verbalizer
9393
| sport.verbalizer
9494
| time.verbalizer
95-
| Whitelist(remove_erhua=self.remove_erhua).verbalizer
95+
| whitelist.verbalizer
9696
).optimize()
9797

9898
postprocessor = PostProcessor(

tn/chinese/rules/cardinal.py

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -64,7 +64,16 @@ def build_tagger(self):
6464
hundred_million = (
6565
(thousand | hundred | ten | digit)
6666
+ insert("亿")
67-
+ (four_nonzero + insert("万") + four_any | rmzero**4 + four_nonzero | rmzero**8)
67+
+ (
68+
four_nonzero + insert("万") + four_any
69+
| rmzero**4 + (
70+
(zero + hundred)
71+
| (rmzero + zero + tens)
72+
| (rmzero**2 + zero + digit)
73+
| (insert("零") + thousand)
74+
)
75+
| rmzero**8
76+
)
6877
)
6978

7079
number = digits | ten | hundred | thousand | ten_thousand | hundred_million

tn/chinese/rules/measure.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -36,7 +36,7 @@ def build_tagger(self):
3636
to = cross("-", "到") | cross("~", "到") | accep("到")
3737

3838
number = self.cardinal.number
39-
strip_comma = self.build_rule(delete(","), self.DIGIT, self.DIGIT)
39+
strip_comma = self.build_rule(delete(",") | delete(","), self.DIGIT, self.DIGIT)
4040
number = strip_comma @ number
4141
number @= self.build_rule(cross("二", "两"), "[BOS]", "[EOS]")
4242
# 1-11个,1个-11个

tn/japanese/normalizer.py

Lines changed: 0 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -84,7 +84,6 @@ def build_tagger_and_verbalizer(self):
8484
tagger = (processor @ tagger).star
8585
self.tagger = tagger @ self.build_rule(delete(" "), r="[EOS]")
8686

87-
transliteration = Transliteration()
8887
verbalizer = (
8988
cardinal.verbalizer
9089
| char.verbalizer

tn/processor.py

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -18,16 +18,16 @@
1818

1919
from pynini import Fst, cdrewrite, cross, difference, escape, invert, shortestpath, union
2020
from pynini.lib import byte, utf8
21+
from pynini.lib.pynutil import delete, insert
22+
23+
from tn.token_parser import TokenParser
2124

2225
logger = logging.getLogger("wetext")
2326
if not logger.handlers:
2427
handler = logging.StreamHandler()
2528
handler.setFormatter(logging.Formatter("%(asctime)s WETEXT %(levelname)s %(message)s"))
2629
logger.addHandler(handler)
2730
logger.setLevel(logging.INFO)
28-
from pynini.lib.pynutil import delete, insert
29-
30-
from tn.token_parser import TokenParser
3131

3232

3333
class Processor:

0 commit comments

Comments
 (0)