ローマ字入力列から日本語文(漢字かな交じり文)への変換を行う RNN-Transducer モデルです。
タイプミス(重複・脱字・隣接キー誤打・転置)を含む入力に対しても頑健に変換できるよう、ノイズを付与したデータで学習しています。
A RNN-Transducer model that converts romaji input sequences into Japanese text (kanji-kana mixed sentences), trained to be robust against typos.
rnnt-trf-v1 の改良版です(同一データセット)。主な変更点は以下のとおりです。
rnnt-trf-v1 と同一のデータセットを使用しています。計
301,429 ペア(input: ローマ字列, target: 日本語文)。約半数にタイプミスを模したノイズを付与(train:valid:test = 8:1:1, seed 0 → 241,143 / 30,142 / 30,144 ペア)。
1git clone https://github.com/takumiecd/kairo-ai
2cd kairo-ai && uv sync
3
4# モデル一式(checkpoint + config + vocab)をダウンロード
5hf download takumiecd/kairo-rnnt-trf-v2 --local-dir artifacts/rnnt-trf-v2
6
7# greedy デコードで推論(デフォルトで checkpoints/best.pt が使われる)
8uv run python -m decode.greedy \
9 --artifact-dir artifacts/rnnt-trf-v2 \
10 --input "wagahaihanekodearu."
11# => 吾輩は猫である。
本モデルの構築に使用したコマンド。詳細なフラグの説明は
kairo-ai リポジトリ の README を参照。データセット構築・分割は v1 と同一のため、
rnnt-trf-v1 のモデルカードを参照してください。
1python -m train.rnnt.train \
2 --data data/combined/all_sources/train.jsonl \
3 --valid-data data/combined/all_sources/valid.jsonl \
4 --output-dir artifacts/rnnt-trf-v2 \
5 --encoder-type transformer --prediction-type transformer \
6 --encoder-layers 4 --prediction-layers 2 \
7 --input-embed-dim 96 --output-embed-dim 256 \
8 --encoder-hidden-dim 256 --prediction-hidden-dim 256 --joint-hidden-dim 256 \
9 --feedforward-dim 1024 --num-heads 4 \
10 --epochs 40 --batch-size 128 \
11 --learning-rate 1e-3 \
12 --amp \
13 --batch-order bucket_lattice --max-batch-lattice-cells 160000 \
14 --max-len 96 --max-positions 256 \
15 --device cuda \
16 --valid-decode greedy --valid-cer-samples 500 --valid-cer-every 1
1python -m eval.run_test \
2 --artifact-dir artifacts/rnnt-trf-v2 \
3 --data data/combined/all_sources/test.jsonl \
4 --device cuda