将Pyroomacoustics的RT60计算代码转PyTorch时遇空列表报错求解
问题描述
我需要把pyroomacoustics中用于分析脉冲响应RT60的NumPy代码转换成PyTorch代码,但一直解决不了torch.where()返回空列表时的报错问题。
原NumPy代码
import numpy as np def measure_rt60(h, fs=1, decay_db=60): """ 分析脉冲响应的RT60值 参数 ---------- h: array_like 脉冲响应数据 fs: float 或 int, 可选 h的采样频率(默认值为1,即按样本数计算) decay_db: float 或 int, 可选 用于估算时间的衰减分贝值 """ h = np.array(h) fs = float(fs) # 计算脉冲响应的功率 power = h**2 # 按照Schroeder方法积分计算能量 energy = np.cumsum(power[::-1])[::-1] # 移除可能存在的全零尾部 i_nz = np.max(np.where(energy > 0)[0]) energy = energy[:i_nz] energy_db = 10 * np.log10(energy) energy_db -= energy_db[0] # 找到-5dB对应的位置 i_5db = np.min(np.where(-5 - energy_db > 0)[0]) e_5db = energy_db[i_5db] t_5db = i_5db / fs # 找到衰减后的目标位置 i_decay = np.min(np.where(-5 - decay_db - energy_db > 0)[0]) t_decay = i_decay / fs # 计算衰减时间和估算的RT60 decay_time = t_decay - t_5db est_rt60 = (60 / decay_db) * decay_time return est_rt60
转换后的PyTorch代码及报错
转换后的代码在计算i_decay时触发报错:RuntimeError: min(): Expected reduction dim to be specified for input.numel() == 0. Specify the reduction dim with the 'dim' argument.
import torch def measure_rt60_torch(h, fs=1, decay_db=60): fs = float(fs) decay_db = float(decay_db) power = h**2 energy = torch.cumsum(power.flip(-1), -1).flip(-1) i_nz = torch.max(torch.nonzero(energy > 0)[-1]) energy = energy[:i_nz] energy_db = 10 * torch.log10(energy) energy_db_adjusted = energy_db.clone() energy_db_adjusted -= energy_db[0] i_5db = torch.min(torch.nonzero(torch.tensor(-5.) - energy_db_adjusted > torch.tensor(0.), as_tuple=True)[0]) e_5db = energy_db_adjusted[i_5db] t_5db = i_5db / fs i_decay = torch.min(torch.nonzero(torch.tensor(-5.) - torch.tensor(decay_db) - energy_db_adjusted > torch.tensor(0.), as_tuple=True)[0]) t_decay = i_decay / fs decay_time = t_decay - t_5db est_rt60 = (60 / decay_db) * decay_time return est_rt60.item() def measure_rtX(x, fs=48000, decay_db=60): """ 获取x分贝的混响时间(能量衰减x分贝所需的时间) :param x: 脉冲响应(IR) :param fs: 采样频率 :param decay_db: 能量衰减的分贝值 :return: x中能量衰减decay_db对应的索引 """ wrapper_list = [] iteration=0 for batch_idx in range(x.size(0)): batch_x = x[batch_idx] print(iteration,batch_x.size()) iteration+=1 rtX = -1 while rtX == -1: try: rtX = measure_rt60_torch(batch_x, fs, decay_db) except ValueError: if decay_db > 10: decay_db -= 10 else: rtX = batch_x.size(0) / fs wrapper_list.append(rtX) return wrapper_list x_tensor = torch.randn(32, 96000) x_tensor = x_tensor.float() print(x_tensor.size()) list_wrapper = measure_rtX(x_tensor, fs=48000)
解决方案
报错核心原因是:PyTorch中torch.min()处理空张量会直接报错,而NumPy的np.min()处理空数组会返回inf。需要在调用torch.min()前先检查索引张量是否为空,同时优化张量创建逻辑避免类型/设备不匹配。
修改后的代码如下:
import torch def measure_rt60_torch(h, fs=1, decay_db=60): fs = float(fs) # 统一设备和数据类型,避免不匹配问题 decay_db = torch.tensor(decay_db, dtype=h.dtype, device=h.device) neg_5 = torch.tensor(-5., dtype=h.dtype, device=h.device) power = h ** 2 # 计算Schroeder能量 energy = torch.cumsum(power.flip(-1), dim=-1).flip(-1) # 处理全零尾部:获取最后一个非零能量的索引 non_zero_mask = energy > 0 non_zero_indices = torch.nonzero(non_zero_mask, as_tuple=True)[0] if len(non_zero_indices) == 0: # 极端情况:整个能量数组都是0,返回最大可能时间 return h.size(-1) / fs i_nz = torch.max(non_zero_indices) energy = energy[:i_nz] energy_db = 10 * torch.log10(energy) energy_db_adjusted = energy_db - energy_db[0] # 处理-5dB位置 mask_5db = (neg_5 - energy_db_adjusted) > 0 idx_5db = torch.nonzero(mask_5db, as_tuple=True)[0] if len(idx_5db) == 0: # 未找到-5dB位置,返回当前能量的最大时间 return i_nz / fs i_5db = torch.min(idx_5db) t_5db = i_5db / fs # 处理衰减目标位置 target_db = neg_5 - decay_db mask_decay = (target_db - energy_db_adjusted) > 0 idx_decay = torch.nonzero(mask_decay, as_tuple=True)[0] if len(idx_decay) == 0: # 无法达到指定衰减分贝,返回当前能量的最大时间 return i_nz / fs i_decay = torch.min(idx_decay) t_decay = i_decay / fs decay_time = t_decay - t_5db est_rt60 = (60 / decay_db.item()) * decay_time return est_rt60.item() def measure_rtX(x, fs=48000, decay_db=60): """ 获取x分贝的混响时间(能量衰减x分贝所需的时间) :param x: 脉冲响应(IR) :param fs: 采样频率 :param decay_db: 能量衰减的分贝值 :return: 每个样本的混响时间列表 """ wrapper_list = [] for batch_idx in range(x.size(0)): batch_x = x[batch_idx] current_decay = decay_db rtX = -1 while rtX == -1: try: rtX = measure_rt60_torch(batch_x, fs, current_decay) except (RuntimeError, ValueError): if current_decay > 10: current_decay -= 10 else: rtX = batch_x.size(0) / fs wrapper_list.append(rtX) return wrapper_list
关键修改点:
- 每次调用
torch.nonzero()后,先检查返回的索引张量是否为空,为空则返回合理默认值 - 将固定值张量的创建统一到函数开头,确保与输入张量的设备、数据类型一致
- 扩展异常捕获类型,覆盖可能出现的
RuntimeError和ValueError
内容的提问来源于stack exchange,提问作者mahnoor.fatima
相关产品推荐
相关产品推荐

