deeeed commited on
Commit
d926b92
·
verified ·
1 Parent(s): 0fe2bca

Upload wasm/runtime/sherpa-onnx-punctuation.js with huggingface_hub

Browse files
wasm/runtime/sherpa-onnx-punctuation.js ADDED
@@ -0,0 +1,150 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ /**
2
+ * sherpa-onnx-punctuation.js
3
+ *
4
+ * Online/Offline Punctuation 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
+ // --- OnlinePunctuation ---
17
+ // SherpaOnnxOnlinePunctuationModelConfig = { cnn_bilstm: ptr, bpe_vocab: ptr, num_threads: i32, debug: i32, provider: ptr }
18
+ // SherpaOnnxOnlinePunctuationConfig = { model: OnlinePunctuationModelConfig }
19
+
20
+ class OnlinePunctuation {
21
+ constructor(configObj, Module) {
22
+ const cnnBilstmStr = (configObj.model && configObj.model.cnnBilstm) || '';
23
+ const bpeVocabStr = (configObj.model && configObj.model.bpeVocab) || '';
24
+ const providerStr = (configObj.model && configObj.model.provider) || 'cpu';
25
+
26
+ const cnnLen = Module.lengthBytesUTF8(cnnBilstmStr) + 1;
27
+ const bpeLen = Module.lengthBytesUTF8(bpeVocabStr) + 1;
28
+ const provLen = Module.lengthBytesUTF8(providerStr) + 1;
29
+
30
+ const strBuf = Module._malloc(cnnLen + bpeLen + provLen);
31
+ let strOff = 0;
32
+ Module.stringToUTF8(cnnBilstmStr, strBuf + strOff, cnnLen); strOff += cnnLen;
33
+ Module.stringToUTF8(bpeVocabStr, strBuf + strOff, bpeLen); strOff += bpeLen;
34
+ Module.stringToUTF8(providerStr, strBuf + strOff, provLen);
35
+
36
+ // Flat struct layout:
37
+ // cnn_bilstm: ptr (4)
38
+ // bpe_vocab: ptr (4)
39
+ // num_threads: i32 (4)
40
+ // debug: i32 (4)
41
+ // provider: ptr (4)
42
+ // Total = 20 bytes
43
+ const ptr = Module._malloc(20);
44
+ let offset = 0;
45
+ Module.setValue(ptr + offset, strBuf, 'i8*'); offset += 4;
46
+ Module.setValue(ptr + offset, strBuf + cnnLen, 'i8*'); offset += 4;
47
+ Module.setValue(ptr + offset, (configObj.model && configObj.model.numThreads) || 1, 'i32'); offset += 4;
48
+ Module.setValue(ptr + offset, (configObj.model && configObj.model.debug) || 0, 'i32'); offset += 4;
49
+ Module.setValue(ptr + offset, strBuf + cnnLen + bpeLen, 'i8*'); offset += 4;
50
+
51
+ const handle = Module._SherpaOnnxCreateOnlinePunctuation(ptr);
52
+ Module._free(strBuf);
53
+ Module._free(ptr);
54
+
55
+ if (!handle) {
56
+ throw new Error('Failed to create online punctuation - null handle');
57
+ }
58
+
59
+ this.handle = handle;
60
+ this.Module = Module;
61
+ }
62
+
63
+ free() {
64
+ if (this.handle) {
65
+ this.Module._SherpaOnnxDestroyOnlinePunctuation(this.handle);
66
+ this.handle = 0;
67
+ }
68
+ }
69
+
70
+ /**
71
+ * Add punctuation to input text
72
+ * @param {string} text - Input text without punctuation
73
+ * @returns {string} - Text with punctuation added
74
+ */
75
+ addPunct(text) {
76
+ const textLen = this.Module.lengthBytesUTF8(text) + 1;
77
+ const textPtr = this.Module._malloc(textLen);
78
+ this.Module.stringToUTF8(text, textPtr, textLen);
79
+
80
+ const resultPtr = this.Module._SherpaOnnxOnlinePunctuationAddPunct(this.handle, textPtr);
81
+ this.Module._free(textPtr);
82
+
83
+ if (!resultPtr) return text;
84
+
85
+ const result = this.Module.UTF8ToString(resultPtr);
86
+ this.Module._SherpaOnnxOnlinePunctuationFreeText(resultPtr);
87
+ return result;
88
+ }
89
+ }
90
+
91
+ // --- Namespace API ---
92
+
93
+ SherpaOnnx.Punctuation = {
94
+ loadModel: async function(modelConfig) {
95
+ const debug = modelConfig.debug || false;
96
+ const modelDir = modelConfig.modelDir || 'punctuation-models';
97
+
98
+ if (debug) console.log(`Punctuation.loadModel: dir=${modelDir}`);
99
+
100
+ SherpaOnnx.FileSystem.removePath(modelDir, debug);
101
+
102
+ const files = [];
103
+ if (modelConfig.cnnBilstm) {
104
+ files.push({ url: modelConfig.cnnBilstm, filename: 'model.onnx' });
105
+ }
106
+ if (modelConfig.bpeVocab) {
107
+ files.push({ url: modelConfig.bpeVocab, filename: 'bpe.vocab' });
108
+ }
109
+
110
+ const result = await SherpaOnnx.FileSystem.prepareModelDirectory(files, modelDir, debug);
111
+ if (!result.success) {
112
+ throw new Error('Failed to load punctuation model files');
113
+ }
114
+
115
+ const modelFile = result.files.find(f => f.success && f.original.filename === 'model.onnx');
116
+ const vocabFile = result.files.find(f => f.success && f.original.filename === 'bpe.vocab');
117
+
118
+ return {
119
+ modelDir: result.modelDir,
120
+ modelPath: modelFile ? modelFile.path : `${result.modelDir}/model.onnx`,
121
+ vocabPath: vocabFile ? vocabFile.path : `${result.modelDir}/bpe.vocab`,
122
+ };
123
+ },
124
+
125
+ createPunctuation: function(loadedModel, options = {}) {
126
+ const debug = options.debug !== undefined ? options.debug : false;
127
+ const config = {
128
+ model: {
129
+ cnnBilstm: loadedModel.modelPath,
130
+ bpeVocab: loadedModel.vocabPath,
131
+ numThreads: options.numThreads || 1,
132
+ debug: debug ? 1 : 0,
133
+ provider: options.provider || 'cpu',
134
+ },
135
+ };
136
+
137
+ const punct = new OnlinePunctuation(config, global.Module);
138
+
139
+ if (SherpaOnnx.trackResource) {
140
+ SherpaOnnx.trackResource('punctuation', punct);
141
+ }
142
+
143
+ return punct;
144
+ },
145
+ };
146
+
147
+ if (typeof module !== 'undefined' && module.exports) {
148
+ module.exports = SherpaOnnx;
149
+ }
150
+ })(typeof window !== 'undefined' ? window : global);