From 860697f2bf40494c06cbdc579e92126bb36142b5 Mon Sep 17 00:00:00 2001 From: pengzhendong <275331498@qq.com> Date: Tue, 9 Jun 2026 20:05:11 +0800 Subject: [PATCH] refactor: code review cleanup and bug fixes MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 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 --- itn/chinese/inverse_normalizer.py | 50 ++++++++++++++++++------------- tn/chinese/normalizer.py | 4 +-- tn/chinese/rules/cardinal.py | 11 ++++++- tn/chinese/rules/measure.py | 2 +- tn/japanese/normalizer.py | 1 - tn/processor.py | 6 ++-- 6 files changed, 46 insertions(+), 28 deletions(-) diff --git a/itn/chinese/inverse_normalizer.py b/itn/chinese/inverse_normalizer.py index eb6249f..3cccba9 100644 --- a/itn/chinese/inverse_normalizer.py +++ b/itn/chinese/inverse_normalizer.py @@ -52,19 +52,29 @@ def __init__( def build_tagger_and_verbalizer(self): cardinal = Cardinal(self.convert_number, self.enable_0_to_9, self.enable_million) + char = Char() + date = Date() + fraction = Fraction(cardinal=cardinal) + train_number = TrainNumber() + math = Math(cardinal=cardinal) + measure = Measure(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal) + money = Money(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal) + time = Time() + license_plate = LicensePlate() + whitelist = Whitelist() tagger = ( - add_weight(Date().tagger, 1.02) - | add_weight(Whitelist().tagger, 1.01) - | add_weight(Fraction(cardinal=cardinal).tagger, 1.05) - | add_weight(Measure(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal).tagger, 1.05) - | add_weight(Money(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal).tagger, 1.04) - | add_weight(Time().tagger, 1.05) + add_weight(date.tagger, 1.02) + | add_weight(whitelist.tagger, 1.01) + | add_weight(fraction.tagger, 1.05) + | add_weight(measure.tagger, 1.05) + | add_weight(money.tagger, 1.04) + | add_weight(time.tagger, 1.05) | add_weight(cardinal.tagger, 1.06) - | add_weight(Math(cardinal=cardinal).tagger, 1.10) - | add_weight(LicensePlate().tagger, 1.0) - | add_weight(TrainNumber().tagger, 1.0) - | add_weight(Char().tagger, 100) + | add_weight(math.tagger, 1.10) + | add_weight(license_plate.tagger, 1.0) + | add_weight(train_number.tagger, 1.0) + | add_weight(char.tagger, 100) ).optimize() tagger = tagger.star @@ -72,16 +82,16 @@ def build_tagger_and_verbalizer(self): verbalizer = ( cardinal.verbalizer - | Char().verbalizer - | Date().verbalizer - | Fraction().verbalizer - | TrainNumber().verbalizer - | Math().verbalizer - | Measure(enable_0_to_9=self.enable_0_to_9).verbalizer - | Money(enable_0_to_9=self.enable_0_to_9).verbalizer - | Time().verbalizer - | LicensePlate().verbalizer - | Whitelist().verbalizer + | char.verbalizer + | date.verbalizer + | fraction.verbalizer + | train_number.verbalizer + | math.verbalizer + | measure.verbalizer + | money.verbalizer + | time.verbalizer + | license_plate.verbalizer + | whitelist.verbalizer ).optimize() postprocessor = PostProcessor(remove_interjections=self.remove_interjections).processor diff --git a/tn/chinese/normalizer.py b/tn/chinese/normalizer.py index 53c66b1..e3845ce 100644 --- a/tn/chinese/normalizer.py +++ b/tn/chinese/normalizer.py @@ -58,7 +58,7 @@ def build_tagger_and_verbalizer(self): processor = PreProcessor(traditional_to_simple=self.traditional_to_simple).processor cardinal = Cardinal() date = Date() - whitelist = Whitelist() + whitelist = Whitelist(remove_erhua=self.remove_erhua) sport = Sport(cardinal=cardinal) fraction = Fraction(cardinal=cardinal) measure = Measure(cardinal=cardinal) @@ -92,7 +92,7 @@ def build_tagger_and_verbalizer(self): | money.verbalizer | sport.verbalizer | time.verbalizer - | Whitelist(remove_erhua=self.remove_erhua).verbalizer + | whitelist.verbalizer ).optimize() postprocessor = PostProcessor( diff --git a/tn/chinese/rules/cardinal.py b/tn/chinese/rules/cardinal.py index 7070f05..f747884 100644 --- a/tn/chinese/rules/cardinal.py +++ b/tn/chinese/rules/cardinal.py @@ -64,7 +64,16 @@ def build_tagger(self): hundred_million = ( (thousand | hundred | ten | digit) + insert("亿") - + (four_nonzero + insert("万") + four_any | rmzero**4 + four_nonzero | rmzero**8) + + ( + four_nonzero + insert("万") + four_any + | rmzero**4 + ( + (zero + hundred) + | (rmzero + zero + tens) + | (rmzero**2 + zero + digit) + | (insert("零") + thousand) + ) + | rmzero**8 + ) ) number = digits | ten | hundred | thousand | ten_thousand | hundred_million diff --git a/tn/chinese/rules/measure.py b/tn/chinese/rules/measure.py index cac5bf1..e6a3710 100644 --- a/tn/chinese/rules/measure.py +++ b/tn/chinese/rules/measure.py @@ -36,7 +36,7 @@ def build_tagger(self): to = cross("-", "到") | cross("~", "到") | accep("到") number = self.cardinal.number - strip_comma = self.build_rule(delete(","), self.DIGIT, self.DIGIT) + strip_comma = self.build_rule(delete(",") | delete(","), self.DIGIT, self.DIGIT) number = strip_comma @ number number @= self.build_rule(cross("二", "两"), "[BOS]", "[EOS]") # 1-11个,1个-11个 diff --git a/tn/japanese/normalizer.py b/tn/japanese/normalizer.py index fd141ef..999cdc7 100644 --- a/tn/japanese/normalizer.py +++ b/tn/japanese/normalizer.py @@ -84,7 +84,6 @@ def build_tagger_and_verbalizer(self): tagger = (processor @ tagger).star self.tagger = tagger @ self.build_rule(delete(" "), r="[EOS]") - transliteration = Transliteration() verbalizer = ( cardinal.verbalizer | char.verbalizer diff --git a/tn/processor.py b/tn/processor.py index a125b67..d7fa698 100644 --- a/tn/processor.py +++ b/tn/processor.py @@ -18,6 +18,9 @@ from pynini import Fst, cdrewrite, cross, difference, escape, invert, shortestpath, union from pynini.lib import byte, utf8 +from pynini.lib.pynutil import delete, insert + +from tn.token_parser import TokenParser logger = logging.getLogger("wetext") if not logger.handlers: @@ -25,9 +28,6 @@ handler.setFormatter(logging.Formatter("%(asctime)s WETEXT %(levelname)s %(message)s")) logger.addHandler(handler) logger.setLevel(logging.INFO) -from pynini.lib.pynutil import delete, insert - -from tn.token_parser import TokenParser class Processor: