test/scripts/analyze_audio_rms.py

329 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters!

This file contains ambiguous Unicode characters that may be confused with others in your current locale. If your use case is intentional and legitimate, you can safely ignore this warning. Use the Escape button to highlight these characters.

#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
音频RMS时序分析工具
用于分析音频<EFBFBD><EFBFBD><EFBFBD>件的RMS能量,帮助确定远场过滤的阈值。
支持立体声、左声道、右声道选择。
"""
import argparse
import numpy as np
import matplotlib.pyplot as plt
from pathlib import Path
import sys
from typing import Optional
# 设置中文显示
plt.rcParams['font.sans-serif'] = ['Arial Unicode MS', 'SimHei', 'DejaVu Sans']
plt.rcParams['axes.unicode_minus'] = False
def load_audio(file_path: str, channel: str = 'stereo') -> tuple:
"""加载音频文件
Args:
file_path: 音频文件路径
channel: 声道选择 ('stereo', 'left', 'right')
Returns:
(audio_data, sample_rate): 音频数据和采样率
"""
file_ext = Path(file_path).suffix.lower()
if file_ext == '.wav':
import wave
with wave.open(file_path, 'rb') as wav_file:
sample_rate = wav_file.getframerate()
n_channels = wav_file.getnchannels()
sample_width = wav_file.getsampwidth()
n_frames = wav_file.getnframes()
# 读取音频数据
audio_bytes = wav_file.readframes(n_frames)
# 转换为numpy数组
if sample_width == 2: # 16-bit
audio_int = np.frombuffer(audio_bytes, dtype=np.int16)
elif sample_width == 4: # 32-bit
audio_int = np.frombuffer(audio_bytes, dtype=np.int32)
else:
raise ValueError(f"不支持的采样位深: {sample_width}")
# 转换为float32 (-1.0 to 1.0)
audio_float = audio_int.astype(np.float32) / (2 ** (8 * sample_width - 1))
# 处理多声道
if n_channels > 1:
audio_float = audio_float.reshape(-1, n_channels)
if channel == 'left':
audio_float = audio_float[:, 0]
print(f"✓ 使用左声道")
elif channel == 'right':
audio_float = audio_float[:, 1]
print(f"✓ 使用右声道")
else: # stereo - 平均
audio_float = np.mean(audio_float, axis=1)
print(f"✓ 使用立体声(双声道平均)")
else:
print(f"✓ 使用单声道")
return audio_float, sample_rate
else:
# 尝试使用 soundfile 或 librosa
try:
import soundfile as sf
audio_float, sample_rate = sf.read(file_path)
if len(audio_float.shape) > 1: # 多声道
if channel == 'left':
audio_float = audio_float[:, 0]
print(f"✓ 使用左声道")
elif channel == 'right':
audio_float = audio_float[:, 1]
print(f"✓ 使用右声道")
else:
audio_float = np.mean(audio_float, axis=1)
print(f"✓ 使用立体声(双声道平均)")
else:
print(f"✓ 使用单声道")
return audio_float, sample_rate
except ImportError:
print("错误: 请先运行 ./scripts/sync_cpu_env.sh 安装 soundfile 依赖")
sys.exit(1)
def calculate_rms_energy(audio_array: np.ndarray) -> float:
"""计算音频RMS能量
Args:
audio_array: float32音频数组,范围-1.0到1.0
Returns:
RMS能量值
"""
if len(audio_array) == 0:
return 0.0
return float(np.sqrt(np.mean(audio_array ** 2)))
def analyze_rms_timeline(audio_data: np.ndarray, sample_rate: int,
chunk_size_ms: int = 240) -> tuple:
"""分析音频的RMS时序
Args:
audio_data: 音频数据
sample_rate: 采样率
chunk_size_ms: 分块大小(毫秒)
Returns:
(time_points, rms_values): 时间点和对应的RMS值
"""
chunk_samples = int(sample_rate * chunk_size_ms / 1000)
n_chunks = len(audio_data) // chunk_samples
time_points = []
rms_values = []
for i in range(n_chunks):
start_idx = i * chunk_samples
end_idx = start_idx + chunk_samples
chunk = audio_data[start_idx:end_idx]
rms = calculate_rms_energy(chunk)
time_s = (start_idx + chunk_samples / 2) / sample_rate
time_points.append(time_s)
rms_values.append(rms)
return np.array(time_points), np.array(rms_values)
def print_statistics(rms_values: np.ndarray, threshold: float = 0.01):
"""打印RMS统计信息
Args:
rms_values: RMS值数组
threshold: 阈值
"""
print("\n" + "="*60)
print("RMS 统计分析")
print("="*60)
print(f"\n📊 基础统计:")
print(f" - 最小值: {np.min(rms_values):.6f}")
print(f" - 最大值: {np.max(rms_values):.6f}")
print(f" - 平均值: {np.mean(rms_values):.6f}")
print(f" - 中位数: {np.median(rms_values):.6f}")
print(f" - 标准差: {np.std(rms_values):.6f}")
print(f"\n📈 百分位数:")
for p in [10, 25, 50, 75, 90, 95, 99]:
value = np.percentile(rms_values, p)
print(f" - P{p:2d}: {value:.6f}")
print(f"\n🎯 阈值分析 (当前阈值: {threshold:.6f}):")
above_threshold = np.sum(rms_values >= threshold)
below_threshold = np.sum(rms_values < threshold)
total = len(rms_values)
print(f" - 超过阈值的帧数: {above_threshold} ({above_threshold/total*100:.1f}%)")
print(f" - 低于阈值的帧数: {below_threshold} ({below_threshold/total*100:.1f}%)")
print(f"\n💡 建议的阈值范围:")
# 基于非零RMS值的统计
non_zero_rms = rms_values[rms_values > 0.001]
if len(non_zero_rms) > 0:
p10 = np.percentile(non_zero_rms, 10)
p25 = np.percentile(non_zero_rms, 25)
mean = np.mean(non_zero_rms)
print(f" - 保守模式 (高灵敏度): {p10:.6f} (P10)")
print(f" - 宽松模式 (推荐): {p25:.6f} (P25)")
print(f" - 严格模式 (低误触): {mean*0.5:.6f} (平均值的50%)")
print("="*60 + "\n")
def plot_rms_timeline(time_points: np.ndarray, rms_values: np.ndarray,
threshold: float = 0.01, save_path: Optional[str] = None):
"""绘制RMS时序图
Args:
time_points: 时间点数组
rms_values: RMS值数组
threshold: 阈值线
save_path: 保存路径
"""
_, (ax1, ax2) = plt.subplots(2, 1, figsize=(14, 10))
# 上图: RMS时序
ax1.plot(time_points, rms_values, linewidth=1, label='RMS Energy', color='steelblue')
ax1.axhline(y=threshold, color='red', linestyle='--', linewidth=2,
label=f'阈值 = {threshold:.6f}')
# 标记超过阈值的区域
above_threshold = rms_values >= threshold
ax1.fill_between(time_points, 0, rms_values, where=above_threshold,
alpha=0.3, color='green', label='近场音频 (>= 阈值)')
ax1.fill_between(time_points, 0, rms_values, where=~above_threshold,
alpha=0.3, color='red', label='远场音频 (< 阈值)')
ax1.set_xlabel('时间 (秒)', fontsize=12)
ax1.set_ylabel('RMS 能量', fontsize=12)
ax1.set_title('音频 RMS 能量时序分析', fontsize=14, fontweight='bold')
ax1.legend(loc='upper right', fontsize=10)
ax1.grid(True, alpha=0.3)
ax1.set_ylim(bottom=0)
# 下图: RMS分布直方图
ax2.hist(rms_values, bins=100, color='steelblue', alpha=0.7, edgecolor='black')
ax2.axvline(x=threshold, color='red', linestyle='--', linewidth=2,
label=f'阈值 = {threshold:.6f}')
ax2.axvline(x=np.mean(rms_values), color='orange', linestyle=':', linewidth=2,
label=f'平均值 = {np.mean(rms_values):.6f}')
ax2.axvline(x=np.median(rms_values), color='green', linestyle=':', linewidth=2,
label=f'中位数 = {np.median(rms_values):.6f}')
ax2.set_xlabel('RMS 能量', fontsize=12)
ax2.set_ylabel('帧数', fontsize=12)
ax2.set_title('RMS 能量分布直方图', fontsize=14, fontweight='bold')
ax2.legend(loc='upper right', fontsize=10)
ax2.grid(True, alpha=0.3, axis='y')
plt.tight_layout()
if save_path:
plt.savefig(save_path, dpi=150, bbox_inches='tight')
print(f"✓ 图表已保存到: {save_path}")
plt.show()
def main():
parser = argparse.ArgumentParser(
description='音频RMS时序分析工具 - 帮助确定远场过滤阈值',
formatter_class=argparse.RawDescriptionHelpFormatter,
epilog="""
示例用法:
# 分析立体声音频(默认)
python analyze_audio_rms.py audio.wav
# 仅分析左声道
python analyze_audio_rms.py audio.wav --channel left
# 仅分析右声道
python analyze_audio_rms.py audio.wav --channel right
# 自定义阈值和分块大小
python analyze_audio_rms.py audio.wav --threshold 0.015 --chunk-size 160
# 保存图表
python analyze_audio_rms.py audio.wav --output rms_analysis.png
"""
)
parser.add_argument('audio_file', type=str,
help='音频文件路径 (支持 WAV, MP3, FLAC 等格式)')
parser.add_argument('--channel', type=str, choices=['stereo', 'left', 'right'],
default='stereo',
help='声道选择: stereo(立体声平均), left(左声道), right(右声道) [默认: stereo]')
parser.add_argument('--threshold', type=float, default=0.01,
help='RMS能量阈值 [默认: 0.01]')
parser.add_argument('--chunk-size', type=int, default=240,
help='分块大小(毫秒) [默认: 240ms,与流式ASR一致]')
parser.add_argument('--output', '-o', type=str, default=None,
help='保存图表的路径 (例如: output.png)')
parser.add_argument('--no-plot', action='store_true',
help='不显示图表,仅输出统计信息')
args = parser.parse_args()
# 检查文件是否存在
if not Path(args.audio_file).exists():
print(f"错误: 文件不存在: {args.audio_file}")
sys.exit(1)
print("="*60)
print("音频 RMS 时序分析工具")
print("="*60)
print(f"\n📁 文件: {args.audio_file}")
print(f"🎚️ 声道: {args.channel}")
print(f"📊 分块大小: {args.chunk_size}ms")
print(f"🎯 阈值: {args.threshold:.6f}")
print()
# 加载音频
print("正在加载音频...")
audio_data, sample_rate = load_audio(args.audio_file, args.channel)
duration = len(audio_data) / sample_rate
print(f"✓ 采样率: {sample_rate} Hz")
print(f"✓ 时长: {duration:.2f} 秒")
print(f"✓ 样本数: {len(audio_data)}")
# 分析RMS时序
print(f"\n正在分析 RMS 时序 (分块大小: {args.chunk_size}ms)...")
time_points, rms_values = analyze_rms_timeline(audio_data, sample_rate, args.chunk_size)
print(f"✓ 分析了 {len(rms_values)} 个音频块")
# 打印统计信息
print_statistics(rms_values, args.threshold)
# 绘制图表
if not args.no_plot:
print("正在生成图表...")
plot_rms_timeline(time_points, rms_values, args.threshold, args.output)
elif args.output:
print("正在保存图表...")
plot_rms_timeline(time_points, rms_values, args.threshold, args.output)
# 关闭显示窗口
plt.close()
if __name__ == '__main__':
main()