~附:安装依赖库及工程源码)
别再死磕枯燥理论了AI 时代拿实战作品说话才是硬道理原创不易哈希望可以帮到还有些许学习劲儿的同学们【进阶版还在创作中耗费精力中……】跳转到专栏目录你学习更有方向和思路……入门实践工程十三基于 CNN 的语音命令识别Keyword Spotting~附:安装依赖库及工程源码简介在 Google Speech Commands 数据集上训练一个小型 CNN识别 10 个语音命令yes/no/up/down/left/right/on/off/stop/go音频转 Mel 频谱图后用 CNN 分类理解音频深度学习「波形→频谱图→分类」全流程。工程详细介绍核心思想把音频分类转化为图像分类——波形转 Mel 频谱图后用 CNN 分类因为语音的时频结构在频谱图上呈二维模式卷积网络能有效捕获是音频深度学习的通用范式。实现方法数据Google Speech Commands v2torchaudio 自动下载约 1GB过滤出 10 个核心命令词。模型3 层卷积 自适应池化 全连接的小型 CNN。流程波形补零/截断到 1 秒 → MelSpectrogram → 转 dB → 标准化 → CNN 分类交叉熵 Adam 训练 5 轮并评估验证准确率–predict 对单个 wav 输出 Top-5 概率。输出验证准确率 80% keyword_cnn.pth 单文件预测。目录结构13_speech_commands/ ├── main.py # 数据加载 训练 预测 ├── requirements.txt ├── data/ # Speech Commands 数据集自动下载 └── keyword_cnn.pth # 训练后的权重正确安装pipinstalltorch torchaudio numpy python main.py若 torchaudio 安装失败可用conda install -c conda-forge torchaudio。首次运行需联网下载 Google Speech Commands 数据集约 1GB。运行方式pipinstall-rrequirements.txt# 训练首次需联网下载约 1GB 数据集python main.py# 训练后对单个 wav 文件预测python main.py--predicttest.wav说明数据集由 torchaudioSPEECHCOMMANDS首次自动下载约 1GB需联网。流程波形 → 补零/截断到 1s → Mel 频谱图 → 转 dB → 标准化 → CNN。CPU 可跑但较慢每轮约 5-10 分钟建议 GPU。如 torchaudio 安装失败可conda install -c conda-forge torchaudio。预期结果5 个 Epoch 后验证准确率通常 80%对 wav 文件输出 Top-5 命令词概率扩展方向增加背景噪声增强、SpecAugment 数据增强换为更深的 ResNet / 音频 Transformer部署为实时唤醒词streaming 推理工程源码main.py 入门实践工程十三基于 CNN 的语音命令识别Keyword Spotting 在 Google Speech Commands 数据集上训练一个小型 CNN识别 10 个语音命令 yes / no / up / down / left / right / on / off / stop / go。 音频转 Mel 频谱图后用 CNN 分类理解音频深度学习全流程。 运行 python main.py python main.py --predict path/to/test.wav 数据集由 torchaudio 首次运行自动下载Speech Commands v2约 1GB。 CPU 可跑但较慢建议 GPU。 importargparseimportosimportnumpyasnpimporttorchimporttorch.nnasnnimporttorch.nn.functionalasFimporttorchaudiofromtorch.utils.dataimportDataset,DataLoader BASE_DIRos.path.dirname(os.path.abspath(__file__))DATA_DIRos.path.join(BASE_DIR,data)DEVICEtorch.device(cudaiftorch.cuda.is_available()elsecpu)# 10 个核心命令词LABELS[down,go,left,no,off,on,right,stop,up,yes]LABEL2IDX{l:ifori,linenumerate(LABELS)}SAMPLE_RATE16000N_MELS40EPOCHS5BATCH_SIZE64defget_mel_transform():returntorchaudio.transforms.MelSpectrogram(sample_rateSAMPLE_RATE,n_melsN_MELS)defpad_or_crop(waveform,lengthSAMPLE_RATE):把波形统一为固定长度不足补零过多截断。sigwaveform[0]ifwaveform.dim()2elsewaveformifsig.size(0)length:padlength-sig.size(0)sigF.pad(sig,(0,pad))else:sigsig[:length]returnsig.unsqueeze(0)# [1, length]classSpeechCommandDataset(Dataset):封装 torchaudio 的 SPEACHCOMMANDS过滤到 10 个核心命令并转 Mel。def__init__(self,rootDATA_DIR,subsettraining):os.makedirs(root,exist_okTrue)try:self.datasettorchaudio.datasets.SPEECHCOMMANDS(rootroot,downloadTrue,subsetsubset)exceptExceptionase:print(f[错误] 下载/加载 Speech Commands 失败:{e})print(请检查网络。数据集约 1GB。)raiseself.melget_mel_transform()# 预先过滤出 10 个核心命令的索引self.indices[]foriinrange(len(self.dataset)):labelself.dataset[i][2]iflabelinLABEL2IDX:self.indices.append(i)def__len__(self):returnlen(self.indices)def__getitem__(self,idx):waveform,sr,label,*_self.dataset[self.indices[idx]]sigpad_or_crop(waveform,SAMPLE_RATE)melself.mel(sig)# [1, n_mels, time]meltorchaudio.transforms.AmplitudeToDB()(mel)# 归一化mel(mel-mel.mean())/(mel.std()1e-6)returnmel,LABEL2IDX[label]classKeywordCNN(nn.Module):def__init__(self,num_classeslen(LABELS)):super().__init__()self.featuresnn.Sequential(nn.Conv2d(1,16,3,padding1),nn.ReLU(),nn.MaxPool2d(2),nn.Conv2d(16,32,3,padding1),nn.ReLU(),nn.MaxPool2d(2),nn.Conv2d(32,64,3,padding1),nn.ReLU(),nn.MaxPool2d(2),)# 自适应池化避免固定尺寸推断self.poolnn.AdaptiveAvgPool2d((1,1))self.headnn.Sequential(nn.Flatten(),nn.Linear(64,64),nn.ReLU(),nn.Dropout(0.3),nn.Linear(64,num_classes),)defforward(self,x):xself.features(x)xself.pool(x)returnself.head(x)torch.no_grad()defevaluate(model,loader):model.eval()correct,total0,0formel,labelinloader:mel,labelmel.to(DEVICE),label.to(DEVICE)outmodel(mel)correct(out.argmax(1)label).sum().item()totallabel.size(0)returncorrect/totaldefmain():parserargparse.ArgumentParser()parser.add_argument(--epochs,typeint,defaultEPOCHS)parser.add_argument(--predict,typestr,defaultNone,help对一个 wav 文件做预测)argsparser.parse_args()ifargs.predict:predict_wav(args.predict)returnprint(f设备:{DEVICE})print(加载 Speech Commands 数据集首次需联网下载约 1GB...)train_setSpeechCommandDataset(subsettraining)val_setSpeechCommandDataset(subsetvalidation)print(f训练样本(10类):{len(train_set)}验证样本:{len(val_set)})train_loaderDataLoader(train_set,batch_sizeBATCH_SIZE,shuffleTrue,num_workers0)val_loaderDataLoader(val_set,batch_sizeBATCH_SIZE,num_workers0)modelKeywordCNN().to(DEVICE)optimizertorch.optim.Adam(model.parameters(),lr1e-3)forepochinrange(1,args.epochs1):model.train()total,correct,loss_sum0,0,0.0formel,labelintrain_loader:mel,labelmel.to(DEVICE),label.to(DEVICE)optimizer.zero_grad()outmodel(mel)lossF.cross_entropy(out,label)loss.backward()optimizer.step()loss_sumloss.item()*mel.size(0)correct(out.argmax(1)label).sum().item()totalmel.size(0)val_accevaluate(model,val_loader)print(fEpoch{epoch}/{args.epochs}训练损失{loss_sum/total:.4f}f训练准确率{correct/total:.4f}验证准确率{val_acc:.4f})pthos.path.join(BASE_DIR,keyword_cnn.pth)torch.save(model.state_dict(),pth)print(f\n模型权重已保存到:{pth})torch.no_grad()defpredict_wav(wav_path:str):print(f加载权重并预测:{wav_path})pthos.path.join(BASE_DIR,keyword_cnn.pth)ifnotos.path.exists(pth):print([错误] 未找到 keyword_cnn.pth请先运行训练python main.py)returnmodelKeywordCNN().to(DEVICE)model.load_state_dict(torch.load(pth,map_locationDEVICE))model.eval()waveform,srtorchaudio.load(wav_path)ifsr!SAMPLE_RATE:waveformtorchaudio.transforms.Resample(sr,SAMPLE_RATE)(waveform)sigpad_or_crop(waveform,SAMPLE_RATE).to(DEVICE)melget_mel_transform()(sig)meltorchaudio.transforms.AmplitudeToDB()(mel)mel(mel-mel.mean())/(mel.std()1e-6)outmodel(mel.unsqueeze(0))probsF.softmax(out,dim1)[0]top5torch.topk(probs,5)print(\n预测 Top-5 命令词)forp,iinzip(top5.values.tolist(),top5.indices.tolist()):print(f{LABELS[i]:8}{p*100:5.1f}%)if__name____main__:main()