Views
No views yet
原始音频 [48000 samples]
↓ Wav2Vec2特征编码器 (7层1D卷积)
局部特征 [1199, 768]
↓ Wav2Vec2上下文网络 (12层Transformer)
上下文特征 [1199, 768]
↓ 全局平均池化
固定特征 [768]
↓ 分类头 (2层全连接)
分类结果 [2] (开关/背景)| 样本类型 | 时长范围 | RMS能量 | 频谱质心 | 过零率 |
|---|---|---|---|---|
| 开关声音 | 3.2-5.2s | 0.0079-0.0115 | 1587-1992Hz | 0.0657-0.1215 |
| 背景噪音 | 2.0-4.0s | 0.005-0.02 | 500-1500Hz | 0.05-0.15 |
| 指标 | 数值 |
|---|---|
| 准确率 | 100% |
| 精确率 | 100% |
| 召回率 | 100% |
| F1分数 | 100% |
| 训练轮数 | 15 epochs |
| 模型大小 | 361MB |
| 推理延迟 | <100ms |
实际\预测 无开关 有开关
无开关 2 0
有开关 0 2pip install torch torchaudio transformers huggingface_hub1from huggingface_hub import hf_hub_download
2import torch
3import torchaudio
4from transformers import Wav2Vec2Model
5
6# 下载模型
7model_path = hf_hub_download(
8 repo_id="lemonhall/heater-switch-detector",
9 filename="switch_detector_model.pth"
10)
11
12# 加载模型
13device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')
14checkpoint = torch.load(model_path, map_location=device)
15
16# 重建模型架构
17wav2vec2_model = Wav2Vec2Model.from_pretrained("facebook/wav2vec2-base")
18classifier = torch.nn.Sequential(
19 torch.nn.Linear(768, 256),
20 torch.nn.ReLU(),
21 torch.nn.Dropout(0.3),
22 torch.nn.Linear(256, 2)
23)
24
25# 加载权重
26classifier.load_state_dict(checkpoint['classifier_state_dict'])
27classifier.eval()1def predict_audio(audio_path):
2 # 加载音频
3 waveform, sample_rate = torchaudio.load(audio_path)
4
5 # 重采样到16kHz
6 if sample_rate != 16000:
7 resampler = torchaudio.transforms.Resample(sample_rate, 16000)
8 waveform = resampler(waveform)
9
10 # 转为单声道
11 if waveform.shape[0] > 1:
12 waveform = waveform.mean(dim=0, keepdim=True)
13
14 # 特征提取
15 with torch.no_grad():
16 features = wav2vec2_model(waveform).last_hidden_state
17 pooled_features = features.mean(dim=1) # 全局平均池化
18
19 # 分类预测
20 logits = classifier(pooled_features)
21 probabilities = torch.softmax(logits, dim=-1)
22 prediction = torch.argmax(probabilities, dim=-1)
23
24 return {
25 'prediction': '开关按下' if prediction.item() == 1 else '背景声音',
26 'confidence': probabilities.max().item(),
27 'probabilities': {
28 '背景声音': probabilities[0][0].item(),
29 '开关按下': probabilities[0][1].item()
30 }
31 }
32
33# 使用示例
34result = predict_audio("test_audio.wav")
35print(f"预测结果: {result['prediction']}")
36print(f"置信度: {result['confidence']:.3f}")1import pyaudio
2import numpy as np
3
4def realtime_detection():
5 # 音频参数
6 SAMPLE_RATE = 16000
7 CHUNK_SIZE = 1024
8 DETECTION_WINDOW = 3.0 # 3秒检测窗口
9
10 # 初始化音频流
11 audio = pyaudio.PyAudio()
12 stream = audio.open(
13 format=pyaudio.paFloat32,
14 channels=1,
15 rate=SAMPLE_RATE,
16 input=True,
17 frames_per_buffer=CHUNK_SIZE
18 )
19
20 print("🎤 开始实时检测...")
21
22 buffer = []
23 window_size = int(DETECTION_WINDOW * SAMPLE_RATE)
24
25 try:
26 while True:
27 # 读取音频数据
28 data = stream.read(CHUNK_SIZE)
29 audio_chunk = np.frombuffer(data, dtype=np.float32)
30 buffer.extend(audio_chunk)
31
32 # 保持窗口大小
33 if len(buffer) > window_size:
34 buffer = buffer[-window_size:]
35
36 # 检测
37 if len(buffer) == window_size:
38 waveform = torch.FloatTensor(buffer).unsqueeze(0)
39
40 with torch.no_grad():
41 features = wav2vec2_model(waveform).last_hidden_state
42 pooled_features = features.mean(dim=1)
43 logits = classifier(pooled_features)
44 probabilities = torch.softmax(logits, dim=-1)
45
46 switch_prob = probabilities[0][1].item()
47
48 if switch_prob > 0.93: # 高置信度阈值
49 print(f"🔥 检测到开关按下! 置信度: {switch_prob:.3f}")
50
51 except KeyboardInterrupt:
52 print("\n⏹️ 检测停止")
53 finally:
54 stream.stop_stream()
55 stream.close()
56 audio.terminate()
57
58# 运行实时检测
59realtime_detection()1@misc{heater-switch-detector-2024,
2 title={基于Wav2Vec2的热水器开关声音检测器},
3 author={lemonhall},
4 year={2024},
5 howpublished={\url{https://huggingface.co/lemonhall/heater-switch-detector}}
6}