From a974ff29e415f161ce62dcb8837f54cf2a843846 Mon Sep 17 00:00:00 2001 From: pengzhendong <275331498@qq.com> Date: Tue, 9 Jun 2026 15:40:45 +0800 Subject: [PATCH] feat: add TrainNumber rule for Chinese ITN (#191) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Convert spoken train numbers to codes, e.g.: 高一二三 => G123 动二三 => D23 特五六七八 => T5678 Requires at least 2 digits to avoid conflicts with 高一/高二/高三 (school grades). --- itn/chinese/data/train_number/prefix.tsv | 7 +++++ itn/chinese/inverse_normalizer.py | 3 ++ itn/chinese/rules/train_number.py | 36 ++++++++++++++++++++++++ itn/chinese/test/data/train_number.txt | 7 +++++ itn/chinese/test/normalizer_test.py | 1 + 5 files changed, 54 insertions(+) create mode 100644 itn/chinese/data/train_number/prefix.tsv create mode 100644 itn/chinese/rules/train_number.py create mode 100644 itn/chinese/test/data/train_number.txt diff --git a/itn/chinese/data/train_number/prefix.tsv b/itn/chinese/data/train_number/prefix.tsv new file mode 100644 index 00000000..0d51f913 --- /dev/null +++ b/itn/chinese/data/train_number/prefix.tsv @@ -0,0 +1,7 @@ +高 G +动 D +特 T +快 K +直 Z +慢 L +临 L diff --git a/itn/chinese/inverse_normalizer.py b/itn/chinese/inverse_normalizer.py index b4516fdd..eb6249f3 100644 --- a/itn/chinese/inverse_normalizer.py +++ b/itn/chinese/inverse_normalizer.py @@ -18,6 +18,7 @@ from itn.chinese.rules.cardinal import Cardinal from itn.chinese.rules.char import Char from itn.chinese.rules.date import Date +from itn.chinese.rules.train_number import TrainNumber from itn.chinese.rules.fraction import Fraction from itn.chinese.rules.license_plate import LicensePlate from itn.chinese.rules.math import Math @@ -62,6 +63,7 @@ def build_tagger_and_verbalizer(self): | 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) ).optimize() @@ -73,6 +75,7 @@ def build_tagger_and_verbalizer(self): | 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 diff --git a/itn/chinese/rules/train_number.py b/itn/chinese/rules/train_number.py new file mode 100644 index 00000000..34c234cb --- /dev/null +++ b/itn/chinese/rules/train_number.py @@ -0,0 +1,36 @@ +# Copyright (c) 2024 Zhendong Peng (pzd17@tsinghua.org.cn) +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from pynini import string_file +from pynini.lib.pynutil import insert + +from tn.processor import Processor +from tn.utils import get_abs_path + + +class TrainNumber(Processor): + + def __init__(self): + super().__init__(name="trainnumber") + self.build_tagger() + self.build_verbalizer() + + def build_tagger(self): + digit = string_file(get_abs_path("../itn/chinese/data/number/digit.tsv")) + zero = string_file(get_abs_path("../itn/chinese/data/number/zero.tsv")) + digits = zero | digit + prefix = string_file(get_abs_path("../itn/chinese/data/train_number/prefix.tsv")) + number = digits + digits + (digits + digits.ques).ques + tagger = insert('value: "') + prefix + number + insert('"') + self.tagger = self.add_tokens(tagger) diff --git a/itn/chinese/test/data/train_number.txt b/itn/chinese/test/data/train_number.txt new file mode 100644 index 00000000..2d9d5743 --- /dev/null +++ b/itn/chinese/test/data/train_number.txt @@ -0,0 +1,7 @@ +高一二 => G12 +高一二三 => G123 +高一二三四 => G1234 +动二三 => D23 +动三四五六 => D3456 +特五六七八 => T5678 +快一零二三 => K1023 diff --git a/itn/chinese/test/normalizer_test.py b/itn/chinese/test/normalizer_test.py index 69b4c5d2..aa958814 100644 --- a/itn/chinese/test/normalizer_test.py +++ b/itn/chinese/test/normalizer_test.py @@ -38,6 +38,7 @@ class TestNormalizer: parse_test_case("data/whitelist.txt"), parse_test_case("data/number.txt"), parse_test_case("data/license_plate.txt"), + parse_test_case("data/train_number.txt"), parse_test_case("data/normalizer.txt"), )