File size: 5,211 Bytes
79a6851
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
30ad92d
79a6851
 
 
30ad92d
79a6851
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
30ad92d
79a6851
30ad92d
79a6851
 
 
 
30ad92d
 
 
 
 
 
 
 
79a6851
 
 
 
 
959a270
 
 
 
 
79a6851
 
 
 
 
30ad92d
35704e7
 
 
 
 
 
 
 
 
 
 
30ad92d
 
79a6851
 
30ad92d
79a6851
 
 
 
 
30ad92d
79a6851
30ad92d
 
 
 
 
 
 
 
 
 
79a6851
 
30ad92d
 
 
 
 
 
 
 
 
 
 
 
79a6851
 
 
 
 
30ad92d
 
79a6851
 
30ad92d
79a6851
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
import sys
sys.stdout.reconfigure(line_buffering=True)

try:
    import spaces
except ImportError:
    # keep @spaces.GPU usable as a no-op; ZeroGPU requires this exact name.
    class spaces:
        class GPU:
            def __init__(self, func=None, duration=60):
                self.func = func

            def __call__(self, *args, **kwargs):
                if self.func is not None:
                    return self.func(*args, **kwargs)
                func = args[0]
                return func

from types import SimpleNamespace

import gradio as gr
import librosa
import numpy as np
import pandas as pd
import torch

from pyharp import ModelCard, AudioLabel, LabelList, build_endpoint
from model.htsat import HTSAT_Swin_Transformer

DEVICE = "cuda" if torch.cuda.is_available() else "cpu"

# AudioSet configuration values from HTS-AT-repo/config.py
CONFIG = SimpleNamespace(
    sample_rate=32000,
    window_size=1024,
    hop_size=320,
    mel_bins=64,
    fmin=50,
    fmax=14000,
    enable_tscam=True,
    htsat_attn_heatmap=False,
    loss_type="clip_bce",
    enable_repeat_mode=False,
)

model = HTSAT_Swin_Transformer(config=CONFIG)
checkpoint = torch.load("HTSAT_AudioSet_Saved_1.ckpt", map_location="cpu")
state_dict = {key.replace("sed_model.", ""): value for key, value in checkpoint["state_dict"].items()}
model.load_state_dict(state_dict)
model.eval()
model_ready = False  # has model been moved onto the device yet?

# The model always produces this many framewise steps, regardless of clip length,
# by stretching shorter clips onto this fixed grid (so frame i maps linearly onto
# real time as i / MODEL_FRAMES * clip_duration). Longer clips instead use a
# sliding-window path that averages overlapping crops together, which no longer
# maps onto real time at all - framewise localization is only meaningful below
# this length until that path is implemented separately.
MODEL_FRAMES = model.spec_size * model.freq_ratio

labels_df = pd.read_csv("class_label_indice.csv")
idx_to_label = dict(zip(labels_df["index"], labels_df["display_name"]))

model_card = ModelCard(
    name="HTS-AT",
    description=(
        "Tags the sounds present in an audio clip across 527 AudioSet classes "
        "(speech, music, instruments, animals, environmental sounds, etc). "
        "Clips longer than 10 seconds are not currently supported."
    ),
    author="Ke Chen, Xingjian Du, Bilei Zhu, Zejun Ma, Taylor Berg-Kirkpatrick, Shlomo Dubnov",
    tags=["classification", "tagging", "audioset"],
)


def find_active_segments(frame_scores: np.ndarray, threshold: float) -> list[tuple[int, int]]:
    segments = []
    start = None
    for i, score in enumerate(frame_scores):
        if score > threshold and start is None:
            start = i
        elif score <= threshold and start is not None:
            segments.append((start, i))
            start = None
    if start is not None:
        segments.append((start, len(frame_scores)))
    return segments


@spaces.GPU
@torch.inference_mode()
def process_fn(input_audio_path: str, detection_threshold: float):
    global model_ready
    if not model_ready:
        model.to(DEVICE)
        model_ready = True

    waveform, _ = librosa.load(input_audio_path, sr=CONFIG.sample_rate)
    duration = len(waveform) / CONFIG.sample_rate
    audio_tensor = torch.from_numpy(waveform).float().to(DEVICE)[None, :]

    mel_frames = model.logmel_extractor(model.spectrogram_extractor(audio_tensor)).shape[2]
    if mel_frames > MODEL_FRAMES:
        raise gr.Error(
            f"Clip is too long ({duration:.1f}s). This model currently supports "
            f"clips up to about {MODEL_FRAMES * CONFIG.hop_size / CONFIG.sample_rate:.1f}s."
        )

    framewise_output = model(audio_tensor, None, False)["framewise_output"][0].cpu().numpy()

    label_list = LabelList()
    for class_idx in range(framewise_output.shape[1]):
        frame_scores = framewise_output[:, class_idx]
        for start, end in find_active_segments(frame_scores, detection_threshold):
            peak_confidence = frame_scores[start:end].max()
            label_list.append(AudioLabel(
                t=start / MODEL_FRAMES * duration,
                duration=(end - start) / MODEL_FRAMES * duration,
                label=idx_to_label[class_idx],
                description=f"confidence {peak_confidence:.0%}",
                amplitude=float(peak_confidence),
            ))
    return input_audio_path, label_list


with gr.Blocks() as demo:
    input_components = [
        gr.Audio(type="filepath", label="Input Audio").harp_required(True),
        gr.Slider(minimum=0.0, maximum=1.0, step=0.05, value=0.5, label="Detection Threshold",
                  info="Minimum confidence for a sound to be tagged"),
    ]
    output_components = [
        gr.Audio(type="filepath", label="Output Audio").set_info("Input audio, unchanged."),
        gr.JSON(label="Detected Sounds").set_info("Top predicted sound tags with confidence scores."),
    ]

    build_endpoint(
        model_card=model_card,
        input_components=input_components,
        output_components=output_components,
        process_fn=process_fn,
    )

demo.queue().launch(pwa=True)