ONNX Model Conversion Script:
import torch
import torch.nn as nn
import warnings
from deeprhythm.model.infer import load_cnn_model
from deeprhythm.audio_proc.hcqm import make_kernels
warnings.filterwarnings("ignore", category=torch.jit.TracerWarning)
warnings.filterwarnings("ignore", category=UserWarning)
class EndToEndDeepRhythm(nn.Module):
def __init__(self, cnn_model, device='cpu'):
super().__init__()
self.cnn_model = cnn_model
sr = 22050
len_audio = sr * 8
self.stft, self.band, self.cqt_specs = make_kernels(len_audio, sr, device=device)
self.cqt_layers = nn.ModuleList(self.cqt_specs)
self.register_buffer('band_filter', self.band)
def forward(self, audio_clip):
# --- HCQM ---
# 1. STFT
stft_out = self.stft(audio_clip)
# 2. Log Filter
# stft.transpose(1, 2) -> matmul -> transpose(1, 2)
stft_transposed = stft_out.transpose(1, 2)
filtered_transposed = torch.matmul(stft_transposed, self.band_filter.transpose(0, 1))
stft_bands = filtered_transposed.transpose(1, 2)
# 3. Flatten [batch, 8, time] -> [batch*8, time]
batch_size = stft_out.shape[0]
stft_bands_flat = stft_bands.reshape(batch_size * 8, -1)
# 4. Onset Strength
# Log magnitude
log_spec = torch.log10(torch.clamp(stft_bands_flat, min=1e-10)) * 20.0
# Diff (lag=1)
onset_env = log_spec[..., 1:] - log_spec[..., :-1]
onset_env = torch.clamp(onset_env, min=0.0)
pad_len = stft_bands_flat.shape[-1] - onset_env.shape[-1]
if pad_len > 0:
onset_env = torch.nn.functional.pad(onset_env, (pad_len, 0), "constant", 0)
# 5. CQT
hcqm_list = []
for cqt_layer in self.cqt_layers:
out = cqt_layer(onset_env)
# Mean over time
out_mean = out.mean(dim=-1)
hcqm_list.append(out_mean)
# Stack & Reshape
# [6, batch*8, 240] -> [batch*8, 240, 6]
hcqm_stacked = torch.stack(hcqm_list, dim=-1)
# [batch*8, 240, 6] -> [batch, 8, 240, 6]
hcqm_reshaped = hcqm_stacked.reshape(batch_size, 8, 240, 6)
# CNN Input format: [batch, 6, 240, 8]
# Permute: (batch, harmonics, bins, bands)
model_input = hcqm_reshaped.permute(0, 3, 2, 1)
logits = self.cnn_model(model_input)
probs = torch.softmax(logits, dim=1)
return probs
def export():
device = 'cpu'
print("Loading original model...")
cnn = load_cnn_model(device=device)
print("Assembling End-to-End model...")
full_model = EndToEndDeepRhythm(cnn, device=device)
full_model.eval()
# Dummy input
dummy_input = torch.randn(1, 22050 * 8).to(device)
print("Exporting to ONNX (Opset 18)...")
torch.onnx.export(
full_model,
dummy_input,
"deeprhythm_web.onnx",
export_params=True,
opset_version=18,
do_constant_folding=True,
input_names=['audio_pcm'],
output_names=['bpm_probs'],
dynamic_axes={
'audio_pcm': {0: 'batch_size'},
'bpm_probs': {0: 'batch_size'}
}
)
print("Export success: deeprhythm_web.onnx")
if __name__ == "__main__":
export()
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support