Как преобразовать код Numpy в код Torch?Python

Программы на Python
Anonymous
Как преобразовать код Numpy в код Torch?

Сообщение Anonymous »

Я хочу преобразовать этот код из пиромакустики в pytorch. И хоть убей, я не могу этого сделать. Я не могу избавиться от этой ошибки, когда torch.where() возвращает пустой список.
def measure_rt60(h, fs=1, decay_db=60):
"""
Analyze the RT60 of an impulse response

Parameters
----------
h: array_likeThe impulse response.
fs: float or int, optional The sampling frequency of h (default
to 1, i.e., samples).
decay_db: float or int, optional
The decay in decibels for which we actually estimate the time.
"""

h = np.array(h)
fs = float(fs)

# The power of the impulse response in dB
power = h**2
energy = np.cumsum(power[::-1])[::-1] # Integration according to Schroeder

# remove the possibly all zero tail
i_nz = np.max(np.where(energy > 0)[0])
energy = energy[:i_nz]
energy_db = 10 * np.log10(energy)
energy_db -= energy_db[0]

# -5 dB headroom
i_5db = np.min(np.where(-5 - energy_db > 0)[0])
e_5db = energy_db[i_5db]
t_5db = i_5db / fs

# after decay
i_decay = np.min(np.where(-5 - decay_db - energy_db > 0)[0])
t_decay = i_decay / fs

# compute the decay time
decay_time = t_decay - t_5db
est_rt60 = (60 / decay_db) * decay_time

return est_rt60

Вот что у меня есть на данный момент. Проблема заключается в вычислении i_decay, где я получаю эту ошибку:
RuntimeError: min(): ожидаемое уменьшение dim должно быть указано для input.numel() == 0. Укажите уменьшение. dim с аргументом 'dim'.
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):
"""
get reverberation time for x dB (time for energy decay by x dB)
:param x: IR
:param fs: sampling frequency
:param decay_db: energy decay in dB
:return: idx in x where the energy decays by 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)


Подробнее здесь: https://stackoverflow.com/questions/783 ... torch-code

Вернуться в «Python»