Skip to content

Commit 6bc6946

Browse files
committed
refactor: harden TN and ITN processing
1 parent b0ce1f9 commit 6bc6946

19 files changed

Lines changed: 349 additions & 64 deletions

.github/workflows/wheels.yml

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -32,12 +32,15 @@ jobs:
3232
python -m tn --text "2.5平方电线" --overwrite_cache --language "zh"
3333
python -m tn --text "2010-03-21" --overwrite_cache --language "en"
3434
python -m itn --text "二点五平方电线" --overwrite_cache
35+
python -m itn --text "one two three" --overwrite_cache --language "en"
3536
3637
- name: Prepare Graph
3738
run: |
3839
mkdir graph
3940
cp tn/*.fst graph
4041
cp itn/*.fst graph
42+
cp tn/*_cache.json graph
43+
cp itn/*_cache.json graph
4144
4245
- name: Get version from setuptools_scm
4346
id: scm_version

.gitignore

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -33,6 +33,7 @@ __pycache__/
3333

3434
# Clangd files
3535
.cache
36+
.venv/
3637
compile_commands.json
3738

3839

README.md

Lines changed: 13 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -34,22 +34,26 @@ weitn --text "二点五平方电线"
3434
Python usage:
3535

3636
```py
37-
from itn.chinese.inverse_normalizer import InverseNormalizer
38-
from tn.chinese.normalizer import Normalizer as ZhNormalizer
39-
from tn.english.normalizer import Normalizer as EnNormalizer
40-
41-
# NOTE(xcsong): 和默认参数不一致时,必须重新构图,要重新构图请务必指定 `overwrite_cache=True`
42-
# When the parameters differ from the defaults, it is mandatory to re-compose. To re-compose, please ensure you specify `overwrite_cache=True`.
37+
from itn.chinese.inverse_normalizer import InverseNormalizer
38+
from itn.english.inverse_normalizer import InverseNormalizer as EnInverseNormalizer
39+
from tn.chinese.normalizer import Normalizer as ZhNormalizer
40+
from tn.english.normalizer import Normalizer as EnNormalizer
41+
42+
# FST 缓存会记录构建参数及规则数据指纹;配置或规则变化时会自动重新构图。
43+
# Set `overwrite_cache=True` only when an unconditional rebuild is required.
4344

4445
zh_tn_text = "你好 WeTextProcessing 1.0,船新版本儿,船新体验儿,简直666,9和10"
4546
zh_itn_text = "你好 WeTextProcessing 一点零,船新版本儿,船新体验儿,简直六六六,九和六"
46-
en_tn_text = "Hello WeTextProcessing 1.0, life is short, just use wetext, 666, 9 and 10"
47+
en_tn_text = "Hello WeTextProcessing 1.0, life is short, just use wetext, 666, 9 and 10"
48+
en_itn_text = "call me at five five five one two three four"
4749
zh_tn_model = ZhNormalizer(remove_erhua=True, overwrite_cache=True)
4850
zh_itn_model = InverseNormalizer(enable_0_to_9=False, overwrite_cache=True)
49-
en_tn_model = EnNormalizer(overwrite_cache=True)
51+
en_tn_model = EnNormalizer(overwrite_cache=True)
52+
en_itn_model = EnInverseNormalizer(overwrite_cache=True)
5053
print("中文 TN (去除儿化音,重新在线构图):\n\t{} => {}".format(zh_tn_text, zh_tn_model.normalize(zh_tn_text)))
5154
print("中文ITN (小于10的单独数字不转换,重新在线构图):\n\t{} => {}".format(zh_itn_text, zh_itn_model.normalize(zh_itn_text)))
52-
print("英文 TN (暂时还没有可控的选项,后面会加...):\n\t{} => {}\n".format(en_tn_text, en_tn_model.normalize(en_tn_text)))
55+
print("英文 TN (暂时还没有可控的选项,后面会加...):\n\t{} => {}\n".format(en_tn_text, en_tn_model.normalize(en_tn_text)))
56+
print("英文 ITN:\n\t{} => {}\n".format(en_itn_text, en_itn_model.normalize(en_itn_text)))
5357

5458
zh_tn_model = ZhNormalizer(overwrite_cache=False)
5559
zh_itn_model = InverseNormalizer(overwrite_cache=False)

itn/chinese/inverse_normalizer.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -48,7 +48,17 @@ def __init__(
4848
self.enable_million = enable_million
4949
if cache_dir is None:
5050
cache_dir = files("itn")
51-
self.build_fst("zh_itn", cache_dir, overwrite_cache)
51+
self.build_fst(
52+
"zh_itn",
53+
cache_dir,
54+
overwrite_cache,
55+
{
56+
"enable_0_to_9": self.enable_0_to_9,
57+
"enable_million": self.enable_million,
58+
"enable_standalone_number": self.convert_number,
59+
"remove_interjections": self.remove_interjections,
60+
},
61+
)
5262

5363
def build_tagger_and_verbalizer(self):
5464
cardinal = Cardinal(self.convert_number, self.enable_0_to_9, self.enable_million)

itn/english/inverse_normalizer.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -38,7 +38,7 @@ def __init__(self, cache_dir=None, overwrite_cache=False):
3838
super().__init__(name="en_inverse_normalizer", ordertype="itn")
3939
if cache_dir is None:
4040
cache_dir = files("itn")
41-
self.build_fst("en_itn", cache_dir, overwrite_cache)
41+
self.build_fst("en_itn", cache_dir, overwrite_cache, {})
4242

4343
def build_tagger_and_verbalizer(self):
4444
cardinal = Cardinal()

itn/japanese/inverse_normalizer.py

Lines changed: 11 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -48,7 +48,17 @@ def __init__(
4848
self.enable_million = enable_million
4949
if cache_dir is None:
5050
cache_dir = files("itn")
51-
self.build_fst("ja_itn", cache_dir, overwrite_cache)
51+
self.build_fst(
52+
"ja_itn",
53+
cache_dir,
54+
overwrite_cache,
55+
{
56+
"enable_0_to_9": self.enable_0_to_9,
57+
"enable_million": self.enable_million,
58+
"enable_standalone_number": self.convert_number,
59+
"full_to_half": self.full_to_half,
60+
},
61+
)
5262

5363
def build_tagger_and_verbalizer(self):
5464
processor = PreProcessor(full_to_half=self.full_to_half).processor

itn/main.py

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -15,6 +15,7 @@
1515
import argparse
1616

1717
from itn.chinese.inverse_normalizer import InverseNormalizer as ZhInverseNormalizer
18+
from itn.english.inverse_normalizer import InverseNormalizer as EnInverseNormalizer
1819
from itn.japanese.inverse_normalizer import InverseNormalizer as JaInverseNormalizer
1920
from tn.utils import str2bool
2021

@@ -28,7 +29,7 @@ def main():
2829
parser.add_argument("--enable_standalone_number", type=str, default="True", help="一百 = 100 if True else 一百")
2930
parser.add_argument("--enable_0_to_9", type=str, default="False", help="零和九 = 0和9 if True else 零和九")
3031
parser.add_argument("--enable_million", type=str, default="False", help="六百万 = 6000000 if True else 600万")
31-
parser.add_argument("--language", type=str, default="zh", choices=["zh", "ja"], help="valid languages")
32+
parser.add_argument("--language", type=str, default="zh", choices=["zh", "en", "ja"], help="valid languages")
3233
args = parser.parse_args()
3334

3435
if args.language == "zh":
@@ -39,6 +40,11 @@ def main():
3940
enable_0_to_9=str2bool(args.enable_0_to_9),
4041
enable_million=str2bool(args.enable_million),
4142
)
43+
elif args.language == "en":
44+
normalizer = EnInverseNormalizer(
45+
cache_dir=args.cache_dir,
46+
overwrite_cache=args.overwrite_cache,
47+
)
4248
elif args.language == "ja":
4349
normalizer = JaInverseNormalizer(
4450
cache_dir=args.cache_dir,

pyproject.toml

Lines changed: 9 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -53,5 +53,12 @@ include = ["tn*", "itn*"]
5353
namespaces = false
5454

5555
[tool.setuptools.package-data]
56-
tn = ["*.fst", "chinese/data/*/*.tsv", "english/data/*/*.tsv", "english/data/*.tsv", "english/data/*/*.far", "japanese/data/*/*.tsv"]
57-
itn = ["*.fst", "chinese/data/*/*.tsv", "japanese/data/*/*.tsv"]
56+
tn = ["*.fst", "*_cache.json", "chinese/data/*/*.tsv", "english/data/*/*.tsv", "english/data/*.tsv", "english/data/*/*.far", "japanese/data/*/*.tsv"]
57+
itn = [
58+
"*.fst",
59+
"*_cache.json",
60+
"chinese/data/*/*.tsv",
61+
"english/data/*.tsv",
62+
"english/data/*/*.tsv",
63+
"japanese/data/*/*.tsv",
64+
]

runtime/processor/wetext_processor.cc

Lines changed: 4 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -30,8 +30,12 @@ Processor::Processor(const std::string& tagger_path,
3030
parse_type_ = ParseType::kZH_ITN;
3131
} else if (tagger_path.find("en_tn_") != tagger_path.npos) {
3232
parse_type_ = ParseType::kEN_TN;
33+
} else if (tagger_path.find("en_itn_") != tagger_path.npos) {
34+
parse_type_ = ParseType::kEN_ITN;
3335
} else if (tagger_path.find("ja_tn_") != tagger_path.npos) {
3436
parse_type_ = ParseType::kJA_TN;
37+
} else if (tagger_path.find("ja_itn_") != tagger_path.npos) {
38+
parse_type_ = ParseType::kJA_ITN;
3539
} else {
3640
LOG(FATAL) << "Invalid fst prefix, prefix should contain"
3741
<< " either \"_tn_\" or \"_itn_\".";

runtime/processor/wetext_token_parser.cc

Lines changed: 57 additions & 23 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,8 @@
1414

1515
#include "processor/wetext_token_parser.h"
1616

17+
#include <stdexcept>
18+
1719
#include "utils/wetext_log.h"
1820
#include "utils/wetext_string.h"
1921

@@ -33,24 +35,26 @@ const std::unordered_map<std::string, std::vector<std::string>> ZH_TN_ORDERS = {
3335
{"money", {"value", "currency"}},
3436
{"time", {"noon", "hour", "minute", "second"}}};
3537
const std::unordered_map<std::string, std::vector<std::string>> JA_TN_ORDERS = {
36-
{"date", {"year", "month", "day"}},
37-
{"money", {"value", "currency"}}};
38+
{"date", {"year", "month", "day"}}, {"money", {"value", "currency"}}};
3839

3940
const std::unordered_map<std::string, std::vector<std::string>> EN_TN_ORDERS = {
4041
{"date", {"preserve_order", "text", "day", "month", "year"}},
4142
{"money", {"integer_part", "fractional_part", "quantity", "currency_maj"}}};
42-
const std::unordered_map<std::string, std::vector<std::string>> ZH_ITN_ORDERS =
43-
{{"date", {"year", "month", "day"}},
44-
{"fraction", {"sign", "numerator", "denominator"}},
45-
{"measure", {"numerator", "denominator", "value"}},
46-
{"money", {"currency", "value", "decimal"}},
47-
{"time", {"hour", "minute", "second", "noon"}}};
43+
const std::unordered_map<std::string, std::vector<std::string>> ITN_ORDERS = {
44+
{"date", {"year", "month", "day", "preserve_order"}},
45+
{"fraction", {"sign", "numerator", "denominator"}},
46+
{"measure", {"numerator", "denominator", "value", "units"}},
47+
{"money", {"currency", "value", "decimal", "quantity"}},
48+
{"time", {"hour", "minute", "second", "noon", "zone"}},
49+
{"telephone", {"country_code", "number_part"}},
50+
{"electronic", {"username", "domain", "protocol"}}};
4851

4952
TokenParser::TokenParser(ParseType type) {
5053
if (type == ParseType::kZH_TN) {
5154
orders_ = ZH_TN_ORDERS;
52-
} else if (type == ParseType::kZH_ITN) {
53-
orders_ = ZH_ITN_ORDERS;
55+
} else if (type == ParseType::kZH_ITN || type == ParseType::kEN_ITN ||
56+
type == ParseType::kJA_ITN) {
57+
orders_ = ITN_ORDERS;
5458
} else if (type == ParseType::kEN_TN) {
5559
orders_ = EN_TN_ORDERS;
5660
} else if (type == ParseType::kJA_TN) {
@@ -62,9 +66,12 @@ TokenParser::TokenParser(ParseType type) {
6266

6367
void TokenParser::Load(const std::string& input) {
6468
wetext::SplitUTF8StringToChars(input, &text_);
65-
CHECK_GT(text_.size(), 0);
69+
if (text_.empty()) {
70+
throw std::invalid_argument("token stream must not be empty");
71+
}
6672
index_ = 0;
6773
ch_ = text_[0];
74+
tokens_.clear();
6875
}
6976

7077
bool TokenParser::Read() {
@@ -94,38 +101,54 @@ bool TokenParser::ParseChar(const std::string& exp) {
94101
}
95102

96103
bool TokenParser::ParseChars(const std::string& exp) {
97-
bool ok = false;
104+
size_t start = index_;
98105
std::vector<std::string> chars;
99106
wetext::SplitUTF8StringToChars(exp, &chars);
100107
for (const auto& x : chars) {
101-
ok |= ParseChar(x);
108+
if (!ParseChar(x)) {
109+
index_ = start;
110+
ch_ = text_[start];
111+
return false;
112+
}
102113
}
103-
return ok;
114+
return true;
104115
}
105116

106117
std::string TokenParser::ParseKey() {
107-
CHECK_NE(ch_, EOS);
108-
CHECK_EQ(UTF8_WHITESPACE.count(ch_), 0);
118+
if (ch_ == EOS || UTF8_WHITESPACE.count(ch_) > 0) {
119+
throw std::invalid_argument("expected token key at position " +
120+
std::to_string(index_));
121+
}
109122

110123
std::string key = "";
111124
while (ASCII_LETTERS.count(ch_) > 0) {
112125
key += ch_;
113126
Read();
114127
}
128+
if (key.empty()) {
129+
throw std::invalid_argument("invalid token key at position " +
130+
std::to_string(index_));
131+
}
115132
return key;
116133
}
117134

118135
std::string TokenParser::ParseValue() {
119-
CHECK_NE(ch_, EOS);
120-
bool escape = false;
136+
if (ch_ == EOS) {
137+
throw std::invalid_argument("expected token value at end of stream");
138+
}
121139

122140
std::string value = "";
123141
while (ch_ != "\"") {
142+
if (ch_ == EOS) {
143+
throw std::invalid_argument("unterminated token value");
144+
}
124145
value += ch_;
125-
escape = ch_ == "\\";
146+
bool escape = ch_ == "\\";
126147
Read();
127148
if (escape) {
128-
escape = false;
149+
if (ch_ == EOS) {
150+
throw std::invalid_argument("unterminated escape in token value");
151+
}
129152
value += ch_;
130153
Read();
131154
}
@@ -137,20 +160,31 @@ void TokenParser::Parse(const std::string& input) {
137160
Load(input);
138161
while (ParseWs()) {
139162
std::string name = ParseKey();
140-
ParseChars(" { ");
163+
if (!ParseChars(" { ")) {
164+
throw std::invalid_argument("expected token opening delimiter");
165+
}
141166

142167
Token token(name);
168+
bool closed = false;
143169
while (ParseWs()) {
144170
if (ch_ == "}") {
145171
ParseChar("}");
172+
closed = true;
146173
break;
147174
}
148175
std::string key = ParseKey();
149-
ParseChars(": \"");
176+
if (!ParseChars(": \"")) {
177+
throw std::invalid_argument("expected token field delimiter");
178+
}
150179
std::string value = ParseValue();
151-
ParseChar("\"");
180+
if (!ParseChar("\"")) {
181+
throw std::invalid_argument("expected closing quote");
182+
}
152183
token.Append(key, value);
153184
}
185+
if (!closed) {
186+
throw std::invalid_argument("unterminated token " + name);
187+
}
154188
tokens_.emplace_back(token);
155189
}
156190
}

0 commit comments

Comments
 (0)