尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

Python训练ai识别方言

发布时间:2026/9/24 13:31:36

资讯中心
01
ARTICLE

Python训练ai识别方言

Python训练ai识别方言
这是一段 Python 代码用来微调 Whisper 模型实现方言语音转写成标准普通话文本相当于把识别和翻译合并成一步。python# train_dialect_asr.py# 方言语音 - 普通话文本端到端识别翻译# pip install torch transformers datasets librosa soundfile evaluate jiwer accelerate peftimport jsonimport torchimport librosafrom dataclasses import dataclassfrom typing import Any, Dict, Listfrom datasets import Dataset, Audiofrom transformers import (WhisperProcessor,WhisperForConditionalGeneration,Seq2SeqTrainingArguments,Seq2SeqTrainer,EarlyStoppingCallback,)import evaluateMODEL_NAME openai/whisper-small # 数据多可换 medium / large-v3LANG zh # Whisper 的语言 tokenTRAIN_JSONL data/train.jsonlVALID_JSONL data/valid.jsonlOUTPUT_DIR ./whisper-dialectdevice cuda if torch.cuda.is_available() else cpu# ---------------------------------------------------------------------------# 1. 数据加载# 每行 JSONL 格式# {audio: data/audio/001.wav, target: 你好吗}# target 写“标准普通话文本” → 模型直接学会 方言语音→普通话# 若想保留方言原文再加一列 dialect_texttarget 换成它即可# ---------------------------------------------------------------------------def load_jsonl(path: str) - Dataset:rows [json.loads(line) for line in open(path, encodingutf-8) if line.strip()]ds Dataset.from_list(rows)# 统一重采样到 16kWhisper 要求ds ds.cast_column(audio, Audio(sampling_rate16000))return ds# ---------------------------------------------------------------------------# 2. 特征提取# ---------------------------------------------------------------------------processor WhisperProcessor.from_pretrained(MODEL_NAME, languageLANG, tasktranscribe)def prepare_dataset(batch):audio batch[audio]batch[input_features] processor.feature_extractor(audio[array], sampling_rateaudio[sampling_rate]).input_features[0]batch[labels] processor.tokenizer(batch[target]).input_idsreturn batch# ---------------------------------------------------------------------------# 3. 动态 padding# ---------------------------------------------------------------------------dataclassclass DataCollatorSpeechSeq2SeqWithPadding:processor: Anydef __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_tensorspt)# 标签label_features [{input_ids: f[labels]} for f in features]labels_batch self.processor.tokenizer.pad(label_features, return_tensorspt)# padding 位置置 -100不参与 losslabels labels_batch[input_ids].masked_fill(labels_batch.attention_mask.ne(1), -100)# 去掉开头的 BOSTrainer 会自己加 decoder_start_token_idif (labels[:, 0] self.processor.tokenizer.bos_token_id).all().cpu().item():labels labels[:, 1:]batch[labels] labelsreturn batch# ---------------------------------------------------------------------------# 4. 评估指标中文用 CER字错率# ---------------------------------------------------------------------------cer_metric evaluate.load(cer)def compute_metrics(pred):pred_ids pred.predictionslabel_ids pred.label_ids.copy()label_ids[label_ids -100] processor.tokenizer.pad_token_idpred_str processor.tokenizer.batch_decode(pred_ids, skip_special_tokensTrue)label_str processor.tokenizer.batch_decode(label_ids, skip_special_tokensTrue)cer cer_metric.compute(predictionspred_str, referenceslabel_str)return {cer: cer}# ---------------------------------------------------------------------------# 5. 训练# ---------------------------------------------------------------------------def main():train_ds load_jsonl(TRAIN_JSONL).map(prepare_dataset, remove_columns[audio, target], num_proc4)valid_ds load_jsonl(VALID_JSONL).map(prepare_dataset, remove_columns[audio, target], num_proc4)model WhisperForConditionalGeneration.from_pretrained(MODEL_NAME)# 关键固定解码语言和任务否则模型可能输出英文model.generation_config.language LANGmodel.generation_config.task transcribemodel.generation_config.forced_decoder_ids None# 只在中文数据上微调可把其他语言的 embedding 冻结省显存model.config.forced_decoder_ids Nonedata_collator DataCollatorSpeechSeq2SeqWithPadding(processorprocessor)args Seq2SeqTrainingArguments(output_dirOUTPUT_DIR,per_device_train_batch_size8,per_device_eval_batch_size8,gradient_accumulation_steps2,learning_rate1e-5,warmup_steps200,num_train_epochs10,gradient_checkpointingTrue,fp16torch.cuda.is_available(),bf16False,eval_strategysteps, # 老版本 transformers 用 evaluation_strategyeval_steps200,save_steps200,logging_steps50,predict_with_generateTrue,generation_max_length225,save_total_limit3,load_best_model_at_endTrue,metric_for_best_modelcer,greater_is_betterFalse,report_to[tensorboard],remove_unused_columnsFalse,)trainer Seq2SeqTrainer(modelmodel,argsargs,train_datasettrain_ds,eval_datasetvalid_ds,data_collatordata_collator,compute_metricscompute_metrics,tokenizerprocessor.feature_extractor,callbacks[EarlyStoppingCallback(early_stopping_patience3)],)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, sr16000)inputs _processor(speech, sampling_rate16000, return_tensorspt).input_features.to(device)with torch.no_grad():ids _model.generate(inputs,languageLANG,tasktranscribe,max_new_tokens225,num_beams5,)return _processor.batch_decode(ids, skip_special_tokensTrue)[0]if __name__ __main__:import sysif 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” 列改成对应的普通话文本模型就会直接学习“方言语音 → 普通话”的映射。仅参考学习用
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。