ローマ字入力列から日本語文(漢字かな交じり文)への変換を行う RNN-Transducer モデルです。
タイプミス(重複・脱字・隣接キー誤打・転置)を含む入力に対しても頑健に変換できるよう、ノイズを付与したデータで学習しています。
A RNN-Transducer model that converts romaji input sequences into Japanese text (kanji-kana mixed sentences), trained to be robust against typos.
学習中の valid CER は検証データ 100 サンプルの抽出値(5 エポックごと、greedy)。テスト全件 CER(7.05%)との差は少数サンプルによるサンプリング誤差であり、最終性能はテストセット全件評価を参照のこと。
1git clone https://github.com/takumiecd/kairo-ai
2cd kairo-ai && uv sync
3
4# モデル一式(checkpoint + config + vocab)をダウンロード
5hf download takumiecd/kairo-rnnt-trf-v1 --local-dir artifacts/rnnt-trf-v1
6
7# greedy デコードで推論(デフォルトで checkpoints/best.pt が使われる)
8uv run python -m decode.greedy \
9 --artifact-dir artifacts/rnnt-trf-v1 \
10 --input "wagahaihanekodearu."
11# => 吾輩は猫である。
本モデルの構築に使用したコマンド。詳細なフラグの説明は
kairo-ai リポジトリ の README を参照。
1# Wikipedia(10万unit × augmentations 1 ≒ 20万ペア)
2python -m dataset.source_wikipedia \
3 --dump data/raw/wiki/jawiki-latest-pages-articles.xml.bz2 \
4 --output data/external/wiki_ja.jsonl \
5 --license cc_by_sa_gfdl \
6 --max-units 100000 \
7 --augmentations 1
8
9# Tatoeba(5万unit × augmentations 1 ≒ 10万ペア)
10python -m dataset.source_tatoeba \
11 --sentences data/raw/tatoeba/sentences.tar.bz2 \
12 --output data/external/tatoeba_ja.jsonl \
13 --lang jpn \
14 --max-units 50000 \
15 --augmentations 1
16
17# 青空文庫『吾輩は猫である』(500unit × augmentations 2 = 1,500ペア)
18python -m dataset.source_text \
19 --source https://www.aozora.gr.jp/cards/000148/files/789_ruby_5639.zip \
20 --output data/external/aozora_wagahai.jsonl \
21 --source-name aozora \
22 --license aozora_public_domain_checked \
23 --format aozora \
24 --max-units 500 \
25 --augmentations 2
1cat data/external/wiki_ja.jsonl \
2 data/external/tatoeba_ja.jsonl \
3 data/external/aozora_wagahai.jsonl > data/combined/all_sources.jsonl
4
5python -m dataset.split \
6 --input data/combined/all_sources.jsonl \
7 --output-dir data/combined/all_sources
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-v1 \
5 --encoder-type transformer --prediction-type transformer \
6 --encoder-layers 4 --prediction-layers 2 \
7 --embed-dim 256 --hidden-dim 256 --num-heads 4 \
8 --epochs 50 --batch-size 8 \
9 --learning-rate 3e-4 \
10 --lr-scheduler cosine --warmup-ratio 0.05 \
11 --device cuda \
12 --valid-decode greedy --valid-cer-samples 100 --valid-cer-every 5 \
13 --max-positions 512
1python -m eval.run_test \
2 --artifact-dir artifacts/rnnt-trf-v1 \
3 --data data/combined/all_sources/test.jsonl \
4 --device cuda