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,98 @@
|
|
|
1
|
+
import { MxArray } from '@mlx-node/core';
|
|
2
|
+
function s(...dims) {
|
|
3
|
+
return BigInt64Array.from(dims.map(BigInt));
|
|
4
|
+
}
|
|
5
|
+
/**
|
|
6
|
+
* Cache for a single Conformer layer — stores attention K/V and conv state.
|
|
7
|
+
*/
|
|
8
|
+
export class ConformerCache {
|
|
9
|
+
keys = null;
|
|
10
|
+
values = null;
|
|
11
|
+
conv = null;
|
|
12
|
+
offset = 0;
|
|
13
|
+
step = 256;
|
|
14
|
+
updateAndFetchKV(k, v) {
|
|
15
|
+
if (this.keys === null || this.values === null) {
|
|
16
|
+
// Allocate initial cache
|
|
17
|
+
const kShape = k.shape();
|
|
18
|
+
const batch = Number(kShape[0]);
|
|
19
|
+
const heads = Number(kShape[1]);
|
|
20
|
+
const headDim = Number(kShape[3]);
|
|
21
|
+
const initLen = Math.ceil(Number(kShape[2]) / this.step) * this.step;
|
|
22
|
+
this.keys = MxArray.zeros(s(batch, heads, initLen, headDim), null);
|
|
23
|
+
this.values = MxArray.zeros(s(batch, heads, initLen, headDim), null);
|
|
24
|
+
}
|
|
25
|
+
const newSeq = Number(k.shape()[2]);
|
|
26
|
+
// Grow if needed
|
|
27
|
+
while (this.offset + newSeq > Number(this.keys.shape()[2])) {
|
|
28
|
+
const batch = Number(this.keys.shape()[0]);
|
|
29
|
+
const heads = Number(this.keys.shape()[1]);
|
|
30
|
+
const headDim = Number(this.keys.shape()[3]);
|
|
31
|
+
const extra = MxArray.zeros(s(batch, heads, this.step, headDim), null);
|
|
32
|
+
this.keys = MxArray.concatenate(this.keys, extra, 2);
|
|
33
|
+
this.values = MxArray.concatenate(this.values, extra, 2);
|
|
34
|
+
}
|
|
35
|
+
// Write new K and V into cache at [offset : offset+newSeq]
|
|
36
|
+
// MLX doesn't have scatter_nd; we rebuild by concatenation
|
|
37
|
+
const before = this.offset;
|
|
38
|
+
const total = Number(this.keys.shape()[2]);
|
|
39
|
+
const keysBefore = this.keys.slice(BigInt64Array.from([0n, 0n, 0n, 0n]), BigInt64Array.from([this.keys.shape()[0], this.keys.shape()[1], BigInt(before), this.keys.shape()[3]]));
|
|
40
|
+
const keysAfter = this.keys.slice(BigInt64Array.from([0n, 0n, BigInt(before + newSeq), 0n]), BigInt64Array.from([this.keys.shape()[0], this.keys.shape()[1], BigInt(total), this.keys.shape()[3]]));
|
|
41
|
+
this.keys = MxArray.concatenateMany([keysBefore, k, keysAfter], 2);
|
|
42
|
+
const valsBefore = this.values.slice(BigInt64Array.from([0n, 0n, 0n, 0n]), BigInt64Array.from([this.values.shape()[0], this.values.shape()[1], BigInt(before), this.values.shape()[3]]));
|
|
43
|
+
const valsAfter = this.values.slice(BigInt64Array.from([0n, 0n, BigInt(before + newSeq), 0n]), BigInt64Array.from([this.values.shape()[0], this.values.shape()[1], BigInt(total), this.values.shape()[3]]));
|
|
44
|
+
this.values = MxArray.concatenateMany([valsBefore, v, valsAfter], 2);
|
|
45
|
+
this.offset += newSeq;
|
|
46
|
+
const cachedK = this.keys.slice(BigInt64Array.from([0n, 0n, 0n, 0n]), BigInt64Array.from([this.keys.shape()[0], this.keys.shape()[1], BigInt(this.offset), this.keys.shape()[3]]));
|
|
47
|
+
const cachedV = this.values.slice(BigInt64Array.from([0n, 0n, 0n, 0n]), BigInt64Array.from([this.values.shape()[0], this.values.shape()[1], BigInt(this.offset), this.values.shape()[3]]));
|
|
48
|
+
return [cachedK, cachedV];
|
|
49
|
+
}
|
|
50
|
+
updateAndFetchConv(x, padding) {
|
|
51
|
+
// x: [batch, seq, channels]
|
|
52
|
+
if (this.conv === null) {
|
|
53
|
+
// Pad with zeros on the left
|
|
54
|
+
const xShape = x.shape();
|
|
55
|
+
const batch = Number(xShape[0]);
|
|
56
|
+
const ch = Number(xShape[2]);
|
|
57
|
+
const pad = MxArray.zeros(s(batch, padding, ch), null);
|
|
58
|
+
this.conv = MxArray.concatenate(pad, x, 1);
|
|
59
|
+
}
|
|
60
|
+
else {
|
|
61
|
+
// Keep only the last `padding` frames plus new input
|
|
62
|
+
const convLen = Number(this.conv.shape()[1]);
|
|
63
|
+
const keep = Math.min(padding, convLen);
|
|
64
|
+
const tail = this.conv.slice(BigInt64Array.from([0n, BigInt(convLen - keep), 0n]), this.conv.shape());
|
|
65
|
+
this.conv = MxArray.concatenate(tail, x, 1);
|
|
66
|
+
}
|
|
67
|
+
return this.conv;
|
|
68
|
+
}
|
|
69
|
+
}
|
|
70
|
+
/**
|
|
71
|
+
* Rotating cache for streaming Conformer: drops old frames beyond keep_size.
|
|
72
|
+
*/
|
|
73
|
+
export class RotatingConformerCache extends ConformerCache {
|
|
74
|
+
keepSize;
|
|
75
|
+
dropSize;
|
|
76
|
+
constructor(keepSize, cacheDrop) {
|
|
77
|
+
super();
|
|
78
|
+
this.keepSize = keepSize;
|
|
79
|
+
this.dropSize = cacheDrop;
|
|
80
|
+
}
|
|
81
|
+
updateAndFetchKV(k, v) {
|
|
82
|
+
const [cachedK, cachedV] = super.updateAndFetchKV(k, v);
|
|
83
|
+
// Trim to keepSize if we've accumulated too many frames
|
|
84
|
+
const currentLen = Number(cachedK.shape()[2]);
|
|
85
|
+
if (currentLen > this.keepSize) {
|
|
86
|
+
const start = currentLen - this.keepSize;
|
|
87
|
+
const trimmedK = cachedK.slice(BigInt64Array.from([0n, 0n, BigInt(start), 0n]), cachedK.shape());
|
|
88
|
+
const trimmedV = cachedV.slice(BigInt64Array.from([0n, 0n, BigInt(start), 0n]), cachedV.shape());
|
|
89
|
+
// Update internal state
|
|
90
|
+
this.keys = trimmedK;
|
|
91
|
+
this.values = trimmedV;
|
|
92
|
+
this.offset = this.keepSize;
|
|
93
|
+
return [trimmedK, trimmedV];
|
|
94
|
+
}
|
|
95
|
+
return [cachedK, cachedV];
|
|
96
|
+
}
|
|
97
|
+
}
|
|
98
|
+
//# sourceMappingURL=cache.js.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"cache.js","sourceRoot":"","sources":["../../src/mlx/cache.ts"],"names":[],"mappings":"AAAA,OAAO,EAAE,OAAO,EAAE,MAAM,gBAAgB,CAAC;AAEzC,SAAS,CAAC,CAAC,GAAG,IAAc;IAC1B,OAAO,aAAa,CAAC,IAAI,CAAC,IAAI,CAAC,GAAG,CAAC,MAAM,CAAC,CAAC,CAAC;AAC9C,CAAC;AAED;;GAEG;AACH,MAAM,OAAO,cAAc;IACzB,IAAI,GAAmB,IAAI,CAAC;IAC5B,MAAM,GAAmB,IAAI,CAAC;IAC9B,IAAI,GAAmB,IAAI,CAAC;IAC5B,MAAM,GAAG,CAAC,CAAC;IACM,IAAI,GAAG,GAAG,CAAC;IAE5B,gBAAgB,CAAC,CAAU,EAAE,CAAU;QACrC,IAAI,IAAI,CAAC,IAAI,KAAK,IAAI,IAAI,IAAI,CAAC,MAAM,KAAK,IAAI,EAAE,CAAC;YAC/C,yBAAyB;YACzB,MAAM,MAAM,GAAG,CAAC,CAAC,KAAK,EAAE,CAAC;YACzB,MAAM,KAAK,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;YAChC,MAAM,KAAK,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;YAChC,MAAM,OAAO,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;YAClC,MAAM,OAAO,GAAG,IAAI,CAAC,IAAI,CAAC,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,GAAG,IAAI,CAAC,IAAI,CAAC,GAAG,IAAI,CAAC,IAAI,CAAC;YAErE,IAAI,CAAC,IAAI,GAAG,OAAO,CAAC,KAAK,CAAC,CAAC,CAAC,KAAK,EAAE,KAAK,EAAE,OAAO,EAAE,OAAO,CAAC,EAAE,IAAI,CAAC,CAAC;YACnE,IAAI,CAAC,MAAM,GAAG,OAAO,CAAC,KAAK,CAAC,CAAC,CAAC,KAAK,EAAE,KAAK,EAAE,OAAO,EAAE,OAAO,CAAC,EAAE,IAAI,CAAC,CAAC;QACvE,CAAC;QAED,MAAM,MAAM,GAAG,MAAM,CAAC,CAAC,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC;QAEpC,iBAAiB;QACjB,OAAO,IAAI,CAAC,MAAM,GAAG,MAAM,GAAG,MAAM,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,EAAE,CAAC;YAC3D,MAAM,KAAK,GAAG,MAAM,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC;YAC3C,MAAM,KAAK,GAAG,MAAM,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC;YAC3C,MAAM,OAAO,GAAG,MAAM,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC;YAC7C,MAAM,KAAK,GAAG,OAAO,CAAC,KAAK,CAAC,CAAC,CAAC,KAAK,EAAE,KAAK,EAAE,IAAI,CAAC,IAAI,EAAE,OAAO,CAAC,EAAE,IAAI,CAAC,CAAC;YACvE,IAAI,CAAC,IAAI,GAAG,OAAO,CAAC,WAAW,CAAC,IAAI,CAAC,IAAI,EAAE,KAAK,EAAE,CAAC,CAAC,CAAC;YACrD,IAAI,CAAC,MAAM,GAAG,OAAO,CAAC,WAAW,CAAC,IAAI,CAAC,MAAM,EAAE,KAAK,EAAE,CAAC,CAAC,CAAC;QAC3D,CAAC;QAED,2DAA2D;QAC3D,2DAA2D;QAC3D,MAAM,MAAM,GAAG,IAAI,CAAC,MAAM,CAAC;QAC3B,MAAM,KAAK,GAAG,MAAM,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC;QAE3C,MAAM,UAAU,GAAG,IAAI,CAAC,IAAI,CAAC,KAAK,CAChC,aAAa,CAAC,IAAI,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,CAAC,CAAC,EACpC,aAAa,CAAC,IAAI,CAAC,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,MAAM,CAAC,MAAM,CAAC,EAAE,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC,CACvG,CAAC;QACF,MAAM,SAAS,GAAG,IAAI,CAAC,IAAI,CAAC,KAAK,CAC/B,aAAa,CAAC,IAAI,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,MAAM,CAAC,MAAM,GAAG,MAAM,CAAC,EAAE,EAAE,CAAC,CAAC,EACzD,aAAa,CAAC,IAAI,CAAC,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,MAAM,CAAC,KAAK,CAAC,EAAE,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC,CACtG,CAAC;QACF,IAAI,CAAC,IAAI,GAAG,OAAO,CAAC,eAAe,CAAC,CAAC,UAAU,EAAE,CAAC,EAAE,SAAS,CAAC,EAAE,CAAC,CAAC,CAAC;QAEnE,MAAM,UAAU,GAAG,IAAI,CAAC,MAAM,CAAC,KAAK,CAClC,aAAa,CAAC,IAAI,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,CAAC,CAAC,EACpC,aAAa,CAAC,IAAI,CAAC,CAAC,IAAI,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,IAAI,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,MAAM,CAAC,MAAM,CAAC,EAAE,IAAI,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC,CAC7G,CAAC;QACF,MAAM,SAAS,GAAG,IAAI,CAAC,MAAM,CAAC,KAAK,CACjC,aAAa,CAAC,IAAI,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,MAAM,CAAC,MAAM,GAAG,MAAM,CAAC,EAAE,EAAE,CAAC,CAAC,EACzD,aAAa,CAAC,IAAI,CAAC,CAAC,IAAI,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,IAAI,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,MAAM,CAAC,KAAK,CAAC,EAAE,IAAI,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC,CAC5G,CAAC;QACF,IAAI,CAAC,MAAM,GAAG,OAAO,CAAC,eAAe,CAAC,CAAC,UAAU,EAAE,CAAC,EAAE,SAAS,CAAC,EAAE,CAAC,CAAC,CAAC;QAErE,IAAI,CAAC,MAAM,IAAI,MAAM,CAAC;QAEtB,MAAM,OAAO,GAAG,IAAI,CAAC,IAAI,CAAC,KAAK,CAC7B,aAAa,CAAC,IAAI,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,CAAC,CAAC,EACpC,aAAa,CAAC,IAAI,CAAC,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,MAAM,CAAC,IAAI,CAAC,MAAM,CAAC,EAAE,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC,CAC5G,CAAC;QACF,MAAM,OAAO,GAAG,IAAI,CAAC,MAAM,CAAC,KAAK,CAC/B,aAAa,CAAC,IAAI,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,EAAE,CAAC,CAAC,EACpC,aAAa,CAAC,IAAI,CAAC,CAAC,IAAI,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,IAAI,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,EAAE,MAAM,CAAC,IAAI,CAAC,MAAM,CAAC,EAAE,IAAI,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC,CAClH,CAAC;QAEF,OAAO,CAAC,OAAO,EAAE,OAAO,CAAC,CAAC;IAC5B,CAAC;IAED,kBAAkB,CAAC,CAAU,EAAE,OAAe;QAC5C,4BAA4B;QAC5B,IAAI,IAAI,CAAC,IAAI,KAAK,IAAI,EAAE,CAAC;YACvB,6BAA6B;YAC7B,MAAM,MAAM,GAAG,CAAC,CAAC,KAAK,EAAE,CAAC;YACzB,MAAM,KAAK,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;YAChC,MAAM,EAAE,GAAG,MAAM,CAAC,MAAM,CAAC,CAAC,CAAC,CAAC,CAAC;YAC7B,MAAM,GAAG,GAAG,OAAO,CAAC,KAAK,CAAC,CAAC,CAAC,KAAK,EAAE,OAAO,EAAE,EAAE,CAAC,EAAE,IAAI,CAAC,CAAC;YACvD,IAAI,CAAC,IAAI,GAAG,OAAO,CAAC,WAAW,CAAC,GAAG,EAAE,CAAC,EAAE,CAAC,CAAC,CAAC;QAC7C,CAAC;aAAM,CAAC;YACN,qDAAqD;YACrD,MAAM,OAAO,GAAG,MAAM,CAAC,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC;YAC7C,MAAM,IAAI,GAAG,IAAI,CAAC,GAAG,CAAC,OAAO,EAAE,OAAO,CAAC,CAAC;YACxC,MAAM,IAAI,GAAG,IAAI,CAAC,IAAI,CAAC,KAAK,CAC1B,aAAa,CAAC,IAAI,CAAC,CAAC,EAAE,EAAE,MAAM,CAAC,OAAO,GAAG,IAAI,CAAC,EAAE,EAAE,CAAC,CAAC,EACpD,IAAI,CAAC,IAAI,CAAC,KAAK,EAAE,CAClB,CAAC;YACF,IAAI,CAAC,IAAI,GAAG,OAAO,CAAC,WAAW,CAAC,IAAI,EAAE,CAAC,EAAE,CAAC,CAAC,CAAC;QAC9C,CAAC;QACD,OAAO,IAAI,CAAC,IAAI,CAAC;IACnB,CAAC;CACF;AAED;;GAEG;AACH,MAAM,OAAO,sBAAuB,SAAQ,cAAc;IACvC,QAAQ,CAAS;IACjB,QAAQ,CAAS;IAElC,YAAY,QAAgB,EAAE,SAAiB;QAC7C,KAAK,EAAE,CAAC;QACR,IAAI,CAAC,QAAQ,GAAG,QAAQ,CAAC;QACzB,IAAI,CAAC,QAAQ,GAAG,SAAS,CAAC;IAC5B,CAAC;IAED,gBAAgB,CAAC,CAAU,EAAE,CAAU;QACrC,MAAM,CAAC,OAAO,EAAE,OAAO,CAAC,GAAG,KAAK,CAAC,gBAAgB,CAAC,CAAC,EAAE,CAAC,CAAC,CAAC;QAExD,wDAAwD;QACxD,MAAM,UAAU,GAAG,MAAM,CAAC,OAAO,CAAC,KAAK,EAAE,CAAC,CAAC,CAAC,CAAC,CAAC;QAC9C,IAAI,UAAU,GAAG,IAAI,CAAC,QAAQ,EAAE,CAAC;YAC/B,MAAM,KAAK,GAAG,UAAU,GAAG,IAAI,CAAC,QAAQ,CAAC;YACzC,MAAM,QAAQ,GAAG,OAAO,CAAC,KAAK,CAC5B,aAAa,CAAC,IAAI,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,MAAM,CAAC,KAAK,CAAC,EAAE,EAAE,CAAC,CAAC,EAC/C,OAAO,CAAC,KAAK,EAAE,CAChB,CAAC;YACF,MAAM,QAAQ,GAAG,OAAO,CAAC,KAAK,CAC5B,aAAa,CAAC,IAAI,CAAC,CAAC,EAAE,EAAE,EAAE,EAAE,MAAM,CAAC,KAAK,CAAC,EAAE,EAAE,CAAC,CAAC,EAC/C,OAAO,CAAC,KAAK,EAAE,CAChB,CAAC;YACF,wBAAwB;YACxB,IAAI,CAAC,IAAI,GAAG,QAAQ,CAAC;YACrB,IAAI,CAAC,MAAM,GAAG,QAAQ,CAAC;YACvB,IAAI,CAAC,MAAM,GAAG,IAAI,CAAC,QAAQ,CAAC;YAC5B,OAAO,CAAC,QAAQ,EAAE,QAAQ,CAAC,CAAC;QAC9B,CAAC;QAED,OAAO,CAAC,OAAO,EAAE,OAAO,CAAC,CAAC;IAC5B,CAAC;CACF"}
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
#!/usr/bin/env node
|
|
2
|
+
/**
|
|
3
|
+
* parakeet.ts CLI
|
|
4
|
+
*
|
|
5
|
+
* Usage:
|
|
6
|
+
* parakeet file.wav
|
|
7
|
+
* parakeet file.wav --json
|
|
8
|
+
* parakeet file.wav --model mlx-community/parakeet-tdt-0.6b-v3
|
|
9
|
+
* parakeet --stream < pcm_f32le_16k.raw
|
|
10
|
+
*/
|
|
11
|
+
export {};
|
|
12
|
+
//# sourceMappingURL=cli.d.ts.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"cli.d.ts","sourceRoot":"","sources":["../../src/mlx/cli.ts"],"names":[],"mappings":";AACA;;;;;;;;GAQG"}
|
package/dist/mlx/cli.js
ADDED
|
@@ -0,0 +1,145 @@
|
|
|
1
|
+
#!/usr/bin/env node
|
|
2
|
+
/**
|
|
3
|
+
* parakeet.ts CLI
|
|
4
|
+
*
|
|
5
|
+
* Usage:
|
|
6
|
+
* parakeet file.wav
|
|
7
|
+
* parakeet file.wav --json
|
|
8
|
+
* parakeet file.wav --model mlx-community/parakeet-tdt-0.6b-v3
|
|
9
|
+
* parakeet --stream < pcm_f32le_16k.raw
|
|
10
|
+
*/
|
|
11
|
+
import { parseArgs } from 'node:util';
|
|
12
|
+
import { fromPretrained } from './load.js';
|
|
13
|
+
import { consumePcmStream } from '../model.js';
|
|
14
|
+
const DEFAULT_MODEL = 'mlx-community/parakeet-tdt-0.6b-v3';
|
|
15
|
+
// ---------------------------------------------------------------------------
|
|
16
|
+
// Arg parsing
|
|
17
|
+
// ---------------------------------------------------------------------------
|
|
18
|
+
const { values, positionals } = parseArgs({
|
|
19
|
+
options: {
|
|
20
|
+
model: { type: 'string', short: 'm', default: DEFAULT_MODEL },
|
|
21
|
+
json: { type: 'boolean', short: 'j', default: false },
|
|
22
|
+
stream: { type: 'boolean', short: 's', default: false },
|
|
23
|
+
help: { type: 'boolean', short: 'h', default: false },
|
|
24
|
+
},
|
|
25
|
+
allowPositionals: true,
|
|
26
|
+
args: process.argv.slice(2),
|
|
27
|
+
});
|
|
28
|
+
if (values.help) {
|
|
29
|
+
process.stdout.write(`
|
|
30
|
+
parakeet — Nvidia Parakeet ASR (parakeet.ts)
|
|
31
|
+
|
|
32
|
+
Usage:
|
|
33
|
+
parakeet <file.wav> [options]
|
|
34
|
+
parakeet --stream [options] < pcm_f32le_16k.raw
|
|
35
|
+
|
|
36
|
+
Options:
|
|
37
|
+
--model, -m <id> HuggingFace repo or local dir (default: ${DEFAULT_MODEL})
|
|
38
|
+
--json, -j Output full AlignedResult JSON instead of plain text
|
|
39
|
+
--stream, -s Read raw f32le mono 16kHz PCM from stdin
|
|
40
|
+
--help, -h Show this help
|
|
41
|
+
|
|
42
|
+
Exit codes: 0 success 1 file/IO error 2 model error
|
|
43
|
+
`.trimStart());
|
|
44
|
+
process.exit(0);
|
|
45
|
+
}
|
|
46
|
+
// ---------------------------------------------------------------------------
|
|
47
|
+
// Progress bar helpers
|
|
48
|
+
// ---------------------------------------------------------------------------
|
|
49
|
+
function progressBar(label, downloaded, total) {
|
|
50
|
+
if (!process.stderr.isTTY) {
|
|
51
|
+
return; // non-TTY: only emit the one-shot line (handled at call site)
|
|
52
|
+
}
|
|
53
|
+
if (total <= 0) {
|
|
54
|
+
process.stderr.write(`\r${label}: ${formatBytes(downloaded)}`);
|
|
55
|
+
return;
|
|
56
|
+
}
|
|
57
|
+
const pct = Math.min(100, Math.round((downloaded / total) * 100));
|
|
58
|
+
const barWidth = 30;
|
|
59
|
+
const filled = Math.round((pct / 100) * barWidth);
|
|
60
|
+
const bar = '█'.repeat(filled) + '░'.repeat(barWidth - filled);
|
|
61
|
+
process.stderr.write(`\r${label}: [${bar}] ${pct}% (${formatBytes(downloaded)}/${formatBytes(total)})`);
|
|
62
|
+
}
|
|
63
|
+
function formatBytes(n) {
|
|
64
|
+
if (n < 1024)
|
|
65
|
+
return `${n} B`;
|
|
66
|
+
if (n < 1024 * 1024)
|
|
67
|
+
return `${(n / 1024).toFixed(1)} KB`;
|
|
68
|
+
return `${(n / 1024 / 1024).toFixed(1)} MB`;
|
|
69
|
+
}
|
|
70
|
+
// ---------------------------------------------------------------------------
|
|
71
|
+
// Main
|
|
72
|
+
// ---------------------------------------------------------------------------
|
|
73
|
+
async function main() {
|
|
74
|
+
const modelId = values.model ?? DEFAULT_MODEL;
|
|
75
|
+
// Load model
|
|
76
|
+
process.stderr.write(`Loading model: ${modelId}\n`);
|
|
77
|
+
let lastFile = '';
|
|
78
|
+
const model = await fromPretrained(modelId, {
|
|
79
|
+
onProgress(file, downloaded, total) {
|
|
80
|
+
if (file !== lastFile) {
|
|
81
|
+
if (lastFile !== '')
|
|
82
|
+
process.stderr.write('\n'); // end previous file's line
|
|
83
|
+
if (!process.stderr.isTTY) {
|
|
84
|
+
process.stderr.write(`Downloading ${file}...\n`);
|
|
85
|
+
}
|
|
86
|
+
lastFile = file;
|
|
87
|
+
}
|
|
88
|
+
progressBar(file, downloaded, total);
|
|
89
|
+
if (downloaded >= total && total > 0) {
|
|
90
|
+
if (process.stderr.isTTY)
|
|
91
|
+
process.stderr.write('\n');
|
|
92
|
+
}
|
|
93
|
+
},
|
|
94
|
+
});
|
|
95
|
+
if (lastFile !== '' && process.stderr.isTTY) {
|
|
96
|
+
// Ensure cursor is on a new line after progress bars
|
|
97
|
+
process.stderr.write('\n');
|
|
98
|
+
}
|
|
99
|
+
process.stderr.write('Model ready.\n');
|
|
100
|
+
let result;
|
|
101
|
+
if (values.stream) {
|
|
102
|
+
// Streaming mode: read raw f32le PCM from stdin
|
|
103
|
+
process.stderr.write('Streaming PCM from stdin (16kHz mono f32le)...\n');
|
|
104
|
+
const stream = model.transcribeStream();
|
|
105
|
+
async function* stdinPcm() {
|
|
106
|
+
for await (const chunk of process.stdin) {
|
|
107
|
+
// Reinterpret raw bytes as Float32 (little-endian f32)
|
|
108
|
+
const buf = chunk instanceof Buffer ? chunk : Buffer.from(chunk);
|
|
109
|
+
// Ensure alignment: process only whole float32 values
|
|
110
|
+
const floatCount = Math.floor(buf.byteLength / 4);
|
|
111
|
+
if (floatCount > 0) {
|
|
112
|
+
yield new Float32Array(buf.buffer, buf.byteOffset, floatCount);
|
|
113
|
+
}
|
|
114
|
+
}
|
|
115
|
+
}
|
|
116
|
+
result = await consumePcmStream(stream, stdinPcm());
|
|
117
|
+
}
|
|
118
|
+
else {
|
|
119
|
+
// File mode
|
|
120
|
+
const filePath = positionals[0];
|
|
121
|
+
if (!filePath) {
|
|
122
|
+
process.stderr.write('Error: provide a WAV file path or use --stream\n');
|
|
123
|
+
process.exit(1);
|
|
124
|
+
}
|
|
125
|
+
process.stderr.write(`Transcribing: ${filePath}\n`);
|
|
126
|
+
result = await model.transcribe(filePath);
|
|
127
|
+
}
|
|
128
|
+
// Output
|
|
129
|
+
if (values.json) {
|
|
130
|
+
process.stdout.write(JSON.stringify(result, null, 2) + '\n');
|
|
131
|
+
}
|
|
132
|
+
else {
|
|
133
|
+
process.stdout.write(result.text + '\n');
|
|
134
|
+
}
|
|
135
|
+
}
|
|
136
|
+
main().catch(err => {
|
|
137
|
+
const msg = err instanceof Error ? err.message : String(err);
|
|
138
|
+
if (msg.includes('ENOENT') || msg.includes('EACCES') || msg.includes('Failed to download')) {
|
|
139
|
+
process.stderr.write(`Error: ${msg}\n`);
|
|
140
|
+
process.exit(1);
|
|
141
|
+
}
|
|
142
|
+
process.stderr.write(`Model error: ${msg}\n`);
|
|
143
|
+
process.exit(2);
|
|
144
|
+
});
|
|
145
|
+
//# sourceMappingURL=cli.js.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"cli.js","sourceRoot":"","sources":["../../src/mlx/cli.ts"],"names":[],"mappings":";AACA;;;;;;;;GAQG;AAEH,OAAO,EAAE,SAAS,EAAE,MAAM,WAAW,CAAC;AACtC,OAAO,EAAE,cAAc,EAAE,MAAM,WAAW,CAAC;AAC3C,OAAO,EAAE,gBAAgB,EAAE,MAAM,aAAa,CAAC;AAG/C,MAAM,aAAa,GAAG,oCAAoC,CAAC;AAE3D,8EAA8E;AAC9E,cAAc;AACd,8EAA8E;AAE9E,MAAM,EAAE,MAAM,EAAE,WAAW,EAAE,GAAG,SAAS,CAAC;IACxC,OAAO,EAAE;QACP,KAAK,EAAE,EAAE,IAAI,EAAE,QAAQ,EAAE,KAAK,EAAE,GAAG,EAAE,OAAO,EAAE,aAAa,EAAE;QAC7D,IAAI,EAAE,EAAE,IAAI,EAAE,SAAS,EAAE,KAAK,EAAE,GAAG,EAAE,OAAO,EAAE,KAAK,EAAE;QACrD,MAAM,EAAE,EAAE,IAAI,EAAE,SAAS,EAAE,KAAK,EAAE,GAAG,EAAE,OAAO,EAAE,KAAK,EAAE;QACvD,IAAI,EAAE,EAAE,IAAI,EAAE,SAAS,EAAE,KAAK,EAAE,GAAG,EAAE,OAAO,EAAE,KAAK,EAAE;KACtD;IACD,gBAAgB,EAAE,IAAI;IACtB,IAAI,EAAE,OAAO,CAAC,IAAI,CAAC,KAAK,CAAC,CAAC,CAAC;CAC5B,CAAC,CAAC;AAEH,IAAI,MAAM,CAAC,IAAI,EAAE,CAAC;IAChB,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC;;;;;;;;+DAQwC,aAAa;;;;;;CAM3E,CAAC,SAAS,EAAE,CAAC,CAAC;IACb,OAAO,CAAC,IAAI,CAAC,CAAC,CAAC,CAAC;AAClB,CAAC;AAED,8EAA8E;AAC9E,uBAAuB;AACvB,8EAA8E;AAE9E,SAAS,WAAW,CAAC,KAAa,EAAE,UAAkB,EAAE,KAAa;IACnE,IAAI,CAAC,OAAO,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC;QAC1B,OAAO,CAAC,8DAA8D;IACxE,CAAC;IACD,IAAI,KAAK,IAAI,CAAC,EAAE,CAAC;QACf,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,KAAK,KAAK,KAAK,WAAW,CAAC,UAAU,CAAC,EAAE,CAAC,CAAC;QAC/D,OAAO;IACT,CAAC;IACD,MAAM,GAAG,GAAG,IAAI,CAAC,GAAG,CAAC,GAAG,EAAE,IAAI,CAAC,KAAK,CAAC,CAAC,UAAU,GAAG,KAAK,CAAC,GAAG,GAAG,CAAC,CAAC,CAAC;IAClE,MAAM,QAAQ,GAAG,EAAE,CAAC;IACpB,MAAM,MAAM,GAAG,IAAI,CAAC,KAAK,CAAC,CAAC,GAAG,GAAG,GAAG,CAAC,GAAG,QAAQ,CAAC,CAAC;IAClD,MAAM,GAAG,GAAG,GAAG,CAAC,MAAM,CAAC,MAAM,CAAC,GAAG,GAAG,CAAC,MAAM,CAAC,QAAQ,GAAG,MAAM,CAAC,CAAC;IAC/D,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,KAAK,KAAK,MAAM,GAAG,KAAK,GAAG,MAAM,WAAW,CAAC,UAAU,CAAC,IAAI,WAAW,CAAC,KAAK,CAAC,GAAG,CAAC,CAAC;AAC1G,CAAC;AAED,SAAS,WAAW,CAAC,CAAS;IAC5B,IAAI,CAAC,GAAG,IAAI;QAAE,OAAO,GAAG,CAAC,IAAI,CAAC;IAC9B,IAAI,CAAC,GAAG,IAAI,GAAG,IAAI;QAAE,OAAO,GAAG,CAAC,CAAC,GAAG,IAAI,CAAC,CAAC,OAAO,CAAC,CAAC,CAAC,KAAK,CAAC;IAC1D,OAAO,GAAG,CAAC,CAAC,GAAG,IAAI,GAAG,IAAI,CAAC,CAAC,OAAO,CAAC,CAAC,CAAC,KAAK,CAAC;AAC9C,CAAC;AAED,8EAA8E;AAC9E,OAAO;AACP,8EAA8E;AAE9E,KAAK,UAAU,IAAI;IACjB,MAAM,OAAO,GAAG,MAAM,CAAC,KAAK,IAAI,aAAa,CAAC;IAE9C,aAAa;IACb,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,kBAAkB,OAAO,IAAI,CAAC,CAAC;IAEpD,IAAI,QAAQ,GAAG,EAAE,CAAC;IAClB,MAAM,KAAK,GAAG,MAAM,cAAc,CAAC,OAAO,EAAE;QAC1C,UAAU,CAAC,IAAI,EAAE,UAAU,EAAE,KAAK;YAChC,IAAI,IAAI,KAAK,QAAQ,EAAE,CAAC;gBACtB,IAAI,QAAQ,KAAK,EAAE;oBAAE,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,IAAI,CAAC,CAAC,CAAC,2BAA2B;gBAC5E,IAAI,CAAC,OAAO,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC;oBAC1B,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,eAAe,IAAI,OAAO,CAAC,CAAC;gBACnD,CAAC;gBACD,QAAQ,GAAG,IAAI,CAAC;YAClB,CAAC;YACD,WAAW,CAAC,IAAI,EAAE,UAAU,EAAE,KAAK,CAAC,CAAC;YACrC,IAAI,UAAU,IAAI,KAAK,IAAI,KAAK,GAAG,CAAC,EAAE,CAAC;gBACrC,IAAI,OAAO,CAAC,MAAM,CAAC,KAAK;oBAAE,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,IAAI,CAAC,CAAC;YACvD,CAAC;QACH,CAAC;KACF,CAAC,CAAC;IAEH,IAAI,QAAQ,KAAK,EAAE,IAAI,OAAO,CAAC,MAAM,CAAC,KAAK,EAAE,CAAC;QAC5C,qDAAqD;QACrD,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,IAAI,CAAC,CAAC;IAC7B,CAAC;IAED,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,gBAAgB,CAAC,CAAC;IAEvC,IAAI,MAAqB,CAAC;IAE1B,IAAI,MAAM,CAAC,MAAM,EAAE,CAAC;QAClB,gDAAgD;QAChD,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,kDAAkD,CAAC,CAAC;QAEzE,MAAM,MAAM,GAAG,KAAK,CAAC,gBAAgB,EAAE,CAAC;QAExC,KAAK,SAAS,CAAC,CAAC,QAAQ;YACtB,IAAI,KAAK,EAAE,MAAM,KAAK,IAAI,OAAO,CAAC,KAA8B,EAAE,CAAC;gBACjE,uDAAuD;gBACvD,MAAM,GAAG,GAAG,KAAK,YAAY,MAAM,CAAC,CAAC,CAAC,KAAK,CAAC,CAAC,CAAC,MAAM,CAAC,IAAI,CAAC,KAAK,CAAC,CAAC;gBACjE,sDAAsD;gBACtD,MAAM,UAAU,GAAG,IAAI,CAAC,KAAK,CAAC,GAAG,CAAC,UAAU,GAAG,CAAC,CAAC,CAAC;gBAClD,IAAI,UAAU,GAAG,CAAC,EAAE,CAAC;oBACnB,MAAM,IAAI,YAAY,CAAC,GAAG,CAAC,MAAM,EAAE,GAAG,CAAC,UAAU,EAAE,UAAU,CAAC,CAAC;gBACjE,CAAC;YACH,CAAC;QACH,CAAC;QAED,MAAM,GAAG,MAAM,gBAAgB,CAAC,MAAM,EAAE,QAAQ,EAAE,CAAC,CAAC;IAEtD,CAAC;SAAM,CAAC;QACN,YAAY;QACZ,MAAM,QAAQ,GAAG,WAAW,CAAC,CAAC,CAAC,CAAC;QAChC,IAAI,CAAC,QAAQ,EAAE,CAAC;YACd,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,kDAAkD,CAAC,CAAC;YACzE,OAAO,CAAC,IAAI,CAAC,CAAC,CAAC,CAAC;QAClB,CAAC;QAED,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,iBAAiB,QAAQ,IAAI,CAAC,CAAC;QACpD,MAAM,GAAG,MAAM,KAAK,CAAC,UAAU,CAAC,QAAQ,CAAC,CAAC;IAC5C,CAAC;IAED,SAAS;IACT,IAAI,MAAM,CAAC,IAAI,EAAE,CAAC;QAChB,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,IAAI,CAAC,SAAS,CAAC,MAAM,EAAE,IAAI,EAAE,CAAC,CAAC,GAAG,IAAI,CAAC,CAAC;IAC/D,CAAC;SAAM,CAAC;QACN,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,MAAM,CAAC,IAAI,GAAG,IAAI,CAAC,CAAC;IAC3C,CAAC;AACH,CAAC;AAED,IAAI,EAAE,CAAC,KAAK,CAAC,GAAG,CAAC,EAAE;IACjB,MAAM,GAAG,GAAW,GAAG,YAAY,KAAK,CAAC,CAAC,CAAC,GAAG,CAAC,OAAO,CAAC,CAAC,CAAC,MAAM,CAAC,GAAG,CAAC,CAAC;IACrE,IAAI,GAAG,CAAC,QAAQ,CAAC,QAAQ,CAAC,IAAI,GAAG,CAAC,QAAQ,CAAC,QAAQ,CAAC,IAAI,GAAG,CAAC,QAAQ,CAAC,oBAAoB,CAAC,EAAE,CAAC;QAC3F,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,UAAU,GAAG,IAAI,CAAC,CAAC;QACxC,OAAO,CAAC,IAAI,CAAC,CAAC,CAAC,CAAC;IAClB,CAAC;IACD,OAAO,CAAC,MAAM,CAAC,KAAK,CAAC,gBAAgB,GAAG,IAAI,CAAC,CAAC;IAC9C,OAAO,CAAC,IAAI,CAAC,CAAC,CAAC,CAAC;AAClB,CAAC,CAAC,CAAC"}
|
|
@@ -0,0 +1,81 @@
|
|
|
1
|
+
import { MxArray } from '@mlx-node/core';
|
|
2
|
+
import { Module, WeightMap, Linear, LayerNorm, BatchNorm, Conv1d, Conv2d } from './nn.js';
|
|
3
|
+
import { MultiHeadAttention, RelPositionMultiHeadAttention, RelPositionMultiHeadLocalAttention, RelPositionalEncoding, LocalRelPositionalEncoding } from './attention.js';
|
|
4
|
+
import { ConformerCache } from './cache.js';
|
|
5
|
+
export interface ConformerArgs {
|
|
6
|
+
featIn: number;
|
|
7
|
+
nLayers: number;
|
|
8
|
+
dModel: number;
|
|
9
|
+
nHeads: number;
|
|
10
|
+
ffExpansionFactor: number;
|
|
11
|
+
subsamplingFactor: number;
|
|
12
|
+
selfAttentionModel: string;
|
|
13
|
+
subsampling: string;
|
|
14
|
+
convKernelSize: number;
|
|
15
|
+
subsamplingConvChannels: number;
|
|
16
|
+
posEmbMaxLen: number;
|
|
17
|
+
causalDownsampling?: boolean;
|
|
18
|
+
useBias?: boolean;
|
|
19
|
+
xscaling?: boolean;
|
|
20
|
+
subsamplingConvChunkingFactor?: number;
|
|
21
|
+
attContextSize?: [number, number] | null;
|
|
22
|
+
}
|
|
23
|
+
declare class FeedForward extends Module {
|
|
24
|
+
linear1: Linear;
|
|
25
|
+
linear2: Linear;
|
|
26
|
+
constructor(dModel: number, dFf: number, useBias: boolean);
|
|
27
|
+
forward(x: MxArray): MxArray;
|
|
28
|
+
loadWeights(weights: WeightMap, prefix: string): void;
|
|
29
|
+
}
|
|
30
|
+
declare class Convolution extends Module {
|
|
31
|
+
readonly padding: number;
|
|
32
|
+
pointwiseConv1: Conv1d;
|
|
33
|
+
depthwiseConv: Conv1d;
|
|
34
|
+
batchNorm: BatchNorm;
|
|
35
|
+
pointwiseConv2: Conv1d;
|
|
36
|
+
constructor(args: ConformerArgs);
|
|
37
|
+
forward(x: MxArray, cache: ConformerCache | null): MxArray;
|
|
38
|
+
loadWeights(weights: WeightMap, prefix: string): void;
|
|
39
|
+
}
|
|
40
|
+
type AttentionModel = 'rel_pos' | 'rel_pos_local_attn' | 'normal';
|
|
41
|
+
export declare class ConformerBlock extends Module {
|
|
42
|
+
normFF1: LayerNorm;
|
|
43
|
+
ff1: FeedForward;
|
|
44
|
+
normSelfAtt: LayerNorm;
|
|
45
|
+
selfAttn: MultiHeadAttention | RelPositionMultiHeadAttention | RelPositionMultiHeadLocalAttention;
|
|
46
|
+
normConv: LayerNorm;
|
|
47
|
+
conv: Convolution;
|
|
48
|
+
normFF2: LayerNorm;
|
|
49
|
+
ff2: FeedForward;
|
|
50
|
+
normOut: LayerNorm;
|
|
51
|
+
private readonly args;
|
|
52
|
+
constructor(args: ConformerArgs);
|
|
53
|
+
private buildAttention;
|
|
54
|
+
setAttentionModel(name: AttentionModel, contextSize?: [number, number]): void;
|
|
55
|
+
forward(x: MxArray, posEmb: MxArray | null, mask: MxArray | null, cache: ConformerCache | null): MxArray;
|
|
56
|
+
loadWeights(weights: WeightMap, prefix: string): void;
|
|
57
|
+
}
|
|
58
|
+
declare class DwStridingSubsampling extends Module {
|
|
59
|
+
private readonly samplingNum;
|
|
60
|
+
private readonly stride;
|
|
61
|
+
private readonly kernelSize;
|
|
62
|
+
private readonly padding;
|
|
63
|
+
convLayers: Array<Conv2d | null>;
|
|
64
|
+
out: Linear;
|
|
65
|
+
constructor(args: ConformerArgs);
|
|
66
|
+
forward(x: MxArray, lengths: MxArray): [MxArray, MxArray];
|
|
67
|
+
loadWeights(weights: WeightMap, prefix: string): void;
|
|
68
|
+
}
|
|
69
|
+
export declare class Conformer extends Module {
|
|
70
|
+
readonly args: ConformerArgs;
|
|
71
|
+
posEnc: RelPositionalEncoding | LocalRelPositionalEncoding | null;
|
|
72
|
+
preEncode: DwStridingSubsampling | Linear;
|
|
73
|
+
layers: ConformerBlock[];
|
|
74
|
+
constructor(args: ConformerArgs);
|
|
75
|
+
setAttentionModel(name: AttentionModel, contextSize?: [number, number]): void;
|
|
76
|
+
forward(x: MxArray, // [batch, seq, mel]
|
|
77
|
+
lengths: MxArray | null, cache: Array<ConformerCache | null> | null): [MxArray, MxArray];
|
|
78
|
+
loadWeights(weights: WeightMap, prefix: string): void;
|
|
79
|
+
}
|
|
80
|
+
export {};
|
|
81
|
+
//# sourceMappingURL=conformer.d.ts.map
|
|
@@ -0,0 +1 @@
|
|
|
1
|
+
{"version":3,"file":"conformer.d.ts","sourceRoot":"","sources":["../../src/mlx/conformer.ts"],"names":[],"mappings":"AAAA,OAAO,EAAE,OAAO,EAAE,MAAM,gBAAgB,CAAC;AACzC,OAAO,EACL,MAAM,EACN,SAAS,EACT,MAAM,EACN,SAAS,EACT,SAAS,EACT,MAAM,EACN,MAAM,EAIP,MAAM,SAAS,CAAC;AACjB,OAAO,EACL,kBAAkB,EAClB,6BAA6B,EAC7B,kCAAkC,EAClC,qBAAqB,EACrB,0BAA0B,EAC3B,MAAM,gBAAgB,CAAC;AACxB,OAAO,EAAE,cAAc,EAAE,MAAM,YAAY,CAAC;AAU5C,MAAM,WAAW,aAAa;IAC5B,MAAM,EAAE,MAAM,CAAC;IACf,OAAO,EAAE,MAAM,CAAC;IAChB,MAAM,EAAE,MAAM,CAAC;IACf,MAAM,EAAE,MAAM,CAAC;IACf,iBAAiB,EAAE,MAAM,CAAC;IAC1B,iBAAiB,EAAE,MAAM,CAAC;IAC1B,kBAAkB,EAAE,MAAM,CAAC;IAC3B,WAAW,EAAE,MAAM,CAAC;IACpB,cAAc,EAAE,MAAM,CAAC;IACvB,uBAAuB,EAAE,MAAM,CAAC;IAChC,YAAY,EAAE,MAAM,CAAC;IACrB,kBAAkB,CAAC,EAAE,OAAO,CAAC;IAC7B,OAAO,CAAC,EAAE,OAAO,CAAC;IAClB,QAAQ,CAAC,EAAE,OAAO,CAAC;IACnB,6BAA6B,CAAC,EAAE,MAAM,CAAC;IACvC,cAAc,CAAC,EAAE,CAAC,MAAM,EAAE,MAAM,CAAC,GAAG,IAAI,CAAC;CAC1C;AAMD,cAAM,WAAY,SAAQ,MAAM;IAC9B,OAAO,EAAE,MAAM,CAAC;IAChB,OAAO,EAAE,MAAM,CAAC;gBAEJ,MAAM,EAAE,MAAM,EAAE,GAAG,EAAE,MAAM,EAAE,OAAO,EAAE,OAAO;IAMzD,OAAO,CAAC,CAAC,EAAE,OAAO,GAAG,OAAO;IAI5B,WAAW,CAAC,OAAO,EAAE,SAAS,EAAE,MAAM,EAAE,MAAM,GAAG,IAAI;CAItD;AAMD,cAAM,WAAY,SAAQ,MAAM;IAC9B,QAAQ,CAAC,OAAO,EAAE,MAAM,CAAC;IACzB,cAAc,EAAE,MAAM,CAAC;IACvB,aAAa,EAAE,MAAM,CAAC;IACtB,SAAS,EAAE,SAAS,CAAC;IACrB,cAAc,EAAE,MAAM,CAAC;gBAEX,IAAI,EAAE,aAAa;IAW/B,OAAO,CAAC,CAAC,EAAE,OAAO,EAAE,KAAK,EAAE,cAAc,GAAG,IAAI,GAAG,OAAO;IAiB1D,WAAW,CAAC,OAAO,EAAE,SAAS,EAAE,MAAM,EAAE,MAAM,GAAG,IAAI;CAMtD;AAMD,KAAK,cAAc,GAAG,SAAS,GAAG,oBAAoB,GAAG,QAAQ,CAAC;AAElE,qBAAa,cAAe,SAAQ,MAAM;IACxC,OAAO,EAAE,SAAS,CAAC;IACnB,GAAG,EAAE,WAAW,CAAC;IAEjB,WAAW,EAAE,SAAS,CAAC;IACvB,QAAQ,EAAE,kBAAkB,GAAG,6BAA6B,GAAG,kCAAkC,CAAC;IAElG,QAAQ,EAAE,SAAS,CAAC;IACpB,IAAI,EAAE,WAAW,CAAC;IAElB,OAAO,EAAE,SAAS,CAAC;IACnB,GAAG,EAAE,WAAW,CAAC;IAEjB,OAAO,EAAE,SAAS,CAAC;IAEnB,OAAO,CAAC,QAAQ,CAAC,IAAI,CAAgB;gBAEzB,IAAI,EAAE,aAAa;IAqB/B,OAAO,CAAC,cAAc;IAmBtB,iBAAiB,CAAC,IAAI,EAAE,cAAc,EAAE,WAAW,GAAE,CAAC,MAAM,EAAE,MAAM,CAAc,GAAG,IAAI;IAOzF,OAAO,CACL,CAAC,EAAE,OAAO,EACV,MAAM,EAAE,OAAO,GAAG,IAAI,EACtB,IAAI,EAAE,OAAO,GAAG,IAAI,EACpB,KAAK,EAAE,cAAc,GAAG,IAAI,GAC3B,OAAO;IAiBV,WAAW,CAAC,OAAO,EAAE,SAAS,EAAE,MAAM,EAAE,MAAM,GAAG,IAAI;CAWtD;AAMD,cAAM,qBAAsB,SAAQ,MAAM;IACxC,OAAO,CAAC,QAAQ,CAAC,WAAW,CAAS;IACrC,OAAO,CAAC,QAAQ,CAAC,MAAM,CAAK;IAC5B,OAAO,CAAC,QAAQ,CAAC,UAAU,CAAK;IAChC,OAAO,CAAC,QAAQ,CAAC,OAAO,CAAS;IAEjC,UAAU,EAAE,KAAK,CAAC,MAAM,GAAG,IAAI,CAAC,CAAC;IACjC,GAAG,EAAE,MAAM,CAAC;gBAEA,IAAI,EAAE,aAAa;IAkC/B,OAAO,CAAC,CAAC,EAAE,OAAO,EAAE,OAAO,EAAE,OAAO,GAAG,CAAC,OAAO,EAAE,OAAO,CAAC;IAiDzD,WAAW,CAAC,OAAO,EAAE,SAAS,EAAE,MAAM,EAAE,MAAM,GAAG,IAAI;CAUtD;AAMD,qBAAa,SAAU,SAAQ,MAAM;IACnC,QAAQ,CAAC,IAAI,EAAE,aAAa,CAAC;IAC7B,MAAM,EAAE,qBAAqB,GAAG,0BAA0B,GAAG,IAAI,CAAC;IAClE,SAAS,EAAE,qBAAqB,GAAG,MAAM,CAAC;IAC1C,MAAM,EAAE,cAAc,EAAE,CAAC;gBAEb,IAAI,EAAE,aAAa;IAiC/B,iBAAiB,CACf,IAAI,EAAE,cAAc,EACpB,WAAW,GAAE,CAAC,MAAM,EAAE,MAAM,CAAc,GACzC,IAAI;IAmBP,OAAO,CACL,CAAC,EAAE,OAAO,EAAE,oBAAoB;IAChC,OAAO,EAAE,OAAO,GAAG,IAAI,EACvB,KAAK,EAAE,KAAK,CAAC,cAAc,GAAG,IAAI,CAAC,GAAG,IAAI,GACzC,CAAC,OAAO,EAAE,OAAO,CAAC;IAkCrB,WAAW,CAAC,OAAO,EAAE,SAAS,EAAE,MAAM,EAAE,MAAM,GAAG,IAAI;CAatD"}
|