Skip to content

Commit 562de68

Browse files
committed
refactor: fix encoding and move str2bool to shared utils
- Add encoding="utf-8" to open() in tn/main.py and itn/main.py for Windows compatibility - Move str2bool from itn/main.py to tn/utils.py to remove TN's dependency on ITN module
1 parent 06575fa commit 562de68

3 files changed

Lines changed: 14 additions & 14 deletions

File tree

itn/main.py

Lines changed: 2 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -14,19 +14,9 @@
1414

1515
import argparse
1616

17-
# TODO(xcsong): multi-language support
1817
from itn.chinese.inverse_normalizer import InverseNormalizer as ZhInverseNormalizer
1918
from itn.japanese.inverse_normalizer import InverseNormalizer as JaInverseNormalizer
20-
21-
22-
def str2bool(s, default=False):
23-
s = s.lower()
24-
if s == "true":
25-
return True
26-
elif s == "false":
27-
return False
28-
else:
29-
return default
19+
from tn.utils import str2bool
3020

3121

3222
def main():
@@ -62,7 +52,7 @@ def main():
6252
print(normalizer.tag(args.text))
6353
print(normalizer.normalize(args.text))
6454
elif args.file:
65-
with open(args.file) as fin:
55+
with open(args.file, encoding="utf-8") as fin:
6656
for line in fin:
6757
print(normalizer.tag(line.strip()))
6858
print(normalizer.normalize(line.strip()))

tn/main.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -14,7 +14,7 @@
1414

1515
import argparse
1616

17-
from itn.main import str2bool
17+
from tn.utils import str2bool
1818

1919
# TODO(pzd17 & sxc19): multi-language support
2020
from tn.chinese.normalizer import Normalizer as ZhNormalizer
@@ -63,7 +63,7 @@ def main():
6363
print(normalizer.tag(args.text))
6464
print(normalizer.normalize(args.text))
6565
elif args.file:
66-
with open(args.file) as fin:
66+
with open(args.file, encoding="utf-8") as fin:
6767
for line in fin:
6868
print(normalizer.tag(line.strip()))
6969
print(normalizer.normalize(line.strip()))

tn/utils.py

Lines changed: 10 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -82,3 +82,13 @@ def get_formats(input_f, input_case="cased", is_default=True):
8282

8383
multiple_formats = pynini.string_map(multiple_formats)
8484
return multiple_formats
85+
86+
87+
def str2bool(s, default=False):
88+
s = s.lower()
89+
if s == "true":
90+
return True
91+
elif s == "false":
92+
return False
93+
else:
94+
return default

0 commit comments

Comments
 (0)