deeeed commited on
Commit
7d78416
·
verified ·
1 Parent(s): 600c1e3

Upload wasm/runtime/sherpa-onnx-audio-tagging.js with huggingface_hub

Browse files
wasm/runtime/sherpa-onnx-audio-tagging.js ADDED
@@ -0,0 +1,229 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /**
2
+ * sherpa-onnx-audio-tagging.js
3
+ *
4
+ * Audio Tagging functionality for SherpaOnnx
5
+ * Requires sherpa-onnx-core.js to be loaded first
6
+ */
7
+
8
+ (function(global) {
9
+ if (!global.SherpaOnnx) {
10
+ console.error('SherpaOnnx namespace not found. Make sure to load sherpa-onnx-core.js first.');
11
+ return;
12
+ }
13
+
14
+ const SherpaOnnx = global.SherpaOnnx;
15
+
16
+ // --- WASM struct init helpers ---
17
+
18
+ // SherpaOnnxOfflineZipformerAudioTaggingModelConfig = { model: ptr }
19
+ // SherpaOnnxAudioTaggingModelConfig = { zipformer: { model: ptr }, ced: ptr, num_threads: i32, debug: i32, provider: ptr }
20
+ // SherpaOnnxAudioTaggingConfig = { model: AudioTaggingModelConfig, labels: ptr, top_k: i32 }
21
+
22
+ function initAudioTaggingConfig(config, Module) {
23
+ // zipformer.model string
24
+ const zipformerModelStr = (config.model && config.model.zipformer && config.model.zipformer.model) || '';
25
+ const zipformerModelLen = Module.lengthBytesUTF8(zipformerModelStr) + 1;
26
+ const zipformerModelBuf = Module._malloc(zipformerModelLen);
27
+ Module.stringToUTF8(zipformerModelStr, zipformerModelBuf, zipformerModelLen);
28
+
29
+ // ced string
30
+ const cedStr = (config.model && config.model.ced) || '';
31
+ const cedLen = Module.lengthBytesUTF8(cedStr) + 1;
32
+ const cedBuf = Module._malloc(cedLen);
33
+ Module.stringToUTF8(cedStr, cedBuf, cedLen);
34
+
35
+ // provider string
36
+ const providerStr = (config.model && config.model.provider) || 'cpu';
37
+ const providerLen = Module.lengthBytesUTF8(providerStr) + 1;
38
+ const providerBuf = Module._malloc(providerLen);
39
+ Module.stringToUTF8(providerStr, providerBuf, providerLen);
40
+
41
+ // labels string
42
+ const labelsStr = config.labels || '';
43
+ const labelsLen = Module.lengthBytesUTF8(labelsStr) + 1;
44
+ const labelsBuf = Module._malloc(labelsLen);
45
+ Module.stringToUTF8(labelsStr, labelsBuf, labelsLen);
46
+
47
+ // Build the flat struct:
48
+ // SherpaOnnxAudioTaggingConfig =
49
+ // zipformer.model: ptr (4)
50
+ // ced: ptr (4)
51
+ // num_threads: i32 (4)
52
+ // debug: i32 (4)
53
+ // provider: ptr (4)
54
+ // labels: ptr (4)
55
+ // top_k: i32 (4)
56
+ // Total = 28 bytes
57
+ const totalLen = 7 * 4;
58
+ const ptr = Module._malloc(totalLen);
59
+ let offset = 0;
60
+
61
+ Module.setValue(ptr + offset, zipformerModelBuf, 'i8*'); offset += 4; // zipformer.model
62
+ Module.setValue(ptr + offset, cedBuf, 'i8*'); offset += 4; // ced
63
+ Module.setValue(ptr + offset, (config.model && config.model.numThreads) || 1, 'i32'); offset += 4;
64
+ Module.setValue(ptr + offset, (config.model && config.model.debug) || 0, 'i32'); offset += 4;
65
+ Module.setValue(ptr + offset, providerBuf, 'i8*'); offset += 4; // provider
66
+ Module.setValue(ptr + offset, labelsBuf, 'i8*'); offset += 4; // labels
67
+ Module.setValue(ptr + offset, config.topK || 5, 'i32'); offset += 4; // top_k
68
+
69
+ return {
70
+ ptr: ptr,
71
+ buffers: [zipformerModelBuf, cedBuf, providerBuf, labelsBuf],
72
+ };
73
+ }
74
+
75
+ function freeAudioTaggingConfig(config, Module) {
76
+ for (const buf of config.buffers) {
77
+ Module._free(buf);
78
+ }
79
+ Module._free(config.ptr);
80
+ }
81
+
82
+ // --- AudioTagging class ---
83
+
84
+ class AudioTagging {
85
+ constructor(configObj, Module) {
86
+ const config = initAudioTaggingConfig(configObj, Module);
87
+ const handle = Module._SherpaOnnxCreateAudioTagging(config.ptr);
88
+ freeAudioTaggingConfig(config, Module);
89
+
90
+ if (!handle) {
91
+ throw new Error('Failed to create audio tagging - null handle');
92
+ }
93
+
94
+ this.handle = handle;
95
+ this.Module = Module;
96
+ this.topK = configObj.topK || 5;
97
+ }
98
+
99
+ free() {
100
+ if (this.handle) {
101
+ this.Module._SherpaOnnxDestroyAudioTagging(this.handle);
102
+ this.handle = 0;
103
+ }
104
+ }
105
+
106
+ createStream() {
107
+ const streamHandle = this.Module._SherpaOnnxAudioTaggingCreateOfflineStream(this.handle);
108
+ if (!streamHandle) {
109
+ throw new Error('Failed to create audio tagging offline stream');
110
+ }
111
+ return streamHandle;
112
+ }
113
+
114
+ /**
115
+ * Feed audio samples into a stream
116
+ * @param {number} stream - Stream handle
117
+ * @param {number} sampleRate - Sample rate
118
+ * @param {Float32Array} samples - Audio samples
119
+ */
120
+ acceptWaveform(stream, sampleRate, samples) {
121
+ const pointer = this.Module._malloc(samples.length * samples.BYTES_PER_ELEMENT);
122
+ this.Module.HEAPF32.set(samples, pointer / samples.BYTES_PER_ELEMENT);
123
+ this.Module._SherpaOnnxAcceptWaveformOffline(stream, sampleRate, pointer, samples.length);
124
+ this.Module._free(pointer);
125
+ }
126
+
127
+ /**
128
+ * Compute audio tags
129
+ * @param {number} stream - Stream handle
130
+ * @param {number} topK - Number of top events to return (-1 for default)
131
+ * @returns {Array<{name: string, index: number, prob: number}>}
132
+ */
133
+ compute(stream, topK) {
134
+ const k = topK !== undefined ? topK : -1;
135
+ const resultsPtr = this.Module._SherpaOnnxAudioTaggingCompute(this.handle, stream, k);
136
+
137
+ if (!resultsPtr) {
138
+ // Free the stream
139
+ this.Module._SherpaOnnxDestroyOfflineStream(stream);
140
+ return [];
141
+ }
142
+
143
+ const events = [];
144
+ // resultsPtr is a pointer to an array of pointers to SherpaOnnxAudioEvent
145
+ // Each SherpaOnnxAudioEvent = { name: ptr (4), index: i32 (4), prob: f32 (4) } = 12 bytes
146
+ let i = 0;
147
+ while (true) {
148
+ const eventPtr = this.Module.HEAP32[(resultsPtr / 4) + i];
149
+ if (!eventPtr) break; // NULL terminator
150
+
151
+ const namePtr = this.Module.HEAP32[eventPtr / 4];
152
+ const index = this.Module.HEAP32[eventPtr / 4 + 1];
153
+ const prob = this.Module.HEAPF32[eventPtr / 4 + 2];
154
+
155
+ const name = namePtr ? this.Module.UTF8ToString(namePtr) : '';
156
+ events.push({ name, index, prob });
157
+ i++;
158
+ }
159
+
160
+ this.Module._SherpaOnnxAudioTaggingFreeResults(resultsPtr);
161
+ this.Module._SherpaOnnxDestroyOfflineStream(stream);
162
+
163
+ return events;
164
+ }
165
+ }
166
+
167
+ // --- Namespace API ---
168
+
169
+ SherpaOnnx.AudioTagging = {
170
+ loadModel: async function(modelConfig) {
171
+ const debug = modelConfig.debug || false;
172
+ const modelDir = modelConfig.modelDir || 'audio-tagging-models';
173
+
174
+ if (debug) console.log(`AudioTagging.loadModel: dir=${modelDir}`);
175
+
176
+ SherpaOnnx.FileSystem.removePath(modelDir, debug);
177
+
178
+ const files = [];
179
+
180
+ if (modelConfig.ced) {
181
+ files.push({ url: modelConfig.ced, filename: 'model.onnx' });
182
+ }
183
+ if (modelConfig.labels) {
184
+ files.push({ url: modelConfig.labels, filename: 'labels.txt' });
185
+ }
186
+
187
+ const result = await SherpaOnnx.FileSystem.prepareModelDirectory(files, modelDir, debug);
188
+ if (!result.success) {
189
+ throw new Error('Failed to load audio tagging model files');
190
+ }
191
+
192
+ const modelFile = result.files.find(f => f.success && f.original.filename === 'model.onnx');
193
+ const labelsFile = result.files.find(f => f.success && f.original.filename === 'labels.txt');
194
+
195
+ return {
196
+ modelDir: result.modelDir,
197
+ modelPath: modelFile ? modelFile.path : `${result.modelDir}/model.onnx`,
198
+ labelsPath: labelsFile ? labelsFile.path : `${result.modelDir}/labels.txt`,
199
+ };
200
+ },
201
+
202
+ createAudioTagging: function(loadedModel, options = {}) {
203
+ const debug = options.debug !== undefined ? options.debug : false;
204
+ const config = {
205
+ model: {
206
+ zipformer: { model: '' },
207
+ ced: loadedModel.modelPath,
208
+ numThreads: options.numThreads || 1,
209
+ debug: debug ? 1 : 0,
210
+ provider: options.provider || 'cpu',
211
+ },
212
+ labels: loadedModel.labelsPath,
213
+ topK: options.topK || 5,
214
+ };
215
+
216
+ const tagger = new AudioTagging(config, global.Module);
217
+
218
+ if (SherpaOnnx.trackResource) {
219
+ SherpaOnnx.trackResource('audioTagging', tagger);
220
+ }
221
+
222
+ return tagger;
223
+ },
224
+ };
225
+
226
+ if (typeof module !== 'undefined' && module.exports) {
227
+ module.exports = SherpaOnnx;
228
+ }
229
+ })(typeof window !== 'undefined' ? window : global);