这是一段 Python 代码,用来微调 Whisper 模型,实现方言语音转写成标准普通话文本,相当于把识别和翻译合并成一步。
```python
# train_dialect_asr.py
# 方言语音 -> 普通话文本(端到端识别+翻译)
# pip install torch transformers datasets librosa soundfile evaluate jiwer accelerate peft
import json
import torch
import librosa
from dataclasses import dataclass
from typing import Any, Dict, List
from datasets import Dataset, Audio
from transformers import (
WhisperProcessor,
WhisperForConditionalGeneration,
Seq2SeqTrainingArguments,
Seq2SeqTrainer,
EarlyStoppingCallback,
)
import evaluate
MODEL_NAME = "openai/whisper-small" # 数据多可换 medium / large-v3
LANG = "zh" # Whisper 的语言 token
TRAIN_JSONL = "data/train.jsonl"
VALID_JSONL = "data/valid.jsonl"
OUTPUT_DIR = "./whisper-dialect"
device = "cuda" if torch.cuda.is_available() else "cpu"
# —————————————————————————
# 1. 数据加载
# 每行 JSONL 格式:
# {"audio": "data/audio/001.wav", "target": "你好吗"}
# target 写“标准普通话文本” → 模型直接学会 方言语音→普通话
# 若想保留方言原文,再加一列 "dialect_text",target 换成它即可
# —————————————————————————
def load_jsonl(path: str) -> Dataset:
rows = [json.loads(line) for line in open(path, encoding="utf-8") if line.strip()]
ds = Dataset.from_list(rows)
# 统一重采样到 16k,Whisper 要求
ds = ds.cast_column("audio", Audio(sampling_rate=16000))
return ds
# —————————————————————————
# 2. 特征提取
# —————————————————————————
processor = WhisperProcessor.from_pretrained(MODEL_NAME, language=LANG, task="transcribe")
def prepare_dataset(batch):
audio = batch["audio"]
batch["input_features"] = processor.feature_extractor(
audio["array"], sampling_rate=audio["sampling_rate"]
).input_features[0]
batch["labels"] = processor.tokenizer(batch["target"]).input_ids
return batch
# —————————————————————————
# 3. 动态 padding
# —————————————————————————
@dataclass
class DataCollatorSpeechSeq2SeqWithPadding:
processor: Any
def __call__(self, features: List[Dict[str, Any]]) -> Dict[str, torch.Tensor]:
# 音频特征
input_features = [{"input_features": f["input_features"]} for f in features]
batch = self.processor.feature_extractor.pad(input_features, return_tensors="pt")
# 标签
label_features = [{"input_ids": f["labels"]} for f in features]
labels_batch = self.processor.tokenizer.pad(label_features, return_tensors="pt")
# padding 位置置 -100,不参与 loss
labels = labels_batch["input_ids"].masked_fill(
labels_batch.attention_mask.ne(1), -100
)
# 去掉开头的 BOS(Trainer 会自己加 decoder_start_token_id)
if (labels[:, 0] == self.processor.tokenizer.bos_token_id).all().cpu().item():
labels = labels[:, 1:]
batch["labels"] = labels
return batch
# —————————————————————————
# 4. 评估指标:中文用 CER(字错率)
# —————————————————————————
cer_metric = evaluate.load("cer")
def compute_metrics(pred):
pred_ids = pred.predictions
label_ids = pred.label_ids.copy()
label_ids[label_ids == -100] = processor.tokenizer.pad_token_id
pred_str = processor.tokenizer.batch_decode(pred_ids, skip_special_tokens=True)
label_str = processor.tokenizer.batch_decode(label_ids, skip_special_tokens=True)
cer = cer_metric.compute(predictions=pred_str, references=label_str)
return {"cer": cer}
# —————————————————————————
# 5. 训练
# —————————————————————————
def main():
train_ds = load_jsonl(TRAIN_JSONL).map(
prepare_dataset, remove_columns=["audio", "target"], num_proc=4
)
valid_ds = load_jsonl(VALID_JSONL).map(
prepare_dataset, remove_columns=["audio", "target"], num_proc=4
)
model = WhisperForConditionalGeneration.from_pretrained(MODEL_NAME)
# 关键:固定解码语言和任务,否则模型可能输出英文
model.generation_config.language = LANG
model.generation_config.task = "transcribe"
model.generation_config.forced_decoder_ids = None
# 只在中文数据上微调,可把其他语言的 embedding 冻结,省显存
model.config.forced_decoder_ids = None
data_collator = DataCollatorSpeechSeq2SeqWithPadding(processor=processor)
args = Seq2SeqTrainingArguments(
output_dir=OUTPUT_DIR,
per_device_train_batch_size=8,
per_device_eval_batch_size=8,
gradient_accumulation_steps=2,
learning_rate=1e-5,
warmup_steps=200,
num_train_epochs=10,
gradient_checkpointing=True,
fp16=torch.cuda.is_available(),
bf16=False,
eval_strategy="steps", # 老版本 transformers 用 evaluation_strategy
eval_steps=200,
save_steps=200,
logging_steps=50,
predict_with_generate=True,
generation_max_length=225,
save_total_limit=3,
load_best_model_at_end=True,
metric_for_best_model="cer",
greater_is_better=False,
report_to=["tensorboard"],
remove_unused_columns=False,
)
trainer = Seq2SeqTrainer(
model=model,
args=args,
train_dataset=train_ds,
eval_dataset=valid_ds,
data_collator=data_collator,
compute_metrics=compute_metrics,
tokenizer=processor.feature_extractor,
callbacks=[EarlyStoppingCallback(early_stopping_patience=3)],
)
trainer.train()
trainer.save_model(OUTPUT_DIR)
processor.save_pretrained(OUTPUT_DIR)
print(f"训练完成,模型保存在 {OUTPUT_DIR}")
# —————————————————————————
# 6. 推理:方言音频 -> 普通话文本
# —————————————————————————
def transcribe(audio_path: str, model_dir: str = OUTPUT_DIR):
_processor = WhisperProcessor.from_pretrained(model_dir)
_model = WhisperForConditionalGeneration.from_pretrained(model_dir).to(device).eval()
speech, _ = librosa.load(audio_path, sr=16000)
inputs = _processor(
speech, sampling_rate=16000, return_tensors="pt"
).input_features.to(device)
with torch.no_grad():
ids = _model.generate(
inputs,
language=LANG,
task="transcribe",
max_new_tokens=225,
num_beams=5,
)
return _processor.batch_decode(ids, skip_special_tokens=True)[0]
if __name__ == "__main__":
import sys
if len(sys.argv) > 1:
# python train_dialect_asr.py 音频路径
print("识别结果:", transcribe(sys.argv[1]))
else:
main()
```
方言识别与翻译的训练流程拆解
从训练到推理,这段代码把方言语音转普通话的流程串了起来。您可以按这几个模块来理解:
· 数据准备:训练数据用 JSONL 格式,每行包含音频路径和对应的普通话文本。音频会自动重采样到 16kHz,并提取成 Whisper 需要的特征。
· 模型训练:基于 openai/whisper-small 微调,冻结解码语言为中文,使用 CER(字错率)作为评估指标。训练时动态 padding,并自动保存验证集上表现最好的模型。
· 推理调用:训练完成后,可以直接用命令行传入方言音频路径,模型会输出普通话文本。推理时用 beam search 解码,识别结果更稳定。
—
优化建议: 如果手头的音频是方言原文标注,可以把 JSONL 里的 “target” 列改成对应的普通话文本,模型就会直接学习“方言语音 → 普通话”的映射。
仅参考学习用,

