Skip to content

Commit ab8a034

Browse files
authored
refactor: share rule instances to eliminate redundant FST construction (#333)
Each rule class now accepts optional dependency params (e.g. cardinal=None) and reuses injected instances instead of creating its own. Normalizer and InverseNormalizer build all rules once in topological order via a new build_tagger_and_verbalizer() method, sharing instances across tagger and verbalizer construction. Before: Cardinal was instantiated 60 times per English Normalizer build. After: Cardinal is instantiated once. Performance improvement: - English Normalizer build: 71s -> 17s (4x faster) - Full test suite: 196s -> 88s (2.2x faster)
1 parent b1c7af4 commit ab8a034

37 files changed

Lines changed: 295 additions & 295 deletions

itn/chinese/inverse_normalizer.py

Lines changed: 9 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -49,27 +49,27 @@ def __init__(
4949
cache_dir = files("itn")
5050
self.build_fst("zh_itn", cache_dir, overwrite_cache)
5151

52-
def build_tagger(self):
52+
def build_tagger_and_verbalizer(self):
53+
cardinal = Cardinal(self.convert_number, self.enable_0_to_9, self.enable_million)
54+
5355
tagger = (
5456
add_weight(Date().tagger, 1.02)
5557
| add_weight(Whitelist().tagger, 1.01)
56-
| add_weight(Fraction().tagger, 1.05)
57-
| add_weight(Measure(enable_0_to_9=self.enable_0_to_9).tagger, 1.05)
58-
| add_weight(Money(enable_0_to_9=self.enable_0_to_9).tagger, 1.04)
58+
| add_weight(Fraction(cardinal=cardinal).tagger, 1.05)
59+
| add_weight(Measure(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal).tagger, 1.05)
60+
| add_weight(Money(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal).tagger, 1.04)
5961
| add_weight(Time().tagger, 1.05)
60-
| add_weight(Cardinal(self.convert_number, self.enable_0_to_9, self.enable_million).tagger, 1.06)
61-
| add_weight(Math().tagger, 1.10)
62+
| add_weight(cardinal.tagger, 1.06)
63+
| add_weight(Math(cardinal=cardinal).tagger, 1.10)
6264
| add_weight(LicensePlate().tagger, 1.0)
6365
| add_weight(Char().tagger, 100)
6466
).optimize()
6567

6668
tagger = tagger.star
67-
# remove the last space
6869
self.tagger = tagger @ self.build_rule(delete(" "), "", "[EOS]")
6970

70-
def build_verbalizer(self):
7171
verbalizer = (
72-
Cardinal(self.convert_number, self.enable_0_to_9, self.enable_million).verbalizer
72+
cardinal.verbalizer
7373
| Char().verbalizer
7474
| Date().verbalizer
7575
| Fraction().verbalizer

itn/chinese/rules/fraction.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,13 +22,14 @@
2222

2323
class Fraction(Processor):
2424

25-
def __init__(self):
25+
def __init__(self, cardinal=None):
2626
super().__init__(name="fraction")
27+
self.cardinal = cardinal or Cardinal()
2728
self.build_tagger()
2829
self.build_verbalizer()
2930

3031
def build_tagger(self):
31-
number = Cardinal().number
32+
number = self.cardinal.number
3233
sign = string_file(get_abs_path("../itn/chinese/data/number/sign.tsv")) # + -
3334

3435
# NOTE(xcsong): default weight = 1.0, set to -1.0 means higher priority

itn/chinese/rules/math.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,15 +22,16 @@
2222

2323
class Math(Processor):
2424

25-
def __init__(self):
25+
def __init__(self, cardinal=None):
2626
super().__init__(name="math")
27+
self.cardinal = cardinal or Cardinal()
2728
self.build_tagger()
2829
self.build_verbalizer()
2930

3031
def build_tagger(self):
3132
operator = string_file(get_abs_path("../itn/chinese/data/math/operator.tsv"))
3233

33-
number = Cardinal().number
34+
number = self.cardinal.number
3435
tagger = number + (operator + number).plus
3536
tagger = insert('value: "') + tagger + insert('"')
3637
self.tagger = self.add_tokens(tagger)

itn/chinese/rules/measure.py

Lines changed: 5 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -22,10 +22,11 @@
2222

2323
class Measure(Processor):
2424

25-
def __init__(self, exclude_one=True, enable_0_to_9=True):
25+
def __init__(self, exclude_one=True, enable_0_to_9=True, cardinal=None):
2626
super().__init__(name="measure")
2727
self.exclude_one = exclude_one
2828
self.enable_0_to_9 = enable_0_to_9
29+
self.cardinal = cardinal or Cardinal()
2930
self.build_tagger()
3031
self.build_verbalizer()
3132

@@ -43,15 +44,15 @@ def build_tagger(self):
4344
add_weight(units_en, -1.0)
4445
)
4546

46-
number = Cardinal().number if self.enable_0_to_9 else Cardinal().number_exclude_0_to_9
47+
number = self.cardinal.number if self.enable_0_to_9 else self.cardinal.number_exclude_0_to_9
4748
# 百分之三十, 百分三十, 百分之百,百分之三十到四十, 百分之三十到百分之五十五
4849
percent = (
4950
(sign + delete("的").ques).ques
5051
+ delete("百分")
5152
+ delete("之").ques
5253
+ (
53-
(Cardinal().number + (to + Cardinal().number).ques)
54-
| ((Cardinal().number + to).ques + cross("百", "100"))
54+
(self.cardinal.number + (to + self.cardinal.number).ques)
55+
| ((self.cardinal.number + to).ques + cross("百", "100"))
5556
)
5657
+ insert("%")
5758
)

itn/chinese/rules/money.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,9 +22,10 @@
2222

2323
class Money(Processor):
2424

25-
def __init__(self, enable_0_to_9=True):
25+
def __init__(self, enable_0_to_9=True, cardinal=None):
2626
super().__init__(name="money")
2727
self.enable_0_to_9 = enable_0_to_9
28+
self.cardinal = cardinal or Cardinal()
2829
self.build_tagger()
2930
self.build_verbalizer()
3031

@@ -33,7 +34,7 @@ def build_tagger(self):
3334
symbol = string_file(get_abs_path("../itn/chinese/data/money/symbol.tsv"))
3435
digit = string_file(get_abs_path("../itn/chinese/data/number/digit.tsv")) # 1 ~ 9
3536

36-
number = Cardinal().number if self.enable_0_to_9 else Cardinal().number_exclude_0_to_9
37+
number = self.cardinal.number if self.enable_0_to_9 else self.cardinal.number_exclude_0_to_9
3738
# 七八美元 => $7~8
3839
number |= digit + insert("~") + digit
3940
# 三千三百八十元五毛八分 => ¥3380.58

itn/japanese/inverse_normalizer.py

Lines changed: 37 additions & 30 deletions
Original file line numberDiff line numberDiff line change
@@ -50,40 +50,47 @@ def __init__(
5050
cache_dir = files("itn")
5151
self.build_fst("ja_itn", cache_dir, overwrite_cache)
5252

53-
def build_tagger(self):
53+
def build_tagger_and_verbalizer(self):
5454
processor = PreProcessor(full_to_half=self.full_to_half).processor
55-
56-
cardinal = add_weight(Cardinal(self.convert_number, self.enable_0_to_9, self.enable_million).tagger, 1.06)
57-
char = add_weight(Char().tagger, 100)
58-
date = add_weight(Date().tagger, 1.02)
59-
fraction = add_weight(Fraction().tagger, 1.05)
60-
math = add_weight(Math().tagger, 90)
61-
measure = add_weight(Measure(enable_0_to_9=self.enable_0_to_9).tagger, 1.05)
62-
money = add_weight(Money(enable_0_to_9=self.enable_0_to_9).tagger, 1.04)
63-
ordinal = add_weight(Ordinal().tagger, 1.04)
64-
time = add_weight(Time().tagger, 1.04)
65-
whitelist = add_weight(Whitelist().tagger, 1.01)
55+
cardinal = Cardinal(self.convert_number, self.enable_0_to_9, self.enable_million)
56+
cardinal_million = Cardinal(enable_million=True)
57+
char = Char()
58+
date = Date(cardinal=cardinal)
59+
fraction = Fraction(cardinal=cardinal_million)
60+
math = Math(cardinal=cardinal)
61+
measure = Measure(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal)
62+
money = Money(enable_0_to_9=self.enable_0_to_9, cardinal=cardinal)
63+
ordinal = Ordinal(cardinal=cardinal)
64+
time = Time()
65+
whitelist = Whitelist()
6666

6767
tagger = (
68-
(cardinal | char | date | fraction | math | measure | money | ordinal | time | whitelist).optimize().star
69-
)
68+
add_weight(cardinal.tagger, 1.06)
69+
| add_weight(char.tagger, 100)
70+
| add_weight(date.tagger, 1.02)
71+
| add_weight(fraction.tagger, 1.05)
72+
| add_weight(math.tagger, 90)
73+
| add_weight(measure.tagger, 1.05)
74+
| add_weight(money.tagger, 1.04)
75+
| add_weight(ordinal.tagger, 1.04)
76+
| add_weight(time.tagger, 1.04)
77+
| add_weight(whitelist.tagger, 1.01)
78+
).optimize().star
7079
tagger = (processor @ tagger).star
71-
# remove the last space
7280
self.tagger = tagger @ self.build_rule(delete(" "), "", "[EOS]")
7381

74-
def build_verbalizer(self):
75-
cardinal = Cardinal(self.convert_number, self.enable_0_to_9).verbalizer
76-
char = Char().verbalizer
77-
date = Date().verbalizer
78-
fraction = Fraction().verbalizer
79-
math = Math().verbalizer
80-
measure = Measure().verbalizer
81-
money = Money().verbalizer
82-
ordinal = Ordinal().verbalizer
83-
time = Time().verbalizer
84-
whitelist = Whitelist().verbalizer
85-
86-
verbalizer = cardinal | char | date | fraction | math | measure | money | ordinal | time | whitelist
82+
verbalizer = (
83+
cardinal.verbalizer
84+
| char.verbalizer
85+
| date.verbalizer
86+
| fraction.verbalizer
87+
| math.verbalizer
88+
| measure.verbalizer
89+
| money.verbalizer
90+
| ordinal.verbalizer
91+
| time.verbalizer
92+
| whitelist.verbalizer
93+
)
8794

88-
processor = PostProcessor().processor
89-
self.verbalizer = (verbalizer @ processor).star
95+
postprocessor = PostProcessor().processor
96+
self.verbalizer = (verbalizer @ postprocessor).star

itn/japanese/rules/date.py

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -22,13 +22,14 @@
2222

2323
class Date(Processor):
2424

25-
def __init__(self):
25+
def __init__(self, cardinal=None):
2626
super().__init__(name="date")
27+
self.cardinal = cardinal or Cardinal()
2728
self.build_tagger()
2829
self.build_verbalizer()
2930

3031
def build_tagger(self):
31-
cardinal = Cardinal().ten_thousand_minus
32+
cardinal = self.cardinal.ten_thousand_minus
3233
day = string_file(get_abs_path("../itn/japanese/data/date/day.tsv"))
3334
month = string_file(get_abs_path("../itn/japanese/data/date/month.tsv"))
3435
to = cross("から", "〜")

itn/japanese/rules/fraction.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -22,14 +22,15 @@
2222

2323
class Fraction(Processor):
2424

25-
def __init__(self):
25+
def __init__(self, cardinal=None):
2626
super().__init__(name="fraction")
27+
self.cardinal = cardinal or Cardinal(enable_million=True)
2728
self.build_tagger()
2829
self.build_verbalizer()
2930

3031
def build_tagger(self):
31-
cardinal = Cardinal(enable_million=True).number
32-
decimal = Cardinal(enable_million=True).decimal
32+
cardinal = self.cardinal.number
33+
decimal = self.cardinal.decimal
3334
sign = string_file(get_abs_path("../itn/japanese/data/number/sign.tsv"))
3435
sign = insert('sign: "') + sign + insert('"')
3536

itn/japanese/rules/math.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -22,16 +22,17 @@
2222

2323
class Math(Processor):
2424

25-
def __init__(self):
25+
def __init__(self, cardinal=None):
2626
super().__init__(name="math")
27+
self.cardinal = cardinal or Cardinal()
2728
self.build_tagger()
2829
self.build_verbalizer()
2930

3031
def build_tagger(self):
3132
operator = string_file(get_abs_path("../itn/japanese/data/math/operator.tsv"))
3233

33-
number = Cardinal().big_integer
34-
decimal = Cardinal().decimal
34+
number = self.cardinal.big_integer
35+
decimal = self.cardinal.decimal
3536
number |= decimal
3637
tagger = number + (operator + number).plus
3738
tagger = insert('value: "') + tagger + insert('"')

itn/japanese/rules/measure.py

Lines changed: 4 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -22,18 +22,19 @@
2222

2323
class Measure(Processor):
2424

25-
def __init__(self, enable_0_to_9=True):
25+
def __init__(self, enable_0_to_9=True, cardinal=None):
2626
super().__init__(name="measure")
2727
self.enable_0_to_9 = enable_0_to_9
28+
self.cardinal = cardinal or Cardinal()
2829
self.build_tagger()
2930
self.build_verbalizer()
3031

3132
def build_tagger(self):
3233
unit_en = string_file(get_abs_path("../itn/japanese/data/measure/unit_en.tsv"))
3334
unit_ja = string_file(get_abs_path("../itn/japanese/data/measure/unit_ja.tsv"))
3435

35-
cardinal = Cardinal().number if self.enable_0_to_9 else Cardinal().number_exclude_0_to_9
36-
decimal = Cardinal().decimal
36+
cardinal = self.cardinal.number if self.enable_0_to_9 else self.cardinal.number_exclude_0_to_9
37+
decimal = self.cardinal.decimal
3738

3839
suffix = (
3940
insert("/")

0 commit comments

Comments
 (0)