File size: 6,460 Bytes
1b0b1dc
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
"""Zero-shot semantic listening with CLAP.

The DSP layer knows a band is 4 dB hot. It does not know the sound is a
reese. CLAP scores the audio against a bank of sound-design descriptors, so
the report can say "gritty distorted reese bass, over-compressed drums"
instead of only quoting numbers.

Loads in a background thread so the Space boots immediately, and fails soft:
if the model never arrives, everything else still works.
"""

from __future__ import annotations

import threading

import numpy as np

MODEL_ID = "laion/clap-htsat-unfused"
CLAP_SR = 48_000

# Grouped so the report can show one line per axis rather than a flat top-k.
BANK: dict[str, list[str]] = {
    "character": [
        "a gritty distorted reese bass",
        "a clean deep sine sub bass",
        "a metallic screaming growl bass",
        "an aggressive detuned saw lead",
        "a warm analog pad",
        "a plucky short synth stab",
        "a bright supersaw chord stack",
        "a wobbling filtered bass",
        "a soft mellow electric piano",
        "an acoustic guitar",
        "a male vocal",
        "a female vocal",
    ],
    "drums": [
        "a punchy tight kick drum",
        "a boomy undamped kick drum",
        "a sharp cracking snare",
        "a boxy resonant snare",
        "crisp hi hats",
        "a heavily compressed drum break",
        "a loose live drum kit",
    ],
    "problem": [
        "a muddy boomy cluttered mix",
        "a harsh sibilant painful mix",
        "an over-compressed lifeless mix",
        "a thin tinny weak mix",
        "a clipping distorted overloaded mix",
        "a clean balanced professional mix",
        "a hissy noisy recording",
        "a phasey hollow comb-filtered sound",
    ],
    "space": [
        "a dry close-miked sound with no reverb",
        "a tight small room reverb",
        "a huge cavernous hall reverb",
        "a long washed-out ambient reverb tail",
        "a slapback delay",
    ],
    "energy": [
        "a quiet sparse intro section",
        "a building tense riser",
        "a full loud drop section",
        "a calm breakdown section",
    ],
}

_FLAT: list[tuple[str, str]] = [(g, t) for g, items in BANK.items() for t in items]


def _features(raw, projection):
    """Normalise CLAP's feature output across transformers versions.

    4.x returns the projected tensor directly. 5.x returns a
    BaseModelOutputWithPooling, so the projection has to be applied here —
    which is exactly what 4.x did internally.
    """
    if hasattr(raw, "shape"):
        return raw
    pooled = getattr(raw, "pooler_output", None)
    if pooled is None:
        pooled = raw.last_hidden_state[:, 0]
    return projection(pooled)


class _Semantic:
    def __init__(self) -> None:
        self.ready = False
        self.error: str | None = None
        self._model = None
        self._processor = None
        self._text_emb = None
        self._lock = threading.Lock()

    def start(self) -> None:
        threading.Thread(target=self._load, daemon=True).start()

    def _load(self) -> None:
        try:
            import torch
            from transformers import ClapModel, ClapProcessor

            torch.set_num_threads(2)
            model = ClapModel.from_pretrained(MODEL_ID)
            model.eval()
            processor = ClapProcessor.from_pretrained(MODEL_ID)

            texts = [t for _, t in _FLAT]
            with torch.no_grad():
                inputs = processor(text=texts, return_tensors="pt", padding=True)
                emb = _features(model.get_text_features(**inputs), model.text_projection)
                emb = emb / emb.norm(dim=-1, keepdim=True)

            self._model, self._processor, self._text_emb = model, processor, emb
            self.ready = True
        except Exception as exc:  # noqa: BLE001 - fail soft, the app still works
            self.error = f"{type(exc).__name__}: {exc}"

    def status(self) -> str:
        if self.ready:
            return "ready"
        if self.error:
            return f"unavailable ({self.error})"
        return "warming up"

    def describe(self, mono48: np.ndarray, top_k: int = 2) -> dict[str, list[tuple[str, float]]]:
        """Score the clip against every descriptor, grouped by axis."""
        if not self.ready or mono48.size < CLAP_SR // 2:
            return {}

        import torch

        # CLAP was trained on 10 s windows; take the loudest one.
        want = CLAP_SR * 10
        if mono48.size > want:
            hop = CLAP_SR
            best_s, best_e = 0, -1.0
            for s in range(0, mono48.size - want + 1, hop):
                e = float(np.mean(mono48[s : s + want] ** 2))
                if e > best_e:
                    best_e, best_s = e, s
            mono48 = mono48[best_s : best_s + want]

        clip = mono48.astype(np.float32)
        with self._lock, torch.no_grad():
            # transformers 4.x takes `audios`, 5.x renamed it to `audio`.
            try:
                inputs = self._processor(audio=clip, sampling_rate=CLAP_SR,
                                         return_tensors="pt")
            except TypeError:
                inputs = self._processor(audios=clip, sampling_rate=CLAP_SR,
                                         return_tensors="pt")
            audio_emb = _features(self._model.get_audio_features(**inputs),
                                  self._model.audio_projection)
            audio_emb = audio_emb / audio_emb.norm(dim=-1, keepdim=True)
            sims = (audio_emb @ self._text_emb.T).squeeze(0).cpu().numpy()

        grouped: dict[str, list[tuple[str, float]]] = {}
        for group in BANK:
            idx = [i for i, (g, _) in enumerate(_FLAT) if g == group]
            local = sims[idx]
            # Softmax within the group — cross-group absolute scores are not
            # comparable, ranking inside a group is.
            e = np.exp((local - local.max()) * 20.0)
            probs = e / e.sum()
            order = np.argsort(-probs)[:top_k]
            grouped[group] = [(_FLAT[idx[o]][1], float(probs[o])) for o in order]
        return grouped


SEMANTIC = _Semantic()


def tags_line(grouped: dict[str, list[tuple[str, float]]], min_conf: float = 0.30) -> str:
    """Flatten the grouped scores into one readable sentence."""
    picks = [items[0][0] for items in grouped.values() if items and items[0][1] >= min_conf]
    return ", ".join(picks) if picks else ""