You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

基于wav2vec2的多任务(分类)训练实现技术咨询

多任务Wav2Vec2模型改造与训练指南

1. 加载预训练基础模型

先加载你已训练好的Wav2Vec2基础编码器(不带原任务头):

from transformers import Wav2Vec2Model

# 替换为你的模型本地路径或模型名
base_model = Wav2Vec2Model.from_pretrained("path/to/your/trained/wav2vec2")

2. 构建多任务模型类

自定义继承自nn.Module的模型,整合基础编码器和多个任务头(比如年龄分类头、ASR的CTC头):

import torch
import torch.nn as nn
from transformers import Wav2Vec2CTCLoss

class MultiTaskWav2Vec2(nn.Module):
    def __init__(self, base_model, num_age_classes, vocab_size):
        super().__init__()
        self.base_model = base_model
        # 年龄分类头:基于<s> token特征(或序列池化特征)
        self.age_head = nn.Sequential(
            nn.Linear(base_model.config.hidden_size, 256),
            nn.ReLU(),
            nn.Linear(256, num_age_classes)
        )
        # ASR的CTC分类头
        self.asr_head = nn.Linear(base_model.config.hidden_size, vocab_size)
        self.ctc_loss = Wav2Vec2CTCLoss()

    def forward(self, input_values, attention_mask=None, age_labels=None, asr_labels=None):
        # 获取基础编码器的序列特征
        encoder_outputs = self.base_model(input_values=input_values, attention_mask=attention_mask)
        last_hidden_state = encoder_outputs.last_hidden_state

        # 年龄分类:取<s> token的特征(序列第一个位置)
        cls_feature = last_hidden_state[:, 0, :]
        age_logits = self.age_head(cls_feature)

        # ASR任务:用整个序列特征做CTC预测
        asr_logits = self.asr_head(last_hidden_state)

        # 计算总损失(多任务加权求和)
        total_loss = None
        if age_labels is not None:
            age_loss = nn.CrossEntropyLoss()(age_logits, age_labels)
            total_loss = age_loss
        if asr_labels is not None:
            # CTC损失需要输入长度和标签长度
            input_lengths = attention_mask.sum(dim=1) if attention_mask else torch.full((input_values.size(0),), input_values.size(1), device=input_values.device)
            label_lengths = torch.tensor([len(lab) for lab in asr_labels], device=input_values.device)
            asr_loss = self.ctc_loss(asr_logits.transpose(0,1), asr_labels, input_lengths, label_lengths)
            total_loss = asr_loss if total_loss is None else total_loss + 0.5 * asr_loss  # 可根据任务优先级调整权重

        return {"age_logits": age_logits, "asr_logits": asr_logits, "loss": total_loss}

3. 数据处理适配多任务

自定义Dataset,同时加载音频输入、年龄标签和ASR文本标签(单任务训练时可忽略对应标签):

from transformers import Wav2Vec2Processor
import librosa

class MultiTaskAudioDataset(torch.utils.data.Dataset):
    def __init__(self, audio_paths, age_labels, asr_texts, processor):
        self.audio_paths = audio_paths
        self.age_labels = age_labels
        self.asr_texts = asr_texts
        self.processor = processor  # 复用原模型的processor

    def __getitem__(self, idx):
        # 加载并处理音频
        audio, _ = librosa.load(self.audio_paths[idx], sr=16000)
        inputs = self.processor(audio, sampling_rate=16000, return_tensors="pt")
        input_values = inputs.input_values.squeeze()
        attention_mask = inputs.attention_mask.squeeze() if "attention_mask" in inputs else None

        # 处理ASR标签
        asr_labels = self.processor.tokenizer(self.asr_texts[idx], return_tensors="pt").input_ids.squeeze()

        return {
            "input_values": input_values,
            "attention_mask": attention_mask,
            "age_labels": torch.tensor(self.age_labels[idx]),
            "asr_labels": asr_labels
        }

    def __len__(self):
        return len(self.audio_paths)

4. 训练与推理配置

训练循环示例

from transformers import AdamW, get_scheduler

# 初始化processor和多任务模型
processor = Wav2Vec2Processor.from_pretrained("path/to/your/processor")
model = MultiTaskWav2Vec2(base_model, num_age_classes=5, vocab_size=len(processor.tokenizer))

# 数据加载
dataset = MultiTaskAudioDataset(your_audio_paths, your_age_labels, your_asr_texts, processor)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=4, shuffle=True)

# 优化器与调度器
optimizer = AdamW(model.parameters(), lr=5e-5)
num_epochs = 3
num_training_steps = num_epochs * len(dataloader)
lr_scheduler = get_scheduler("linear", optimizer=optimizer, num_warmup_steps=0, num_training_steps=num_training_steps)

# 设备配置
device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")
model.to(device)

# 开始训练
model.train()
for epoch in range(num_epochs):
    for batch in dataloader:
        batch = {k: v.to(device) for k, v in batch.items()}
        outputs = model(**batch)
        loss = outputs["loss"]
        loss.backward()

        optimizer.step()
        lr_scheduler.step()
        optimizer.zero_grad()

        print(f"Epoch {epoch+1}, Batch Loss: {loss.item():.4f}")

推理示例

model.eval()
with torch.no_grad():
    # 年龄分类推理
    audio, _ = librosa.load("test_audio.wav", sr=16000)
    inputs = processor(audio, sampling_rate=16000, return_tensors="pt").to(device)
    outputs = model(**inputs)
    age_pred = outputs["age_logits"].argmax(dim=-1).item()
    print(f"Predicted Age Class: {age_pred}")

    # ASR推理
    asr_logits = outputs["asr_logits"]
    predicted_ids = torch.argmax(asr_logits, dim=-1)
    transcription = processor.decode(predicted_ids[0])
    print(f"Transcription: {transcription}")

关键注意事项

  • 任务权重:根据任务优先级调整损失的加权系数,比如ASR任务更重要可加大权重
  • 特征选择:年龄分类可选择序列平均池化特征替代 token,根据效果调整
  • 微调策略:数据量较小时可冻结基础编码器的部分层,只训练任务头;数据充足时可全量微调
  • 单任务推理:不需要某任务输出时,可在forward中跳过对应头的计算,或只提取需要的logits

内容的提问来源于stack exchange,提问作者padmalcom

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.05 00:05:23