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()
Downloads last month

-

Downloads are not tracked for this model. How to track
Inference Providers NEW
This model isn't deployed by any Inference Provider. 🙋 Ask for provider support