欢迎光临
我们一直在努力

Python训练ai识别方言

这是一段  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” 列改成对应的普通话文本,模型就会直接学习“方言语音 → 普通话”的映射。

 

 

 

仅参考学习用,

赞(0)
未经允许不得转载:171主机测评 » Python训练ai识别方言
分享到: 更多 (0)

评论 抢沙发

  • 昵称 (必填)
  • 邮箱 (必填)
  • 网址