基于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
相关产品推荐
相关产品推荐

