如何微调mBART或mT5实现罗马化印地语到原生脚本转写?
罗马化印地语转原生印地语脚本:mBART/mT5微调实践方案
一、数据准备
任务定位
完全可以将这个任务视为序列到序列的翻译任务——输入是拉丁字母书写的罗马化印地语,输出是Devanagari脚本的原生印地语,二者的序列对应关系和机器翻译的范式完全匹配,mBART和mT5的预训练架构天然适合这类任务。
数据集格式化与预处理
- 拆分数据集:按8:1:1的比例将平行数据拆分为训练集、验证集、测试集,确保三个集合的句子长度、场景分布一致(比如不要让训练集全是日常对话,验证集全是正式文本)。
- 存储格式:
- 推荐用CSV或JSON格式存储,每条数据对应一组输入输出对:
- CSV:设置两列,列名分别为
source(罗马化文本,如"aap kaise hain?")、target(原生印地语,如"आप कैसे हैं?") - JSON:每条数据为
{"source": "aap kaise hain?", "target": "आप कैसे हैं?"}
- CSV:设置两列,列名分别为
- 推荐用CSV或JSON格式存储,每条数据对应一组输入输出对:
- 预处理要点:
- 统一输入文本大小写:将罗马化文本全部转为小写(如
"Aap"转"aap"),因为罗马化印地语的大小写不影响语义,避免模型学习不必要的差异 - 清理冗余字符:移除多余的连续空格、重复标点,保留原句的必要标点(如问号、感叹号),确保输入输出的标点对应
- 校验目标文本:确认所有原生印地语都是标准的Devanagari脚本,无乱码、拼写错误
- 统一输入文本大小写:将罗马化文本全部转为小写(如
二、模型微调步骤与专属设置
通用前提
使用Hugging Face的transformers和datasets库完成微调,确保环境安装了torch、transformers、datasets、evaluate依赖。
针对mBART的微调步骤
- 加载模型与Tokenizer:
from transformers import MBartForConditionalGeneration, MBartTokenizer model_name = "facebook/mbart-large-50" tokenizer = MBartTokenizer.from_pretrained(model_name) model = MBartForConditionalGeneration.from_pretrained(model_name) # 指定目标语言为印地语,罗马化文本用拉丁字母兼容编码 tokenizer.src_lang = "en" tokenizer.tgt_lang = "hi_IN" - 数据编码:
定义编码函数,将输入输出文本转为模型可处理的张量,注意将pad token设为-100以避免计算损失:def preprocess_function(examples): inputs = examples["source"] targets = examples["target"] model_inputs = tokenizer(inputs, max_length=256, padding="max_length", truncation=True, return_tensors="pt") # 处理目标文本 with tokenizer.as_target_tokenizer(): labels = tokenizer(targets, max_length=256, padding="max_length", truncation=True, return_tensors="pt") # 将pad token替换为-100,排除损失计算 labels["input_ids"] = [[(l if l != tokenizer.pad_token_id else -100) for l in label] for label in labels["input_ids"]] model_inputs["labels"] = labels["input_ids"] return model_inputs # 对数据集批量应用编码 encoded_dataset = raw_dataset.map(preprocess_function, batched=True) - 训练配置与启动:
使用Seq2SeqTrainingArguments配置参数,用Seq2SeqTrainer启动训练:from transformers import Seq2SeqTrainingArguments, Seq2SeqTrainer import evaluate import numpy as np metric = evaluate.load("cer") # 用字符错误率评估转写准确率 def compute_metrics(eval_pred): predictions, labels = eval_pred decoded_preds = tokenizer.batch_decode(predictions, skip_special_tokens=True) # 将labels中的-100转回pad token id,用于解码 labels = np.where(labels != -100, labels, tokenizer.pad_token_id) decoded_labels = tokenizer.batch_decode(labels, skip_special_tokens=True) cer = metric.compute(predictions=decoded_preds, references=decoded_labels) return {"cer": cer} training_args = Seq2SeqTrainingArguments( output_dir="./mbart_roman_hindi", per_device_train_batch_size=8, per_device_eval_batch_size=8, learning_rate=2e-5, num_train_epochs=4, fp16=True, # 混合精度训练,加速训练进程 save_total_limit=3, # 只保留最优的3个模型 checkpoint evaluation_strategy="epoch", save_strategy="epoch", load_best_model_at_end=True, metric_for_best_model="cer", greater_is_better=False # CER越低,模型效果越好 ) trainer = Seq2SeqTrainer( model=model, args=training_args, train_dataset=encoded_dataset["train"], eval_dataset=encoded_dataset["validation"], tokenizer=tokenizer, compute_metrics=compute_metrics ) trainer.train()
针对mT5的微调步骤
mT5是统一多语言模型,无需指定语言代码,只需给输入添加任务前缀即可明确任务目标:
- 加载模型与Tokenizer:
from transformers import T5ForConditionalGeneration, T5Tokenizer model_name = "google/mt5-base" tokenizer = T5Tokenizer.from_pretrained(model_name) model = T5ForConditionalGeneration.from_pretrained(model_name) - 数据编码:
给输入文本添加"roman_to_hindi: "前缀,帮助模型聚焦转写任务:def preprocess_function(examples): inputs = ["roman_to_hindi: " + src for src in examples["source"]] targets = examples["target"] model_inputs = tokenizer(inputs, max_length=256, padding="max_length", truncation=True, return_tensors="pt") with tokenizer.as_target_tokenizer(): labels = tokenizer(targets, max_length=256, padding="max_length", truncation=True, return_tensors="pt") labels["input_ids"] = [[(l if l != tokenizer.pad_token_id else -100) for l in label] for label in labels["input_ids"]] model_inputs["labels"] = labels["input_ids"] return model_inputs encoded_dataset = raw_dataset.map(preprocess_function, batched=True) - 训练配置与启动:
训练参数和mBART类似,仅调整学习率(mT5适合稍高的学习率):training_args = Seq2SeqTrainingArguments( output_dir="./mt5_roman_hindi", per_device_train_batch_size=8, per_device_eval_batch_size=8, learning_rate=3e-5, # mT5适配稍高的学习率 num_train_epochs=4, fp16=True, save_total_limit=3, evaluation_strategy="epoch", save_strategy="epoch", load_best_model_at_end=True, metric_for_best_model="cer", greater_is_better=False ) trainer = Seq2SeqTrainer( model=model, args=training_args, train_dataset=encoded_dataset["train"], eval_dataset=encoded_dataset["validation"], tokenizer=tokenizer, compute_metrics=compute_metrics # 复用mBART的评估函数 ) trainer.train()
转写任务专属超参数与优化技巧
- 批次大小:根据GPU显存调整,单16GB显存GPU推荐
per_device_train_batch_size=8,显存不足时开启gradient_accumulation_steps=2(等效于批次大小翻倍) - 学习率:mBART用
2e-5,mT5用3e-5,转写任务比通用翻译简单,无需过高学习率 - 训练轮数:3-5轮即可,依赖验证集CER指标早停,避免过拟合
- 评估指标:优先用CER(字符错误率),它直接衡量字符级的转写准确率,比BLEU更适合转写任务;可搭配SacreBLEU作为辅助指标
- 解码策略:推理时用
beam_search(beam size=3-5),比贪心解码能生成更准确的结果,尤其是长句子 - 其他优化:开启
fp16混合精度训练,大幅提升训练速度;设置load_best_model_at_end=True,保留验证集效果最好的模型
内容的提问来源于stack exchange,提问作者sameera perera
相关产品推荐
相关产品推荐

