You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.06.25 18:14:53