parakeet.ts 1.0.0
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- package/LICENSE +21 -0
- package/README.md +286 -0
- package/dist/alignment.d.ts +36 -0
- package/dist/alignment.d.ts.map +1 -0
- package/dist/alignment.js +202 -0
- package/dist/alignment.js.map +1 -0
- package/dist/audio.d.ts +59 -0
- package/dist/audio.d.ts.map +1 -0
- package/dist/audio.js +324 -0
- package/dist/audio.js.map +1 -0
- package/dist/backend.d.ts +61 -0
- package/dist/backend.d.ts.map +1 -0
- package/dist/backend.js +14 -0
- package/dist/backend.js.map +1 -0
- package/dist/decode.d.ts +38 -0
- package/dist/decode.d.ts.map +1 -0
- package/dist/decode.js +124 -0
- package/dist/decode.js.map +1 -0
- package/dist/hub.d.ts +18 -0
- package/dist/hub.d.ts.map +1 -0
- package/dist/hub.js +86 -0
- package/dist/hub.js.map +1 -0
- package/dist/index.d.ts +13 -0
- package/dist/index.d.ts.map +1 -0
- package/dist/index.js +19 -0
- package/dist/index.js.map +1 -0
- package/dist/load.d.ts +51 -0
- package/dist/load.d.ts.map +1 -0
- package/dist/load.js +34 -0
- package/dist/load.js.map +1 -0
- package/dist/mlx/attention.d.ts +61 -0
- package/dist/mlx/attention.d.ts.map +1 -0
- package/dist/mlx/attention.js +330 -0
- package/dist/mlx/attention.js.map +1 -0
- package/dist/mlx/audio.d.ts +46 -0
- package/dist/mlx/audio.d.ts.map +1 -0
- package/dist/mlx/audio.js +309 -0
- package/dist/mlx/audio.js.map +1 -0
- package/dist/mlx/backend.d.ts +22 -0
- package/dist/mlx/backend.d.ts.map +1 -0
- package/dist/mlx/backend.js +67 -0
- package/dist/mlx/backend.js.map +1 -0
- package/dist/mlx/cache.d.ts +23 -0
- package/dist/mlx/cache.d.ts.map +1 -0
- package/dist/mlx/cache.js +98 -0
- package/dist/mlx/cache.js.map +1 -0
- package/dist/mlx/cli.d.ts +12 -0
- package/dist/mlx/cli.d.ts.map +1 -0
- package/dist/mlx/cli.js +145 -0
- package/dist/mlx/cli.js.map +1 -0
- package/dist/mlx/conformer.d.ts +81 -0
- package/dist/mlx/conformer.d.ts.map +1 -0
- package/dist/mlx/conformer.js +315 -0
- package/dist/mlx/conformer.js.map +1 -0
- package/dist/mlx/index.d.ts +6 -0
- package/dist/mlx/index.d.ts.map +1 -0
- package/dist/mlx/index.js +9 -0
- package/dist/mlx/index.js.map +1 -0
- package/dist/mlx/load.d.ts +27 -0
- package/dist/mlx/load.d.ts.map +1 -0
- package/dist/mlx/load.js +161 -0
- package/dist/mlx/load.js.map +1 -0
- package/dist/mlx/nn.d.ts +121 -0
- package/dist/mlx/nn.d.ts.map +1 -0
- package/dist/mlx/nn.js +511 -0
- package/dist/mlx/nn.js.map +1 -0
- package/dist/mlx/rnnt.d.ts +59 -0
- package/dist/mlx/rnnt.d.ts.map +1 -0
- package/dist/mlx/rnnt.js +233 -0
- package/dist/mlx/rnnt.js.map +1 -0
- package/dist/mlx/server.d.ts +34 -0
- package/dist/mlx/server.d.ts.map +1 -0
- package/dist/mlx/server.js +115 -0
- package/dist/mlx/server.js.map +1 -0
- package/dist/mlx/utils.d.ts +19 -0
- package/dist/mlx/utils.d.ts.map +1 -0
- package/dist/mlx/utils.js +103 -0
- package/dist/mlx/utils.js.map +1 -0
- package/dist/model.d.ts +112 -0
- package/dist/model.d.ts.map +1 -0
- package/dist/model.js +196 -0
- package/dist/model.js.map +1 -0
- package/dist/onnx/backend.d.ts +55 -0
- package/dist/onnx/backend.d.ts.map +1 -0
- package/dist/onnx/backend.js +111 -0
- package/dist/onnx/backend.js.map +1 -0
- package/dist/onnx/index.d.ts +7 -0
- package/dist/onnx/index.d.ts.map +1 -0
- package/dist/onnx/index.js +6 -0
- package/dist/onnx/index.js.map +1 -0
- package/dist/onnx/parakeet.d.ts +44 -0
- package/dist/onnx/parakeet.d.ts.map +1 -0
- package/dist/onnx/parakeet.js +135 -0
- package/dist/onnx/parakeet.js.map +1 -0
- package/dist/tokenizer.d.ts +6 -0
- package/dist/tokenizer.d.ts.map +1 -0
- package/dist/tokenizer.js +8 -0
- package/dist/tokenizer.js.map +1 -0
- package/package.json +85 -0
|
@@ -0,0 +1,315 @@
|
|
|
1
|
+
import { MxArray } from '@mlx-node/core';
|
|
2
|
+
import { Module, Linear, LayerNorm, BatchNorm, Conv1d, Conv2d, silu, relu, glu, } from './nn.js';
|
|
3
|
+
import { MultiHeadAttention, RelPositionMultiHeadAttention, RelPositionMultiHeadLocalAttention, RelPositionalEncoding, LocalRelPositionalEncoding, } from './attention.js';
|
|
4
|
+
function s(...dims) {
|
|
5
|
+
return BigInt64Array.from(dims.map(BigInt));
|
|
6
|
+
}
|
|
7
|
+
// ---------------------------------------------------------------------------
|
|
8
|
+
// Feed-forward block
|
|
9
|
+
// ---------------------------------------------------------------------------
|
|
10
|
+
class FeedForward extends Module {
|
|
11
|
+
linear1;
|
|
12
|
+
linear2;
|
|
13
|
+
constructor(dModel, dFf, useBias) {
|
|
14
|
+
super();
|
|
15
|
+
this.linear1 = new Linear(dModel, dFf, useBias);
|
|
16
|
+
this.linear2 = new Linear(dFf, dModel, useBias);
|
|
17
|
+
}
|
|
18
|
+
forward(x) {
|
|
19
|
+
return this.linear2.forward(silu(this.linear1.forward(x)));
|
|
20
|
+
}
|
|
21
|
+
loadWeights(weights, prefix) {
|
|
22
|
+
this.linear1.loadWeights(weights, `${prefix}.linear1`);
|
|
23
|
+
this.linear2.loadWeights(weights, `${prefix}.linear2`);
|
|
24
|
+
}
|
|
25
|
+
}
|
|
26
|
+
// ---------------------------------------------------------------------------
|
|
27
|
+
// Convolution block
|
|
28
|
+
// ---------------------------------------------------------------------------
|
|
29
|
+
class Convolution extends Module {
|
|
30
|
+
padding;
|
|
31
|
+
pointwiseConv1;
|
|
32
|
+
depthwiseConv;
|
|
33
|
+
batchNorm;
|
|
34
|
+
pointwiseConv2;
|
|
35
|
+
constructor(args) {
|
|
36
|
+
super();
|
|
37
|
+
const useBias = args.useBias ?? true;
|
|
38
|
+
this.padding = Math.floor((args.convKernelSize - 1) / 2);
|
|
39
|
+
this.pointwiseConv1 = new Conv1d(args.dModel, args.dModel * 2, 1, 1, 0, 1, useBias);
|
|
40
|
+
this.depthwiseConv = new Conv1d(args.dModel, args.dModel, args.convKernelSize, 1, 0, args.dModel, useBias);
|
|
41
|
+
this.batchNorm = new BatchNorm(args.dModel);
|
|
42
|
+
this.pointwiseConv2 = new Conv1d(args.dModel, args.dModel, 1, 1, 0, 1, useBias);
|
|
43
|
+
}
|
|
44
|
+
forward(x, cache) {
|
|
45
|
+
x = this.pointwiseConv1.forward(x);
|
|
46
|
+
x = glu(x, 2); // split along last axis, gate with sigmoid
|
|
47
|
+
if (cache !== null) {
|
|
48
|
+
x = cache.updateAndFetchConv(x, this.padding);
|
|
49
|
+
}
|
|
50
|
+
else {
|
|
51
|
+
x = x.pad(new Int32Array([0, 0, this.padding, this.padding, 0, 0]), 0.0);
|
|
52
|
+
}
|
|
53
|
+
x = this.depthwiseConv.forward(x);
|
|
54
|
+
x = this.batchNorm.forward(x);
|
|
55
|
+
x = silu(x);
|
|
56
|
+
x = this.pointwiseConv2.forward(x);
|
|
57
|
+
return x;
|
|
58
|
+
}
|
|
59
|
+
loadWeights(weights, prefix) {
|
|
60
|
+
this.pointwiseConv1.loadWeights(weights, `${prefix}.pointwise_conv1`);
|
|
61
|
+
this.depthwiseConv.loadWeights(weights, `${prefix}.depthwise_conv`);
|
|
62
|
+
this.batchNorm.loadWeights(weights, `${prefix}.batch_norm`);
|
|
63
|
+
this.pointwiseConv2.loadWeights(weights, `${prefix}.pointwise_conv2`);
|
|
64
|
+
}
|
|
65
|
+
}
|
|
66
|
+
export class ConformerBlock extends Module {
|
|
67
|
+
normFF1;
|
|
68
|
+
ff1;
|
|
69
|
+
normSelfAtt;
|
|
70
|
+
selfAttn;
|
|
71
|
+
normConv;
|
|
72
|
+
conv;
|
|
73
|
+
normFF2;
|
|
74
|
+
ff2;
|
|
75
|
+
normOut;
|
|
76
|
+
args;
|
|
77
|
+
constructor(args) {
|
|
78
|
+
super();
|
|
79
|
+
this.args = args;
|
|
80
|
+
const useBias = args.useBias ?? true;
|
|
81
|
+
const ffDim = args.dModel * args.ffExpansionFactor;
|
|
82
|
+
this.normFF1 = new LayerNorm(args.dModel);
|
|
83
|
+
this.ff1 = new FeedForward(args.dModel, ffDim, useBias);
|
|
84
|
+
this.normSelfAtt = new LayerNorm(args.dModel);
|
|
85
|
+
this.selfAttn = this.buildAttention(args.selfAttentionModel, args.attContextSize ?? null);
|
|
86
|
+
this.normConv = new LayerNorm(args.dModel);
|
|
87
|
+
this.conv = new Convolution(args);
|
|
88
|
+
this.normFF2 = new LayerNorm(args.dModel);
|
|
89
|
+
this.ff2 = new FeedForward(args.dModel, ffDim, useBias);
|
|
90
|
+
this.normOut = new LayerNorm(args.dModel);
|
|
91
|
+
}
|
|
92
|
+
buildAttention(name, contextSize) {
|
|
93
|
+
const useBias = this.args.useBias ?? true;
|
|
94
|
+
if (name === 'rel_pos') {
|
|
95
|
+
return new RelPositionMultiHeadAttention(this.args.nHeads, this.args.dModel, useBias);
|
|
96
|
+
}
|
|
97
|
+
else if (name === 'rel_pos_local_attn') {
|
|
98
|
+
return new RelPositionMultiHeadLocalAttention(this.args.nHeads, this.args.dModel, useBias, contextSize ?? [256, 256]);
|
|
99
|
+
}
|
|
100
|
+
else {
|
|
101
|
+
return new MultiHeadAttention(this.args.nHeads, this.args.dModel, true);
|
|
102
|
+
}
|
|
103
|
+
}
|
|
104
|
+
setAttentionModel(name, contextSize = [256, 256]) {
|
|
105
|
+
const newAttn = this.buildAttention(name, contextSize);
|
|
106
|
+
// Copy weights from old attention if possible
|
|
107
|
+
// (In a real implementation we'd need to transfer parameters)
|
|
108
|
+
this.selfAttn = newAttn;
|
|
109
|
+
}
|
|
110
|
+
forward(x, posEmb, mask, cache) {
|
|
111
|
+
// FF1
|
|
112
|
+
x = x.add(this.ff1.forward(this.normFF1.forward(x)).mulScalar(0.5));
|
|
113
|
+
// Self-attention
|
|
114
|
+
const xNorm = this.normSelfAtt.forward(x);
|
|
115
|
+
x = x.add(this.selfAttn.forward(xNorm, xNorm, xNorm, posEmb, mask, cache));
|
|
116
|
+
// Convolution
|
|
117
|
+
x = x.add(this.conv.forward(this.normConv.forward(x), cache));
|
|
118
|
+
// FF2
|
|
119
|
+
x = x.add(this.ff2.forward(this.normFF2.forward(x)).mulScalar(0.5));
|
|
120
|
+
return this.normOut.forward(x);
|
|
121
|
+
}
|
|
122
|
+
loadWeights(weights, prefix) {
|
|
123
|
+
this.normFF1.loadWeights(weights, `${prefix}.norm_feed_forward1`);
|
|
124
|
+
this.ff1.loadWeights(weights, `${prefix}.feed_forward1`);
|
|
125
|
+
this.normSelfAtt.loadWeights(weights, `${prefix}.norm_self_att`);
|
|
126
|
+
this.selfAttn.loadWeights(weights, `${prefix}.self_attn`);
|
|
127
|
+
this.normConv.loadWeights(weights, `${prefix}.norm_conv`);
|
|
128
|
+
this.conv.loadWeights(weights, `${prefix}.conv`);
|
|
129
|
+
this.normFF2.loadWeights(weights, `${prefix}.norm_feed_forward2`);
|
|
130
|
+
this.ff2.loadWeights(weights, `${prefix}.feed_forward2`);
|
|
131
|
+
this.normOut.loadWeights(weights, `${prefix}.norm_out`);
|
|
132
|
+
}
|
|
133
|
+
}
|
|
134
|
+
// ---------------------------------------------------------------------------
|
|
135
|
+
// DW-striding subsampling (Conv2D-based)
|
|
136
|
+
// ---------------------------------------------------------------------------
|
|
137
|
+
class DwStridingSubsampling extends Module {
|
|
138
|
+
samplingNum;
|
|
139
|
+
stride = 2;
|
|
140
|
+
kernelSize = 3;
|
|
141
|
+
padding;
|
|
142
|
+
convLayers; // null = ReLU placeholder
|
|
143
|
+
out;
|
|
144
|
+
constructor(args) {
|
|
145
|
+
super();
|
|
146
|
+
this.padding = Math.floor((this.kernelSize - 1) / 2);
|
|
147
|
+
this.samplingNum = Math.round(Math.log2(args.subsamplingFactor));
|
|
148
|
+
const convChannels = args.subsamplingConvChannels;
|
|
149
|
+
// Compute final frequency dimension
|
|
150
|
+
let finalFreqDim = args.featIn;
|
|
151
|
+
for (let i = 0; i < this.samplingNum; i++) {
|
|
152
|
+
finalFreqDim =
|
|
153
|
+
Math.floor((finalFreqDim + 2 * this.padding - this.kernelSize) / this.stride) + 1;
|
|
154
|
+
}
|
|
155
|
+
// Build convolution layers
|
|
156
|
+
this.convLayers = [];
|
|
157
|
+
let inCh = 1;
|
|
158
|
+
// First layer: standard conv2d
|
|
159
|
+
this.convLayers.push(new Conv2d(inCh, convChannels, this.kernelSize, this.stride, this.padding, 1, true));
|
|
160
|
+
this.convLayers.push(null); // ReLU
|
|
161
|
+
inCh = convChannels;
|
|
162
|
+
for (let i = 1; i < this.samplingNum; i++) {
|
|
163
|
+
// Depthwise
|
|
164
|
+
this.convLayers.push(new Conv2d(inCh, inCh, this.kernelSize, this.stride, this.padding, inCh, true));
|
|
165
|
+
// Pointwise
|
|
166
|
+
this.convLayers.push(new Conv2d(inCh, convChannels, 1, 1, 0, 1, true));
|
|
167
|
+
this.convLayers.push(null); // ReLU
|
|
168
|
+
}
|
|
169
|
+
this.out = new Linear(convChannels * finalFreqDim, args.dModel);
|
|
170
|
+
}
|
|
171
|
+
forward(x, lengths) {
|
|
172
|
+
// x: [batch, seq, mel] → [batch, 1, seq, mel] then to MLX NHWC: [batch, seq, mel, 1]
|
|
173
|
+
const xShape = x.shape();
|
|
174
|
+
const batch = Number(xShape[0]);
|
|
175
|
+
// lengths update
|
|
176
|
+
let outLengths = lengths;
|
|
177
|
+
for (let i = 0; i < this.samplingNum; i++) {
|
|
178
|
+
const pad = this.padding;
|
|
179
|
+
const k = this.kernelSize;
|
|
180
|
+
const st = this.stride;
|
|
181
|
+
// floor((len + 2*pad - k) / stride) + 1
|
|
182
|
+
outLengths = outLengths
|
|
183
|
+
.addScalar(2 * pad - k)
|
|
184
|
+
.divScalar(st)
|
|
185
|
+
.floor()
|
|
186
|
+
.addScalar(1);
|
|
187
|
+
}
|
|
188
|
+
outLengths = outLengths.astype(3); // 3 = int32
|
|
189
|
+
// Reshape x: [batch, seq, mel] → [batch, seq, mel, 1] (NHWC with C=1)
|
|
190
|
+
let cur = x.expandDims(3); // [batch, seq, mel, 1]
|
|
191
|
+
for (const layer of this.convLayers) {
|
|
192
|
+
if (layer === null) {
|
|
193
|
+
cur = relu(cur);
|
|
194
|
+
}
|
|
195
|
+
else {
|
|
196
|
+
cur = layer.forward(cur);
|
|
197
|
+
}
|
|
198
|
+
}
|
|
199
|
+
// cur: [batch, outSeq, outMel, convChannels] (NHWC layout)
|
|
200
|
+
// Python flattens as [batch, outSeq, convChannels*outMel] (C-major, F-minor):
|
|
201
|
+
// it transposes NHWC -> NCHW then swapaxes(1,2) -> [B, T', C', F'] before
|
|
202
|
+
// reshaping. To match that memory ordering (which the Linear weights expect),
|
|
203
|
+
// transpose channel axis before freq axis here.
|
|
204
|
+
const cShape = cur.shape();
|
|
205
|
+
const outSeq = Number(cShape[1]);
|
|
206
|
+
const outMel = Number(cShape[2]);
|
|
207
|
+
const ch = Number(cShape[3]);
|
|
208
|
+
cur = cur
|
|
209
|
+
.transpose(new Int32Array([0, 1, 3, 2]))
|
|
210
|
+
.reshape(s(batch, outSeq, ch * outMel));
|
|
211
|
+
cur = this.out.forward(cur);
|
|
212
|
+
return [cur, outLengths];
|
|
213
|
+
}
|
|
214
|
+
loadWeights(weights, prefix) {
|
|
215
|
+
let convIdx = 0;
|
|
216
|
+
for (const layer of this.convLayers) {
|
|
217
|
+
if (layer !== null) {
|
|
218
|
+
layer.loadWeights(weights, `${prefix}.conv.${convIdx}`);
|
|
219
|
+
}
|
|
220
|
+
convIdx++;
|
|
221
|
+
}
|
|
222
|
+
this.out.loadWeights(weights, `${prefix}.out`);
|
|
223
|
+
}
|
|
224
|
+
}
|
|
225
|
+
// ---------------------------------------------------------------------------
|
|
226
|
+
// Conformer encoder
|
|
227
|
+
// ---------------------------------------------------------------------------
|
|
228
|
+
export class Conformer extends Module {
|
|
229
|
+
args;
|
|
230
|
+
posEnc;
|
|
231
|
+
preEncode;
|
|
232
|
+
layers;
|
|
233
|
+
constructor(args) {
|
|
234
|
+
super();
|
|
235
|
+
this.args = args;
|
|
236
|
+
const selfAttModel = args.selfAttentionModel;
|
|
237
|
+
const ctxSize = args.attContextSize ?? null;
|
|
238
|
+
if (selfAttModel === 'rel_pos') {
|
|
239
|
+
this.posEnc = new RelPositionalEncoding(args.dModel, args.posEmbMaxLen, args.xscaling ?? false);
|
|
240
|
+
}
|
|
241
|
+
else if (selfAttModel === 'rel_pos_local_attn') {
|
|
242
|
+
this.posEnc = new LocalRelPositionalEncoding(args.dModel, args.posEmbMaxLen, args.xscaling ?? false, ctxSize ?? [256, 256]);
|
|
243
|
+
}
|
|
244
|
+
else {
|
|
245
|
+
this.posEnc = null;
|
|
246
|
+
}
|
|
247
|
+
if (args.subsamplingFactor > 1) {
|
|
248
|
+
if (args.subsampling === 'dw_striding' && !(args.causalDownsampling ?? false)) {
|
|
249
|
+
this.preEncode = new DwStridingSubsampling(args);
|
|
250
|
+
}
|
|
251
|
+
else {
|
|
252
|
+
throw new Error('Only dw_striding non-causal subsampling is supported');
|
|
253
|
+
}
|
|
254
|
+
}
|
|
255
|
+
else {
|
|
256
|
+
this.preEncode = new Linear(args.featIn, args.dModel);
|
|
257
|
+
}
|
|
258
|
+
this.layers = Array.from({ length: args.nLayers }, () => new ConformerBlock(args));
|
|
259
|
+
}
|
|
260
|
+
setAttentionModel(name, contextSize = [256, 256]) {
|
|
261
|
+
if (name === 'rel_pos') {
|
|
262
|
+
this.posEnc = new RelPositionalEncoding(this.args.dModel, this.args.posEmbMaxLen, this.args.xscaling ?? false);
|
|
263
|
+
}
|
|
264
|
+
else if (name === 'rel_pos_local_attn') {
|
|
265
|
+
this.posEnc = new LocalRelPositionalEncoding(this.args.dModel, this.args.posEmbMaxLen, this.args.xscaling ?? false, contextSize);
|
|
266
|
+
}
|
|
267
|
+
else {
|
|
268
|
+
this.posEnc = null;
|
|
269
|
+
}
|
|
270
|
+
for (const layer of this.layers) {
|
|
271
|
+
layer.setAttentionModel(name, contextSize);
|
|
272
|
+
}
|
|
273
|
+
}
|
|
274
|
+
forward(x, // [batch, seq, mel]
|
|
275
|
+
lengths, cache) {
|
|
276
|
+
const xShape = x.shape();
|
|
277
|
+
const batch = Number(xShape[0]);
|
|
278
|
+
const seq = Number(xShape[1]);
|
|
279
|
+
if (lengths === null) {
|
|
280
|
+
const lenData = new Int32Array(batch).fill(seq);
|
|
281
|
+
lengths = MxArray.fromInt32(lenData, BigInt64Array.from([BigInt(batch)]));
|
|
282
|
+
}
|
|
283
|
+
let outLengths;
|
|
284
|
+
if (this.preEncode instanceof DwStridingSubsampling) {
|
|
285
|
+
[x, outLengths] = this.preEncode.forward(x, lengths);
|
|
286
|
+
}
|
|
287
|
+
else {
|
|
288
|
+
x = this.preEncode.forward(x);
|
|
289
|
+
outLengths = lengths;
|
|
290
|
+
}
|
|
291
|
+
const effectiveCache = cache ?? new Array(this.layers.length).fill(null);
|
|
292
|
+
let posEmb = null;
|
|
293
|
+
if (this.posEnc !== null) {
|
|
294
|
+
const offset = effectiveCache[0]?.offset ?? 0;
|
|
295
|
+
[x, posEmb] = this.posEnc.forward(x, offset);
|
|
296
|
+
}
|
|
297
|
+
for (let i = 0; i < this.layers.length; i++) {
|
|
298
|
+
x = this.layers[i].forward(x, posEmb, null, effectiveCache[i]);
|
|
299
|
+
}
|
|
300
|
+
return [x, outLengths];
|
|
301
|
+
}
|
|
302
|
+
loadWeights(weights, prefix) {
|
|
303
|
+
if (this.preEncode instanceof DwStridingSubsampling) {
|
|
304
|
+
this.preEncode.loadWeights(weights, `${prefix}.pre_encode`);
|
|
305
|
+
}
|
|
306
|
+
else {
|
|
307
|
+
this.preEncode.loadWeights(weights, `${prefix}.pre_encode`);
|
|
308
|
+
}
|
|
309
|
+
for (let i = 0; i < this.layers.length; i++) {
|
|
310
|
+
this.layers[i].loadWeights(weights, `${prefix}.layers.${i}`);
|
|
311
|
+
}
|
|
312
|
+
// pos_enc has no learnable weights (pe is computed; pos_bias_{u,v} live in attention)
|
|
313
|
+
}
|
|
314
|
+
}
|
|
315
|
+
//# sourceMappingURL=conformer.js.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"conformer.js","sourceRoot":"","sources":["../../src/mlx/conformer.ts"],"names":[],"mappings":"AAAA,OAAO,EAAE,OAAO,EAAE,MAAM,gBAAgB,CAAC;AACzC,OAAO,EACL,MAAM,EAEN,MAAM,EACN,SAAS,EACT,SAAS,EACT,MAAM,EACN,MAAM,EACN,IAAI,EACJ,IAAI,EACJ,GAAG,GACJ,MAAM,SAAS,CAAC;AACjB,OAAO,EACL,kBAAkB,EAClB,6BAA6B,EAC7B,kCAAkC,EAClC,qBAAqB,EACrB,0BAA0B,GAC3B,MAAM,gBAAgB,CAAC;AAGxB,SAAS,CAAC,CAAC,GAAG,IAAc;IAC1B,OAAO,aAAa,CAAC,IAAI,CAAC,IAAI,CAAC,GAAG,CAAC,MAAM,CAAC,CAAC,CAAC;AAC9C,CAAC;AAyBD,8EAA8E;AAC9E,qBAAqB;AACrB,8EAA8E;AAE9E,MAAM,WAAY,SAAQ,MAAM;IAC9B,OAAO,CAAS;IAChB,OAAO,CAAS;IAEhB,YAAY,MAAc,EAAE,GAAW,EAAE,OAAgB;QACvD,KAAK,EAAE,CAAC;QACR,IAAI,CAAC,OAAO,GAAG,IAAI,MAAM,CAAC,MAAM,EAAE,GAAG,EAAE,OAAO,CAAC,CAAC;QAChD,IAAI,CAAC,OAAO,GAAG,IAAI,MAAM,CAAC,GAAG,EAAE,MAAM,EAAE,OAAO,CAAC,CAAC;IAClD,CAAC;IAED,OAAO,CAAC,CAAU;QAChB,OAAO,IAAI,CAAC,OAAO,CAAC,OAAO,CAAC,IAAI,CAAC,IAAI,CAAC,OAAO,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC,CAAC,CAAC;IAC7D,CAAC;IAED,WAAW,CAAC,OAAkB,EAAE,MAAc;QAC5C,IAAI,CAAC,OAAO,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,UAAU,CAAC,CAAC;QACvD,IAAI,CAAC,OAAO,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,UAAU,CAAC,CAAC;IACzD,CAAC;CACF;AAED,8EAA8E;AAC9E,oBAAoB;AACpB,8EAA8E;AAE9E,MAAM,WAAY,SAAQ,MAAM;IACrB,OAAO,CAAS;IACzB,cAAc,CAAS;IACvB,aAAa,CAAS;IACtB,SAAS,CAAY;IACrB,cAAc,CAAS;IAEvB,YAAY,IAAmB;QAC7B,KAAK,EAAE,CAAC;QACR,MAAM,OAAO,GAAG,IAAI,CAAC,OAAO,IAAI,IAAI,CAAC;QACrC,IAAI,CAAC,OAAO,GAAG,IAAI,CAAC,KAAK,CAAC,CAAC,IAAI,CAAC,cAAc,GAAG,CAAC,CAAC,GAAG,CAAC,CAAC,CAAC;QAEzD,IAAI,CAAC,cAAc,GAAG,IAAI,MAAM,CAAC,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,MAAM,GAAG,CAAC,EAAE,CAAC,EAAE,CAAC,EAAE,CAAC,EAAE,CAAC,EAAE,OAAO,CAAC,CAAC;QACpF,IAAI,CAAC,aAAa,GAAG,IAAI,MAAM,CAAC,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,cAAc,EAAE,CAAC,EAAE,CAAC,EAAE,IAAI,CAAC,MAAM,EAAE,OAAO,CAAC,CAAC;QAC3G,IAAI,CAAC,SAAS,GAAG,IAAI,SAAS,CAAC,IAAI,CAAC,MAAM,CAAC,CAAC;QAC5C,IAAI,CAAC,cAAc,GAAG,IAAI,MAAM,CAAC,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,MAAM,EAAE,CAAC,EAAE,CAAC,EAAE,CAAC,EAAE,CAAC,EAAE,OAAO,CAAC,CAAC;IAClF,CAAC;IAED,OAAO,CAAC,CAAU,EAAE,KAA4B;QAC9C,CAAC,GAAG,IAAI,CAAC,cAAc,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC;QACnC,CAAC,GAAG,GAAG,CAAC,CAAC,EAAE,CAAC,CAAC,CAAC,CAAC,2CAA2C;QAE1D,IAAI,KAAK,KAAK,IAAI,EAAE,CAAC;YACnB,CAAC,GAAG,KAAK,CAAC,kBAAkB,CAAC,CAAC,EAAE,IAAI,CAAC,OAAO,CAAC,CAAC;QAChD,CAAC;aAAM,CAAC;YACN,CAAC,GAAG,CAAC,CAAC,GAAG,CAAC,IAAI,UAAU,CAAC,CAAC,CAAC,EAAE,CAAC,EAAE,IAAI,CAAC,OAAO,EAAE,IAAI,CAAC,OAAO,EAAE,CAAC,EAAE,CAAC,CAAC,CAAC,EAAE,GAAG,CAAC,CAAC;QAC3E,CAAC;QAED,CAAC,GAAG,IAAI,CAAC,aAAa,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC;QAClC,CAAC,GAAG,IAAI,CAAC,SAAS,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC;QAC9B,CAAC,GAAG,IAAI,CAAC,CAAC,CAAC,CAAC;QACZ,CAAC,GAAG,IAAI,CAAC,cAAc,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC;QACnC,OAAO,CAAC,CAAC;IACX,CAAC;IAED,WAAW,CAAC,OAAkB,EAAE,MAAc;QAC5C,IAAI,CAAC,cAAc,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,kBAAkB,CAAC,CAAC;QACtE,IAAI,CAAC,aAAa,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,iBAAiB,CAAC,CAAC;QACpE,IAAI,CAAC,SAAS,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,aAAa,CAAC,CAAC;QAC5D,IAAI,CAAC,cAAc,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,kBAAkB,CAAC,CAAC;IACxE,CAAC;CACF;AAQD,MAAM,OAAO,cAAe,SAAQ,MAAM;IACxC,OAAO,CAAY;IACnB,GAAG,CAAc;IAEjB,WAAW,CAAY;IACvB,QAAQ,CAA0F;IAElG,QAAQ,CAAY;IACpB,IAAI,CAAc;IAElB,OAAO,CAAY;IACnB,GAAG,CAAc;IAEjB,OAAO,CAAY;IAEF,IAAI,CAAgB;IAErC,YAAY,IAAmB;QAC7B,KAAK,EAAE,CAAC;QACR,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;QACjB,MAAM,OAAO,GAAG,IAAI,CAAC,OAAO,IAAI,IAAI,CAAC;QACrC,MAAM,KAAK,GAAG,IAAI,CAAC,MAAM,GAAG,IAAI,CAAC,iBAAiB,CAAC;QAEnD,IAAI,CAAC,OAAO,GAAG,IAAI,SAAS,CAAC,IAAI,CAAC,MAAM,CAAC,CAAC;QAC1C,IAAI,CAAC,GAAG,GAAG,IAAI,WAAW,CAAC,IAAI,CAAC,MAAM,EAAE,KAAK,EAAE,OAAO,CAAC,CAAC;QAExD,IAAI,CAAC,WAAW,GAAG,IAAI,SAAS,CAAC,IAAI,CAAC,MAAM,CAAC,CAAC;QAC9C,IAAI,CAAC,QAAQ,GAAG,IAAI,CAAC,cAAc,CAAC,IAAI,CAAC,kBAAkB,EAAE,IAAI,CAAC,cAAc,IAAI,IAAI,CAAC,CAAC;QAE1F,IAAI,CAAC,QAAQ,GAAG,IAAI,SAAS,CAAC,IAAI,CAAC,MAAM,CAAC,CAAC;QAC3C,IAAI,CAAC,IAAI,GAAG,IAAI,WAAW,CAAC,IAAI,CAAC,CAAC;QAElC,IAAI,CAAC,OAAO,GAAG,IAAI,SAAS,CAAC,IAAI,CAAC,MAAM,CAAC,CAAC;QAC1C,IAAI,CAAC,GAAG,GAAG,IAAI,WAAW,CAAC,IAAI,CAAC,MAAM,EAAE,KAAK,EAAE,OAAO,CAAC,CAAC;QAExD,IAAI,CAAC,OAAO,GAAG,IAAI,SAAS,CAAC,IAAI,CAAC,MAAM,CAAC,CAAC;IAC5C,CAAC;IAEO,cAAc,CACpB,IAAY,EACZ,WAAoC;QAEpC,MAAM,OAAO,GAAG,IAAI,CAAC,IAAI,CAAC,OAAO,IAAI,IAAI,CAAC;QAC1C,IAAI,IAAI,KAAK,SAAS,EAAE,CAAC;YACvB,OAAO,IAAI,6BAA6B,CAAC,IAAI,CAAC,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,IAAI,CAAC,MAAM,EAAE,OAAO,CAAC,CAAC;QACxF,CAAC;aAAM,IAAI,IAAI,KAAK,oBAAoB,EAAE,CAAC;YACzC,OAAO,IAAI,kCAAkC,CAC3C,IAAI,CAAC,IAAI,CAAC,MAAM,EAChB,IAAI,CAAC,IAAI,CAAC,MAAM,EAChB,OAAO,EACP,WAAW,IAAI,CAAC,GAAG,EAAE,GAAG,CAAC,CAC1B,CAAC;QACJ,CAAC;aAAM,CAAC;YACN,OAAO,IAAI,kBAAkB,CAAC,IAAI,CAAC,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,CAAC;QAC1E,CAAC;IACH,CAAC;IAED,iBAAiB,CAAC,IAAoB,EAAE,cAAgC,CAAC,GAAG,EAAE,GAAG,CAAC;QAChF,MAAM,OAAO,GAAG,IAAI,CAAC,cAAc,CAAC,IAAI,EAAE,WAAW,CAAC,CAAC;QACvD,8CAA8C;QAC9C,8DAA8D;QAC9D,IAAI,CAAC,QAAQ,GAAG,OAAO,CAAC;IAC1B,CAAC;IAED,OAAO,CACL,CAAU,EACV,MAAsB,EACtB,IAAoB,EACpB,KAA4B;QAE5B,MAAM;QACN,CAAC,GAAG,CAAC,CAAC,GAAG,CAAC,IAAI,CAAC,GAAG,CAAC,OAAO,CAAC,IAAI,CAAC,OAAO,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC,CAAC,SAAS,CAAC,GAAG,CAAC,CAAC,CAAC;QAEpE,iBAAiB;QACjB,MAAM,KAAK,GAAG,IAAI,CAAC,WAAW,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC;QAC1C,CAAC,GAAG,CAAC,CAAC,GAAG,CAAC,IAAI,CAAC,QAAQ,CAAC,OAAO,CAAC,KAAK,EAAE,KAAK,EAAE,KAAK,EAAE,MAAM,EAAE,IAAI,EAAE,KAAK,CAAC,CAAC,CAAC;QAE3E,cAAc;QACd,CAAC,GAAG,CAAC,CAAC,GAAG,CAAC,IAAI,CAAC,IAAI,CAAC,OAAO,CAAC,IAAI,CAAC,QAAQ,CAAC,OAAO,CAAC,CAAC,CAAC,EAAE,KAAK,CAAC,CAAC,CAAC;QAE9D,MAAM;QACN,CAAC,GAAG,CAAC,CAAC,GAAG,CAAC,IAAI,CAAC,GAAG,CAAC,OAAO,CAAC,IAAI,CAAC,OAAO,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC,CAAC,SAAS,CAAC,GAAG,CAAC,CAAC,CAAC;QAEpE,OAAO,IAAI,CAAC,OAAO,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC;IACjC,CAAC;IAED,WAAW,CAAC,OAAkB,EAAE,MAAc;QAC5C,IAAI,CAAC,OAAO,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,qBAAqB,CAAC,CAAC;QAClE,IAAI,CAAC,GAAG,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,gBAAgB,CAAC,CAAC;QACzD,IAAI,CAAC,WAAW,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,gBAAgB,CAAC,CAAC;QACjE,IAAI,CAAC,QAAQ,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,YAAY,CAAC,CAAC;QAC1D,IAAI,CAAC,QAAQ,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,YAAY,CAAC,CAAC;QAC1D,IAAI,CAAC,IAAI,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,OAAO,CAAC,CAAC;QACjD,IAAI,CAAC,OAAO,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,qBAAqB,CAAC,CAAC;QAClE,IAAI,CAAC,GAAG,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,gBAAgB,CAAC,CAAC;QACzD,IAAI,CAAC,OAAO,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,WAAW,CAAC,CAAC;IAC1D,CAAC;CACF;AAED,8EAA8E;AAC9E,yCAAyC;AACzC,8EAA8E;AAE9E,MAAM,qBAAsB,SAAQ,MAAM;IACvB,WAAW,CAAS;IACpB,MAAM,GAAG,CAAC,CAAC;IACX,UAAU,GAAG,CAAC,CAAC;IACf,OAAO,CAAS;IAEjC,UAAU,CAAuB,CAAC,0BAA0B;IAC5D,GAAG,CAAS;IAEZ,YAAY,IAAmB;QAC7B,KAAK,EAAE,CAAC;QACR,IAAI,CAAC,OAAO,GAAG,IAAI,CAAC,KAAK,CAAC,CAAC,IAAI,CAAC,UAAU,GAAG,CAAC,CAAC,GAAG,CAAC,CAAC,CAAC;QACrD,IAAI,CAAC,WAAW,GAAG,IAAI,CAAC,KAAK,CAAC,IAAI,CAAC,IAAI,CAAC,IAAI,CAAC,iBAAiB,CAAC,CAAC,CAAC;QAEjE,MAAM,YAAY,GAAG,IAAI,CAAC,uBAAuB,CAAC;QAElD,oCAAoC;QACpC,IAAI,YAAY,GAAG,IAAI,CAAC,MAAM,CAAC;QAC/B,KAAK,IAAI,CAAC,GAAG,CAAC,EAAE,CAAC,GAAG,IAAI,CAAC,WAAW,EAAE,CAAC,EAAE,EAAE,CAAC;YAC1C,YAAY;gBACV,IAAI,CAAC,KAAK,CAAC,CAAC,YAAY,GAAG,CAAC,GAAG,IAAI,CAAC,OAAO,GAAG,IAAI,CAAC,UAAU,CAAC,GAAG,IAAI,CAAC,MAAM,CAAC,GAAG,CAAC,CAAC;QACtF,CAAC;QAED,2BAA2B;QAC3B,IAAI,CAAC,UAAU,GAAG,EAAE,CAAC;QACrB,IAAI,IAAI,GAAG,CAAC,CAAC;QAEb,+BAA+B;QAC/B,IAAI,CAAC,UAAU,CAAC,IAAI,CAAC,IAAI,MAAM,CAAC,IAAI,EAAE,YAAY,EAAE,IAAI,CAAC,UAAU,EAAE,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,OAAO,EAAE,CAAC,EAAE,IAAI,CAAC,CAAC,CAAC;QAC1G,IAAI,CAAC,UAAU,CAAC,IAAI,CAAC,IAAI,CAAC,CAAC,CAAC,OAAO;QAEnC,IAAI,GAAG,YAAY,CAAC;QACpB,KAAK,IAAI,CAAC,GAAG,CAAC,EAAE,CAAC,GAAG,IAAI,CAAC,WAAW,EAAE,CAAC,EAAE,EAAE,CAAC;YAC1C,YAAY;YACZ,IAAI,CAAC,UAAU,CAAC,IAAI,CAAC,IAAI,MAAM,CAAC,IAAI,EAAE,IAAI,EAAE,IAAI,CAAC,UAAU,EAAE,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,OAAO,EAAE,IAAI,EAAE,IAAI,CAAC,CAAC,CAAC;YACrG,YAAY;YACZ,IAAI,CAAC,UAAU,CAAC,IAAI,CAAC,IAAI,MAAM,CAAC,IAAI,EAAE,YAAY,EAAE,CAAC,EAAE,CAAC,EAAE,CAAC,EAAE,CAAC,EAAE,IAAI,CAAC,CAAC,CAAC;YACvE,IAAI,CAAC,UAAU,CAAC,IAAI,CAAC,IAAI,CAAC,CAAC,CAAC,OAAO;QACrC,CAAC;QAED,IAAI,CAAC,GAAG,GAAG,IAAI,MAAM,CAAC,YAAY,GAAG,YAAY,EAAE,IAAI,CAAC,MAAM,CAAC,CAAC;IAClE,CAAC;IAED,OAAO,CAAC,CAAU,EAAE,OAAgB;QAClC,qFAAqF;QACrF,MAAM,MAAM,GAAG,CAAC,CAAC,KAAK,EAAE,CAAC;QACzB,MAAM,KAAK,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;QAEhC,iBAAiB;QACjB,IAAI,UAAU,GAAG,OAAO,CAAC;QACzB,KAAK,IAAI,CAAC,GAAG,CAAC,EAAE,CAAC,GAAG,IAAI,CAAC,WAAW,EAAE,CAAC,EAAE,EAAE,CAAC;YAC1C,MAAM,GAAG,GAAG,IAAI,CAAC,OAAO,CAAC;YACzB,MAAM,CAAC,GAAG,IAAI,CAAC,UAAU,CAAC;YAC1B,MAAM,EAAE,GAAG,IAAI,CAAC,MAAM,CAAC;YACvB,wCAAwC;YACxC,UAAU,GAAG,UAAU;iBACpB,SAAS,CAAC,CAAC,GAAG,GAAG,GAAG,CAAC,CAAC;iBACtB,SAAS,CAAC,EAAE,CAAC;iBACb,KAAK,EAAE;iBACP,SAAS,CAAC,CAAC,CAAC,CAAC;QAClB,CAAC;QACD,UAAU,GAAG,UAAU,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC,YAAY;QAE/C,sEAAsE;QACtE,IAAI,GAAG,GAAG,CAAC,CAAC,UAAU,CAAC,CAAC,CAAC,CAAC,CAAC,uBAAuB;QAElD,KAAK,MAAM,KAAK,IAAI,IAAI,CAAC,UAAU,EAAE,CAAC;YACpC,IAAI,KAAK,KAAK,IAAI,EAAE,CAAC;gBACnB,GAAG,GAAG,IAAI,CAAC,GAAG,CAAC,CAAC;YAClB,CAAC;iBAAM,CAAC;gBACN,GAAG,GAAG,KAAK,CAAC,OAAO,CAAC,GAAG,CAAC,CAAC;YAC3B,CAAC;QACH,CAAC;QAED,2DAA2D;QAC3D,8EAA8E;QAC9E,0EAA0E;QAC1E,8EAA8E;QAC9E,gDAAgD;QAChD,MAAM,MAAM,GAAG,GAAG,CAAC,KAAK,EAAE,CAAC;QAC3B,MAAM,MAAM,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;QACjC,MAAM,MAAM,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;QACjC,MAAM,EAAE,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;QAE7B,GAAG,GAAG,GAAG;aACN,SAAS,CAAC,IAAI,UAAU,CAAC,CAAC,CAAC,EAAE,CAAC,EAAE,CAAC,EAAE,CAAC,CAAC,CAAC,CAAC;aACvC,OAAO,CAAC,CAAC,CAAC,KAAK,EAAE,MAAM,EAAE,EAAE,GAAG,MAAM,CAAC,CAAC,CAAC;QAC1C,GAAG,GAAG,IAAI,CAAC,GAAG,CAAC,OAAO,CAAC,GAAG,CAAC,CAAC;QAE5B,OAAO,CAAC,GAAG,EAAE,UAAU,CAAC,CAAC;IAC3B,CAAC;IAED,WAAW,CAAC,OAAkB,EAAE,MAAc;QAC5C,IAAI,OAAO,GAAG,CAAC,CAAC;QAChB,KAAK,MAAM,KAAK,IAAI,IAAI,CAAC,UAAU,EAAE,CAAC;YACpC,IAAI,KAAK,KAAK,IAAI,EAAE,CAAC;gBACnB,KAAK,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,SAAS,OAAO,EAAE,CAAC,CAAC;YAC1D,CAAC;YACD,OAAO,EAAE,CAAC;QACZ,CAAC;QACD,IAAI,CAAC,GAAG,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,MAAM,CAAC,CAAC;IACjD,CAAC;CACF;AAED,8EAA8E;AAC9E,oBAAoB;AACpB,8EAA8E;AAE9E,MAAM,OAAO,SAAU,SAAQ,MAAM;IAC1B,IAAI,CAAgB;IAC7B,MAAM,CAA4D;IAClE,SAAS,CAAiC;IAC1C,MAAM,CAAmB;IAEzB,YAAY,IAAmB;QAC7B,KAAK,EAAE,CAAC;QACR,IAAI,CAAC,IAAI,GAAG,IAAI,CAAC;QAEjB,MAAM,YAAY,GAAG,IAAI,CAAC,kBAAkB,CAAC;QAC7C,MAAM,OAAO,GAAG,IAAI,CAAC,cAAc,IAAI,IAAI,CAAC;QAE5C,IAAI,YAAY,KAAK,SAAS,EAAE,CAAC;YAC/B,IAAI,CAAC,MAAM,GAAG,IAAI,qBAAqB,CAAC,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,YAAY,EAAE,IAAI,CAAC,QAAQ,IAAI,KAAK,CAAC,CAAC;QAClG,CAAC;aAAM,IAAI,YAAY,KAAK,oBAAoB,EAAE,CAAC;YACjD,IAAI,CAAC,MAAM,GAAG,IAAI,0BAA0B,CAC1C,IAAI,CAAC,MAAM,EACX,IAAI,CAAC,YAAY,EACjB,IAAI,CAAC,QAAQ,IAAI,KAAK,EACtB,OAAO,IAAI,CAAC,GAAG,EAAE,GAAG,CAAC,CACtB,CAAC;QACJ,CAAC;aAAM,CAAC;YACN,IAAI,CAAC,MAAM,GAAG,IAAI,CAAC;QACrB,CAAC;QAED,IAAI,IAAI,CAAC,iBAAiB,GAAG,CAAC,EAAE,CAAC;YAC/B,IAAI,IAAI,CAAC,WAAW,KAAK,aAAa,IAAI,CAAC,CAAC,IAAI,CAAC,kBAAkB,IAAI,KAAK,CAAC,EAAE,CAAC;gBAC9E,IAAI,CAAC,SAAS,GAAG,IAAI,qBAAqB,CAAC,IAAI,CAAC,CAAC;YACnD,CAAC;iBAAM,CAAC;gBACN,MAAM,IAAI,KAAK,CAAC,sDAAsD,CAAC,CAAC;YAC1E,CAAC;QACH,CAAC;aAAM,CAAC;YACN,IAAI,CAAC,SAAS,GAAG,IAAI,MAAM,CAAC,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,MAAM,CAAC,CAAC;QACxD,CAAC;QAED,IAAI,CAAC,MAAM,GAAG,KAAK,CAAC,IAAI,CAAC,EAAE,MAAM,EAAE,IAAI,CAAC,OAAO,EAAE,EAAE,GAAG,EAAE,CAAC,IAAI,cAAc,CAAC,IAAI,CAAC,CAAC,CAAC;IACrF,CAAC;IAED,iBAAiB,CACf,IAAoB,EACpB,cAAgC,CAAC,GAAG,EAAE,GAAG,CAAC;QAE1C,IAAI,IAAI,KAAK,SAAS,EAAE,CAAC;YACvB,IAAI,CAAC,MAAM,GAAG,IAAI,qBAAqB,CAAC,IAAI,CAAC,IAAI,CAAC,MAAM,EAAE,IAAI,CAAC,IAAI,CAAC,YAAY,EAAE,IAAI,CAAC,IAAI,CAAC,QAAQ,IAAI,KAAK,CAAC,CAAC;QACjH,CAAC;aAAM,IAAI,IAAI,KAAK,oBAAoB,EAAE,CAAC;YACzC,IAAI,CAAC,MAAM,GAAG,IAAI,0BAA0B,CAC1C,IAAI,CAAC,IAAI,CAAC,MAAM,EAChB,IAAI,CAAC,IAAI,CAAC,YAAY,EACtB,IAAI,CAAC,IAAI,CAAC,QAAQ,IAAI,KAAK,EAC3B,WAAW,CACZ,CAAC;QACJ,CAAC;aAAM,CAAC;YACN,IAAI,CAAC,MAAM,GAAG,IAAI,CAAC;QACrB,CAAC;QAED,KAAK,MAAM,KAAK,IAAI,IAAI,CAAC,MAAM,EAAE,CAAC;YAChC,KAAK,CAAC,iBAAiB,CAAC,IAAI,EAAE,WAAW,CAAC,CAAC;QAC7C,CAAC;IACH,CAAC;IAED,OAAO,CACL,CAAU,EAAE,oBAAoB;IAChC,OAAuB,EACvB,KAA0C;QAE1C,MAAM,MAAM,GAAG,CAAC,CAAC,KAAK,EAAE,CAAC;QACzB,MAAM,KAAK,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;QAChC,MAAM,GAAG,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;QAE9B,IAAI,OAAO,KAAK,IAAI,EAAE,CAAC;YACrB,MAAM,OAAO,GAAG,IAAI,UAAU,CAAC,KAAK,CAAC,CAAC,IAAI,CAAC,GAAG,CAAC,CAAC;YAChD,OAAO,GAAG,OAAO,CAAC,SAAS,CAAC,OAAO,EAAE,aAAa,CAAC,IAAI,CAAC,CAAC,MAAM,CAAC,KAAK,CAAC,CAAC,CAAC,CAAC,CAAC;QAC5E,CAAC;QAED,IAAI,UAAmB,CAAC;QAExB,IAAI,IAAI,CAAC,SAAS,YAAY,qBAAqB,EAAE,CAAC;YACpD,CAAC,CAAC,EAAE,UAAU,CAAC,GAAG,IAAI,CAAC,SAAS,CAAC,OAAO,CAAC,CAAC,EAAE,OAAO,CAAC,CAAC;QACvD,CAAC;aAAM,CAAC;YACN,CAAC,GAAG,IAAI,CAAC,SAAS,CAAC,OAAO,CAAC,CAAC,CAAC,CAAC;YAC9B,UAAU,GAAG,OAAO,CAAC;QACvB,CAAC;QAED,MAAM,cAAc,GAAG,KAAK,IAAI,IAAI,KAAK,CAAC,IAAI,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,IAAI,CAAC,IAAI,CAAC,CAAC;QAEzE,IAAI,MAAM,GAAmB,IAAI,CAAC;QAClC,IAAI,IAAI,CAAC,MAAM,KAAK,IAAI,EAAE,CAAC;YACzB,MAAM,MAAM,GAAG,cAAc,CAAC,CAAC,CAAC,EAAE,MAAM,IAAI,CAAC,CAAC;YAC9C,CAAC,CAAC,EAAE,MAAM,CAAC,GAAG,IAAI,CAAC,MAAM,CAAC,OAAO,CAAC,CAAC,EAAE,MAAM,CAAC,CAAC;QAC/C,CAAC;QAED,KAAK,IAAI,CAAC,GAAG,CAAC,EAAE,CAAC,GAAG,IAAI,CAAC,MAAM,CAAC,MAAM,EAAE,CAAC,EAAE,EAAE,CAAC;YAC5C,CAAC,GAAG,IAAI,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,OAAO,CAAC,CAAC,EAAE,MAAM,EAAE,IAAI,EAAE,cAAc,CAAC,CAAC,CAAC,CAAC,CAAC;QACjE,CAAC;QAED,OAAO,CAAC,CAAC,EAAE,UAAU,CAAC,CAAC;IACzB,CAAC;IAED,WAAW,CAAC,OAAkB,EAAE,MAAc;QAC5C,IAAI,IAAI,CAAC,SAAS,YAAY,qBAAqB,EAAE,CAAC;YACpD,IAAI,CAAC,SAAS,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,aAAa,CAAC,CAAC;QAC9D,CAAC;aAAM,CAAC;YACL,IAAI,CAAC,SAAoB,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,aAAa,CAAC,CAAC;QAC1E,CAAC;QAED,KAAK,IAAI,CAAC,GAAG,CAAC,EAAE,CAAC,GAAG,IAAI,CAAC,MAAM,CAAC,MAAM,EAAE,CAAC,EAAE,EAAE,CAAC;YAC5C,IAAI,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,WAAW,CAAC,OAAO,EAAE,GAAG,MAAM,WAAW,CAAC,EAAE,CAAC,CAAC;QAC/D,CAAC;QAED,sFAAsF;IACxF,CAAC;CACF"}
|
|
@@ -0,0 +1,6 @@
|
|
|
1
|
+
export { fromLocal, fromPretrained } from './load.js';
|
|
2
|
+
export type { MlxModelOptions, FromPretrainedOptions } from './load.js';
|
|
3
|
+
export { MlxBackend } from './backend.js';
|
|
4
|
+
export { ParakeetModel, StreamingParakeet, consumePcmStream, decodeTDTGreedy, decodeRNNTGreedy, getLogMel, loadAudioRaw, makePreprocessArgs, computeMelFilterbanks, computeMelFilterbanksInterpolated, tokensToSentences, sentencesToResult, decode, } from '../index.js';
|
|
5
|
+
export type { ParakeetBackend, EncoderOutput, TranscribeOptions, StreamOptions, DecoderState, PreprocessArgs, LogMel, AlignedToken, AlignedSentence, AlignedResult, SentenceConfig, } from '../index.js';
|
|
6
|
+
//# sourceMappingURL=index.d.ts.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"index.d.ts","sourceRoot":"","sources":["../../src/mlx/index.ts"],"names":[],"mappings":"AAKA,OAAO,EAAE,SAAS,EAAE,cAAc,EAAE,MAAM,WAAW,CAAC;AACtD,YAAY,EAAE,eAAe,EAAE,qBAAqB,EAAE,MAAM,WAAW,CAAC;AACxE,OAAO,EAAE,UAAU,EAAE,MAAM,cAAc,CAAC;AAG1C,OAAO,EACL,aAAa,EACb,iBAAiB,EACjB,gBAAgB,EAChB,eAAe,EACf,gBAAgB,EAChB,SAAS,EACT,YAAY,EACZ,kBAAkB,EAClB,qBAAqB,EACrB,iCAAiC,EACjC,iBAAiB,EACjB,iBAAiB,EACjB,MAAM,GACP,MAAM,aAAa,CAAC;AACrB,YAAY,EACV,eAAe,EACf,aAAa,EACb,iBAAiB,EACjB,aAAa,EACb,YAAY,EACZ,cAAc,EACd,MAAM,EACN,YAAY,EACZ,eAAe,EACf,aAAa,EACb,cAAc,GACf,MAAM,aAAa,CAAC"}
|
|
@@ -0,0 +1,9 @@
|
|
|
1
|
+
// parakeet.ts/mlx — MLX backend (Apple Silicon)
|
|
2
|
+
//
|
|
3
|
+
// Loads a safetensors checkpoint and returns the shared, backend-agnostic
|
|
4
|
+
// `ParakeetModel` — the same class the ONNX backend returns.
|
|
5
|
+
export { fromLocal, fromPretrained } from './load.js';
|
|
6
|
+
export { MlxBackend } from './backend.js';
|
|
7
|
+
// Shared API, re-exported for convenience
|
|
8
|
+
export { ParakeetModel, StreamingParakeet, consumePcmStream, decodeTDTGreedy, decodeRNNTGreedy, getLogMel, loadAudioRaw, makePreprocessArgs, computeMelFilterbanks, computeMelFilterbanksInterpolated, tokensToSentences, sentencesToResult, decode, } from '../index.js';
|
|
9
|
+
//# sourceMappingURL=index.js.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"index.js","sourceRoot":"","sources":["../../src/mlx/index.ts"],"names":[],"mappings":"AAAA,gDAAgD;AAChD,EAAE;AACF,0EAA0E;AAC1E,6DAA6D;AAE7D,OAAO,EAAE,SAAS,EAAE,cAAc,EAAE,MAAM,WAAW,CAAC;AAEtD,OAAO,EAAE,UAAU,EAAE,MAAM,cAAc,CAAC;AAE1C,0CAA0C;AAC1C,OAAO,EACL,aAAa,EACb,iBAAiB,EACjB,gBAAgB,EAChB,eAAe,EACf,gBAAgB,EAChB,SAAS,EACT,YAAY,EACZ,kBAAkB,EAClB,qBAAqB,EACrB,iCAAiC,EACjC,iBAAiB,EACjB,iBAAiB,EACjB,MAAM,GACP,MAAM,aAAa,CAAC"}
|
|
@@ -0,0 +1,27 @@
|
|
|
1
|
+
import { ParakeetModel } from '../model.js';
|
|
2
|
+
export interface MlxModelOptions {
|
|
3
|
+
/**
|
|
4
|
+
* 'interpolated' (default) matches NVIDIA's reference preprocessor.
|
|
5
|
+
* 'floor' reproduces parakeet-mlx's filterbank, which collapses 13 of 128 mel
|
|
6
|
+
* bins to zero — see docs/cuda.md.
|
|
7
|
+
*/
|
|
8
|
+
filterbank?: 'floor' | 'interpolated';
|
|
9
|
+
}
|
|
10
|
+
/**
|
|
11
|
+
* Load a safetensors Parakeet checkpoint from a local directory.
|
|
12
|
+
*
|
|
13
|
+
* @param dir directory holding `config.json` and `model.safetensors`
|
|
14
|
+
*/
|
|
15
|
+
export declare function fromLocal(dir: string, options?: MlxModelOptions): ParakeetModel;
|
|
16
|
+
export interface FromPretrainedOptions extends MlxModelOptions {
|
|
17
|
+
cacheDir?: string;
|
|
18
|
+
onProgress?: (file: string, downloaded: number, total: number) => void;
|
|
19
|
+
}
|
|
20
|
+
/**
|
|
21
|
+
* Load a Parakeet checkpoint from a HuggingFace Hub repo or a local directory.
|
|
22
|
+
*
|
|
23
|
+
* @param hfIdOrPath HuggingFace repo id (e.g. "mlx-community/parakeet-tdt-0.6b-v3")
|
|
24
|
+
* or a local directory path.
|
|
25
|
+
*/
|
|
26
|
+
export declare function fromPretrained(hfIdOrPath: string, options?: FromPretrainedOptions): Promise<ParakeetModel>;
|
|
27
|
+
//# sourceMappingURL=load.d.ts.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"load.d.ts","sourceRoot":"","sources":["../../src/mlx/load.ts"],"names":[],"mappings":"AAoBA,OAAO,EAAE,aAAa,EAAE,MAAM,aAAa,CAAC;AAG5C,MAAM,WAAW,eAAe;IAC9B;;;;OAIG;IACH,UAAU,CAAC,EAAE,OAAO,GAAG,cAAc,CAAC;CACvC;AAwID;;;;GAIG;AACH,wBAAgB,SAAS,CAAC,GAAG,EAAE,MAAM,EAAE,OAAO,GAAE,eAAoB,GAAG,aAAa,CAQnF;AAED,MAAM,WAAW,qBAAsB,SAAQ,eAAe;IAC5D,QAAQ,CAAC,EAAE,MAAM,CAAC;IAClB,UAAU,CAAC,EAAE,CAAC,IAAI,EAAE,MAAM,EAAE,UAAU,EAAE,MAAM,EAAE,KAAK,EAAE,MAAM,KAAK,IAAI,CAAC;CACxE;AAED;;;;;GAKG;AACH,wBAAsB,cAAc,CAClC,UAAU,EAAE,MAAM,EAClB,OAAO,GAAE,qBAA0B,GAClC,OAAO,CAAC,aAAa,CAAC,CAexB"}
|
package/dist/mlx/load.js
ADDED
|
@@ -0,0 +1,161 @@
|
|
|
1
|
+
/**
|
|
2
|
+
* Loading MLX-weight Parakeet checkpoints into the backend-agnostic
|
|
3
|
+
* `ParakeetModel`.
|
|
4
|
+
*
|
|
5
|
+
* This is the single MLX entry point: it builds the Conformer encoder, the
|
|
6
|
+
* prediction network and the joint network directly from `config.json`, loads
|
|
7
|
+
* the safetensors weights into them, and wraps them behind `MlxBackend` so the
|
|
8
|
+
* MLX path drives the same decode loop, tokenizer, alignment and audio
|
|
9
|
+
* front-end as the ONNX path. MLX and ONNX therefore produce identical output
|
|
10
|
+
* from identical features.
|
|
11
|
+
*
|
|
12
|
+
* TDT and RNN-T checkpoints are supported (this is also all the ONNX export
|
|
13
|
+
* covers). CTC / hybrid TDT-CTC checkpoints are not.
|
|
14
|
+
*/
|
|
15
|
+
import fs from 'node:fs';
|
|
16
|
+
import path from 'node:path';
|
|
17
|
+
import { Conformer } from './conformer.js';
|
|
18
|
+
import { PredictNetwork, JointNetwork } from './rnnt.js';
|
|
19
|
+
import { MlxBackend } from './backend.js';
|
|
20
|
+
import { loadSafetensors, downloadFromHub } from './utils.js';
|
|
21
|
+
import { ParakeetModel } from '../model.js';
|
|
22
|
+
import { makePreprocessArgs } from '../audio.js';
|
|
23
|
+
/** Build the MLX modules and metadata from a parsed NeMo-style config. */
|
|
24
|
+
function buildFromConfig(config) {
|
|
25
|
+
const modelDefaults = config['model_defaults'] ?? {};
|
|
26
|
+
const encoderRaw = config['encoder'];
|
|
27
|
+
const decoderRaw = config['decoder'];
|
|
28
|
+
const jointRaw = config['joint'];
|
|
29
|
+
const decodingRaw = config['decoding'] ?? {};
|
|
30
|
+
if (!jointRaw || !decoderRaw?.['prednet']) {
|
|
31
|
+
throw new Error('Only TDT and RNN-T checkpoints are supported by the MLX loader ' +
|
|
32
|
+
'(config needs decoder.prednet and a joint network). ' +
|
|
33
|
+
'CTC / hybrid TDT-CTC checkpoints are not supported.');
|
|
34
|
+
}
|
|
35
|
+
const attContextSize = encoderRaw['att_context_size'] ?? null;
|
|
36
|
+
const encoderArgs = {
|
|
37
|
+
featIn: encoderRaw['feat_in'],
|
|
38
|
+
nLayers: encoderRaw['n_layers'],
|
|
39
|
+
dModel: encoderRaw['d_model'],
|
|
40
|
+
nHeads: encoderRaw['n_heads'],
|
|
41
|
+
ffExpansionFactor: encoderRaw['ff_expansion_factor'],
|
|
42
|
+
subsamplingFactor: encoderRaw['subsampling_factor'],
|
|
43
|
+
selfAttentionModel: encoderRaw['self_attention_model'],
|
|
44
|
+
subsampling: encoderRaw['subsampling'],
|
|
45
|
+
convKernelSize: encoderRaw['conv_kernel_size'],
|
|
46
|
+
subsamplingConvChannels: encoderRaw['subsampling_conv_channels'],
|
|
47
|
+
posEmbMaxLen: encoderRaw['pos_emb_max_len'],
|
|
48
|
+
causalDownsampling: encoderRaw['causal_downsampling'] ?? false,
|
|
49
|
+
useBias: encoderRaw['use_bias'] ?? true,
|
|
50
|
+
xscaling: encoderRaw['xscaling'] ?? false,
|
|
51
|
+
subsamplingConvChunkingFactor: encoderRaw['subsampling_conv_chunking_factor'] ?? 1,
|
|
52
|
+
attContextSize,
|
|
53
|
+
};
|
|
54
|
+
const prednetRaw = decoderRaw['prednet'] ?? {};
|
|
55
|
+
const predictArgs = {
|
|
56
|
+
blankAsPad: decoderRaw['blank_as_pad'],
|
|
57
|
+
vocabSize: decoderRaw['vocab_size'],
|
|
58
|
+
prednet: {
|
|
59
|
+
predHidden: prednetRaw['pred_hidden'],
|
|
60
|
+
predRnnLayers: prednetRaw['pred_rnn_layers'],
|
|
61
|
+
rnnHiddenSize: prednetRaw['rnn_hidden_size'],
|
|
62
|
+
},
|
|
63
|
+
};
|
|
64
|
+
const jointnetRaw = jointRaw['jointnet'] ?? {};
|
|
65
|
+
const jointArgs = {
|
|
66
|
+
numClasses: jointRaw['num_classes'],
|
|
67
|
+
vocabulary: jointRaw['vocabulary'],
|
|
68
|
+
jointnet: {
|
|
69
|
+
jointHidden: jointnetRaw['joint_hidden'],
|
|
70
|
+
activation: jointnetRaw['activation'],
|
|
71
|
+
encoderHidden: jointnetRaw['encoder_hidden'],
|
|
72
|
+
predHidden: jointnetRaw['pred_hidden'],
|
|
73
|
+
},
|
|
74
|
+
numExtraOutputs: jointRaw['num_extra_outputs'] ?? 0,
|
|
75
|
+
};
|
|
76
|
+
const durations = modelDefaults['tdt_durations'] ?? null;
|
|
77
|
+
const greedy = decodingRaw['greedy'];
|
|
78
|
+
const maxSymbols = greedy?.['max_symbols'] != null ? Number(greedy['max_symbols']) : null;
|
|
79
|
+
return {
|
|
80
|
+
encoder: new Conformer(encoderArgs),
|
|
81
|
+
predict: new PredictNetwork(predictArgs),
|
|
82
|
+
joint: new JointNetwork(jointArgs),
|
|
83
|
+
vocabulary: jointArgs.vocabulary,
|
|
84
|
+
durations,
|
|
85
|
+
maxSymbols,
|
|
86
|
+
encoderDim: encoderArgs.dModel,
|
|
87
|
+
subsamplingFactor: encoderArgs.subsamplingFactor,
|
|
88
|
+
};
|
|
89
|
+
}
|
|
90
|
+
/** Assemble a `ParakeetModel` from a config, a weights file, and options. */
|
|
91
|
+
function assemble(config, weightsPath, options) {
|
|
92
|
+
const built = buildFromConfig(config);
|
|
93
|
+
const weights = loadSafetensors(weightsPath);
|
|
94
|
+
built.encoder.loadWeights(weights, 'encoder');
|
|
95
|
+
built.predict.loadWeights(weights, 'decoder');
|
|
96
|
+
built.joint.loadWeights(weights, 'joint');
|
|
97
|
+
const pre = config['preprocessor'];
|
|
98
|
+
const preprocessor = makePreprocessArgs({
|
|
99
|
+
sampleRate: pre['sample_rate'],
|
|
100
|
+
normalize: pre['normalize'],
|
|
101
|
+
windowSize: pre['window_size'],
|
|
102
|
+
windowStride: pre['window_stride'],
|
|
103
|
+
window: pre['window'],
|
|
104
|
+
features: pre['features'],
|
|
105
|
+
nFft: pre['n_fft'],
|
|
106
|
+
dither: pre['dither'],
|
|
107
|
+
padTo: pre['pad_to'] ?? 0,
|
|
108
|
+
padValue: pre['pad_value'] ?? 0,
|
|
109
|
+
preemph: pre['preemph'],
|
|
110
|
+
magPower: pre['mag_power'] ?? 2.0,
|
|
111
|
+
filterbank: options.filterbank ?? 'interpolated',
|
|
112
|
+
});
|
|
113
|
+
const backend = new MlxBackend({
|
|
114
|
+
encoder: built.encoder,
|
|
115
|
+
predict: built.predict,
|
|
116
|
+
joint: built.joint,
|
|
117
|
+
encoderDim: built.encoderDim,
|
|
118
|
+
});
|
|
119
|
+
return new ParakeetModel({
|
|
120
|
+
backend,
|
|
121
|
+
preprocessor,
|
|
122
|
+
vocabulary: built.vocabulary,
|
|
123
|
+
durations: built.durations,
|
|
124
|
+
maxSymbols: built.maxSymbols,
|
|
125
|
+
subsamplingFactor: built.subsamplingFactor,
|
|
126
|
+
});
|
|
127
|
+
}
|
|
128
|
+
/**
|
|
129
|
+
* Load a safetensors Parakeet checkpoint from a local directory.
|
|
130
|
+
*
|
|
131
|
+
* @param dir directory holding `config.json` and `model.safetensors`
|
|
132
|
+
*/
|
|
133
|
+
export function fromLocal(dir, options = {}) {
|
|
134
|
+
const configPath = path.join(dir, 'config.json');
|
|
135
|
+
const weightsPath = path.join(dir, 'model.safetensors');
|
|
136
|
+
if (!fs.existsSync(configPath))
|
|
137
|
+
throw new Error(`config.json not found in ${dir}`);
|
|
138
|
+
if (!fs.existsSync(weightsPath))
|
|
139
|
+
throw new Error(`model.safetensors not found in ${dir}`);
|
|
140
|
+
const config = JSON.parse(fs.readFileSync(configPath, 'utf8'));
|
|
141
|
+
return assemble(config, weightsPath, options);
|
|
142
|
+
}
|
|
143
|
+
/**
|
|
144
|
+
* Load a Parakeet checkpoint from a HuggingFace Hub repo or a local directory.
|
|
145
|
+
*
|
|
146
|
+
* @param hfIdOrPath HuggingFace repo id (e.g. "mlx-community/parakeet-tdt-0.6b-v3")
|
|
147
|
+
* or a local directory path.
|
|
148
|
+
*/
|
|
149
|
+
export async function fromPretrained(hfIdOrPath, options = {}) {
|
|
150
|
+
if (fs.existsSync(hfIdOrPath) && fs.statSync(hfIdOrPath).isDirectory()) {
|
|
151
|
+
return fromLocal(hfIdOrPath, options);
|
|
152
|
+
}
|
|
153
|
+
const makeProgress = (file) => options.onProgress
|
|
154
|
+
? (downloaded, total) => options.onProgress(file, downloaded, total)
|
|
155
|
+
: undefined;
|
|
156
|
+
const configPath = await downloadFromHub(hfIdOrPath, 'config.json', options.cacheDir, makeProgress('config.json'));
|
|
157
|
+
const weightsPath = await downloadFromHub(hfIdOrPath, 'model.safetensors', options.cacheDir, makeProgress('model.safetensors'));
|
|
158
|
+
const config = JSON.parse(fs.readFileSync(configPath, 'utf8'));
|
|
159
|
+
return assemble(config, weightsPath, options);
|
|
160
|
+
}
|
|
161
|
+
//# sourceMappingURL=load.js.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"load.js","sourceRoot":"","sources":["../../src/mlx/load.ts"],"names":[],"mappings":"AAAA;;;;;;;;;;;;;GAaG;AACH,OAAO,EAAE,MAAM,SAAS,CAAC;AACzB,OAAO,IAAI,MAAM,WAAW,CAAC;AAC7B,OAAO,EAAE,SAAS,EAAiB,MAAM,gBAAgB,CAAC;AAC1D,OAAO,EAAE,cAAc,EAAE,YAAY,EAA0B,MAAM,WAAW,CAAC;AACjF,OAAO,EAAE,UAAU,EAAE,MAAM,cAAc,CAAC;AAC1C,OAAO,EAAE,eAAe,EAAE,eAAe,EAAE,MAAM,YAAY,CAAC;AAC9D,OAAO,EAAE,aAAa,EAAE,MAAM,aAAa,CAAC;AAC5C,OAAO,EAAE,kBAAkB,EAAE,MAAM,aAAa,CAAC;AAsBjD,0EAA0E;AAC1E,SAAS,eAAe,CAAC,MAA+B;IACtD,MAAM,aAAa,GAAI,MAAM,CAAC,gBAAgB,CAA6B,IAAI,EAAE,CAAC;IAClF,MAAM,UAAU,GAAG,MAAM,CAAC,SAAS,CAA4B,CAAC;IAChE,MAAM,UAAU,GAAG,MAAM,CAAC,SAAS,CAA4B,CAAC;IAChE,MAAM,QAAQ,GAAG,MAAM,CAAC,OAAO,CAAwC,CAAC;IACxE,MAAM,WAAW,GAAI,MAAM,CAAC,UAAU,CAA6B,IAAI,EAAE,CAAC;IAE1E,IAAI,CAAC,QAAQ,IAAI,CAAC,UAAU,EAAE,CAAC,SAAS,CAAC,EAAE,CAAC;QAC1C,MAAM,IAAI,KAAK,CACb,iEAAiE;YACjE,sDAAsD;YACtD,qDAAqD,CACtD,CAAC;IACJ,CAAC;IAED,MAAM,cAAc,GAAI,UAAU,CAAC,kBAAkB,CAA6B,IAAI,IAAI,CAAC;IAC3F,MAAM,WAAW,GAAkB;QACjC,MAAM,EAAE,UAAU,CAAC,SAAS,CAAW;QACvC,OAAO,EAAE,UAAU,CAAC,UAAU,CAAW;QACzC,MAAM,EAAE,UAAU,CAAC,SAAS,CAAW;QACvC,MAAM,EAAE,UAAU,CAAC,SAAS,CAAW;QACvC,iBAAiB,EAAE,UAAU,CAAC,qBAAqB,CAAW;QAC9D,iBAAiB,EAAE,UAAU,CAAC,oBAAoB,CAAW;QAC7D,kBAAkB,EAAE,UAAU,CAAC,sBAAsB,CAAW;QAChE,WAAW,EAAE,UAAU,CAAC,aAAa,CAAW;QAChD,cAAc,EAAE,UAAU,CAAC,kBAAkB,CAAW;QACxD,uBAAuB,EAAE,UAAU,CAAC,2BAA2B,CAAW;QAC1E,YAAY,EAAE,UAAU,CAAC,iBAAiB,CAAW;QACrD,kBAAkB,EAAG,UAAU,CAAC,qBAAqB,CAAa,IAAI,KAAK;QAC3E,OAAO,EAAG,UAAU,CAAC,UAAU,CAAa,IAAI,IAAI;QACpD,QAAQ,EAAG,UAAU,CAAC,UAAU,CAAa,IAAI,KAAK;QACtD,6BAA6B,EAAG,UAAU,CAAC,kCAAkC,CAAY,IAAI,CAAC;QAC9F,cAAc;KACf,CAAC;IAEF,MAAM,UAAU,GAAI,UAAU,CAAC,SAAS,CAA6B,IAAI,EAAE,CAAC;IAC5E,MAAM,WAAW,GAAgB;QAC/B,UAAU,EAAE,UAAU,CAAC,cAAc,CAAY;QACjD,SAAS,EAAE,UAAU,CAAC,YAAY,CAAW;QAC7C,OAAO,EAAE;YACP,UAAU,EAAE,UAAU,CAAC,aAAa,CAAW;YAC/C,aAAa,EAAE,UAAU,CAAC,iBAAiB,CAAW;YACtD,aAAa,EAAE,UAAU,CAAC,iBAAiB,CAAuB;SACnE;KACF,CAAC;IAEF,MAAM,WAAW,GAAI,QAAQ,CAAC,UAAU,CAA6B,IAAI,EAAE,CAAC;IAC5E,MAAM,SAAS,GAAc;QAC3B,UAAU,EAAE,QAAQ,CAAC,aAAa,CAAW;QAC7C,UAAU,EAAE,QAAQ,CAAC,YAAY,CAAa;QAC9C,QAAQ,EAAE;YACR,WAAW,EAAE,WAAW,CAAC,cAAc,CAAW;YAClD,UAAU,EAAE,WAAW,CAAC,YAAY,CAAW;YAC/C,aAAa,EAAE,WAAW,CAAC,gBAAgB,CAAW;YACtD,UAAU,EAAE,WAAW,CAAC,aAAa,CAAW;SACjD;QACD,eAAe,EAAG,QAAQ,CAAC,mBAAmB,CAAY,IAAI,CAAC;KAChE,CAAC;IAEF,MAAM,SAAS,GAAI,aAAa,CAAC,eAAe,CAA0B,IAAI,IAAI,CAAC;IACnF,MAAM,MAAM,GAAG,WAAW,CAAC,QAAQ,CAAwC,CAAC;IAC5E,MAAM,UAAU,GAAG,MAAM,EAAE,CAAC,aAAa,CAAC,IAAI,IAAI,CAAC,CAAC,CAAC,MAAM,CAAC,MAAM,CAAC,aAAa,CAAC,CAAC,CAAC,CAAC,CAAC,IAAI,CAAC;IAE1F,OAAO;QACL,OAAO,EAAE,IAAI,SAAS,CAAC,WAAW,CAAC;QACnC,OAAO,EAAE,IAAI,cAAc,CAAC,WAAW,CAAC;QACxC,KAAK,EAAE,IAAI,YAAY,CAAC,SAAS,CAAC;QAClC,UAAU,EAAE,SAAS,CAAC,UAAU;QAChC,SAAS;QACT,UAAU;QACV,UAAU,EAAE,WAAW,CAAC,MAAM;QAC9B,iBAAiB,EAAE,WAAW,CAAC,iBAAiB;KACjD,CAAC;AACJ,CAAC;AAED,6EAA6E;AAC7E,SAAS,QAAQ,CACf,MAA+B,EAC/B,WAAmB,EACnB,OAAwB;IAExB,MAAM,KAAK,GAAG,eAAe,CAAC,MAAM,CAAC,CAAC;IAEtC,MAAM,OAAO,GAAG,eAAe,CAAC,WAAW,CAAC,CAAC;IAC7C,KAAK,CAAC,OAAO,CAAC,WAAW,CAAC,OAAO,EAAE,SAAS,CAAC,CAAC;IAC9C,KAAK,CAAC,OAAO,CAAC,WAAW,CAAC,OAAO,EAAE,SAAS,CAAC,CAAC;IAC9C,KAAK,CAAC,KAAK,CAAC,WAAW,CAAC,OAAO,EAAE,OAAO,CAAC,CAAC;IAE1C,MAAM,GAAG,GAAG,MAAM,CAAC,cAAc,CAA4B,CAAC;IAC9D,MAAM,YAAY,GAAG,kBAAkB,CAAC;QACtC,UAAU,EAAE,GAAG,CAAC,aAAa,CAAW;QACxC,SAAS,EAAE,GAAG,CAAC,WAAW,CAAW;QACrC,UAAU,EAAE,GAAG,CAAC,aAAa,CAAW;QACxC,YAAY,EAAE,GAAG,CAAC,eAAe,CAAW;QAC5C,MAAM,EAAE,GAAG,CAAC,QAAQ,CAAW;QAC/B,QAAQ,EAAE,GAAG,CAAC,UAAU,CAAW;QACnC,IAAI,EAAE,GAAG,CAAC,OAAO,CAAW;QAC5B,MAAM,EAAE,GAAG,CAAC,QAAQ,CAAW;QAC/B,KAAK,EAAG,GAAG,CAAC,QAAQ,CAAY,IAAI,CAAC;QACrC,QAAQ,EAAG,GAAG,CAAC,WAAW,CAAY,IAAI,CAAC;QAC3C,OAAO,EAAE,GAAG,CAAC,SAAS,CAAkB;QACxC,QAAQ,EAAG,GAAG,CAAC,WAAW,CAAY,IAAI,GAAG;QAC7C,UAAU,EAAE,OAAO,CAAC,UAAU,IAAI,cAAc;KACjD,CAAC,CAAC;IAEH,MAAM,OAAO,GAAG,IAAI,UAAU,CAAC;QAC7B,OAAO,EAAE,KAAK,CAAC,OAAO;QACtB,OAAO,EAAE,KAAK,CAAC,OAAO;QACtB,KAAK,EAAE,KAAK,CAAC,KAAK;QAClB,UAAU,EAAE,KAAK,CAAC,UAAU;KAC7B,CAAC,CAAC;IAEH,OAAO,IAAI,aAAa,CAAC;QACvB,OAAO;QACP,YAAY;QACZ,UAAU,EAAE,KAAK,CAAC,UAAU;QAC5B,SAAS,EAAE,KAAK,CAAC,SAAS;QAC1B,UAAU,EAAE,KAAK,CAAC,UAAU;QAC5B,iBAAiB,EAAE,KAAK,CAAC,iBAAiB;KAC3C,CAAC,CAAC;AACL,CAAC;AAED;;;;GAIG;AACH,MAAM,UAAU,SAAS,CAAC,GAAW,EAAE,UAA2B,EAAE;IAClE,MAAM,UAAU,GAAG,IAAI,CAAC,IAAI,CAAC,GAAG,EAAE,aAAa,CAAC,CAAC;IACjD,MAAM,WAAW,GAAG,IAAI,CAAC,IAAI,CAAC,GAAG,EAAE,mBAAmB,CAAC,CAAC;IACxD,IAAI,CAAC,EAAE,CAAC,UAAU,CAAC,UAAU,CAAC;QAAE,MAAM,IAAI,KAAK,CAAC,4BAA4B,GAAG,EAAE,CAAC,CAAC;IACnF,IAAI,CAAC,EAAE,CAAC,UAAU,CAAC,WAAW,CAAC;QAAE,MAAM,IAAI,KAAK,CAAC,kCAAkC,GAAG,EAAE,CAAC,CAAC;IAE1F,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,EAAE,CAAC,YAAY,CAAC,UAAU,EAAE,MAAM,CAAC,CAA4B,CAAC;IAC1F,OAAO,QAAQ,CAAC,MAAM,EAAE,WAAW,EAAE,OAAO,CAAC,CAAC;AAChD,CAAC;AAOD;;;;;GAKG;AACH,MAAM,CAAC,KAAK,UAAU,cAAc,CAClC,UAAkB,EAClB,UAAiC,EAAE;IAEnC,IAAI,EAAE,CAAC,UAAU,CAAC,UAAU,CAAC,IAAI,EAAE,CAAC,QAAQ,CAAC,UAAU,CAAC,CAAC,WAAW,EAAE,EAAE,CAAC;QACvE,OAAO,SAAS,CAAC,UAAU,EAAE,OAAO,CAAC,CAAC;IACxC,CAAC;IAED,MAAM,YAAY,GAAG,CAAC,IAAY,EAAE,EAAE,CACpC,OAAO,CAAC,UAAU;QAChB,CAAC,CAAC,CAAC,UAAkB,EAAE,KAAa,EAAE,EAAE,CAAC,OAAO,CAAC,UAAW,CAAC,IAAI,EAAE,UAAU,EAAE,KAAK,CAAC;QACrF,CAAC,CAAC,SAAS,CAAC;IAEhB,MAAM,UAAU,GAAG,MAAM,eAAe,CAAC,UAAU,EAAE,aAAa,EAAE,OAAO,CAAC,QAAQ,EAAE,YAAY,CAAC,aAAa,CAAC,CAAC,CAAC;IACnH,MAAM,WAAW,GAAG,MAAM,eAAe,CAAC,UAAU,EAAE,mBAAmB,EAAE,OAAO,CAAC,QAAQ,EAAE,YAAY,CAAC,mBAAmB,CAAC,CAAC,CAAC;IAEhI,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,EAAE,CAAC,YAAY,CAAC,UAAU,EAAE,MAAM,CAAC,CAA4B,CAAC;IAC1F,OAAO,QAAQ,CAAC,MAAM,EAAE,WAAW,EAAE,OAAO,CAAC,CAAC;AAChD,CAAC"}
|