@tryhamster/gerbil 1.11.4 → 1.13.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/README.md +104 -0
- package/dist/{architectures-DmZMEFsA.mjs → architectures-DHwj9AQD.mjs} +50 -38
- package/dist/architectures-DHwj9AQD.mjs.map +1 -0
- package/dist/browser/index.d.ts.map +1 -1
- package/dist/browser/index.js +11 -0
- package/dist/browser/index.js.map +1 -1
- package/dist/cli.mjs +8 -8
- package/dist/cli.mjs.map +1 -1
- package/dist/{defaults-DfGx4d1m.mjs → defaults-B0aQZJTM.mjs} +3 -2
- package/dist/defaults-B0aQZJTM.mjs.map +1 -0
- package/dist/frameworks/express.mjs +1 -1
- package/dist/frameworks/fastify.mjs +1 -1
- package/dist/frameworks/hono.mjs +1 -1
- package/dist/frameworks/next.d.mts +2 -2
- package/dist/frameworks/next.mjs +1 -1
- package/dist/frameworks/trpc.mjs +1 -1
- package/dist/gerbil-Bw6do78d.mjs +4 -0
- package/dist/{gerbil-5_80K0gA.d.mts → gerbil-CD_skiL3.d.mts} +2 -2
- package/dist/{gerbil-5_80K0gA.d.mts.map → gerbil-CD_skiL3.d.mts.map} +1 -1
- package/dist/{gerbil-CYVmU8sQ.mjs → gerbil-DAz_a4Jh.mjs} +33 -13
- package/dist/gerbil-DAz_a4Jh.mjs.map +1 -0
- package/dist/gpu/hooks.d.mts +1 -1
- package/dist/gpu/hooks.mjs +1 -1
- package/dist/gpu/index.d.mts +2 -2
- package/dist/gpu/index.mjs +5 -5
- package/dist/{gpu-CrzjQHv2.mjs → gpu-CU_Mldk0.mjs} +1959 -198
- package/dist/gpu-CU_Mldk0.mjs.map +1 -0
- package/dist/index-B3tjyDJI.d.mts.map +1 -1
- package/dist/{index-h8TDu1qm.d.mts → index-Fj2XkP-o.d.mts} +1946 -938
- package/dist/index-Fj2XkP-o.d.mts.map +1 -0
- package/dist/index.d.mts +3 -3
- package/dist/index.d.mts.map +1 -1
- package/dist/index.mjs +6 -6
- package/dist/index.mjs.map +1 -1
- package/dist/integrations/ai-sdk.mjs +1 -1
- package/dist/integrations/langchain.mjs +1 -1
- package/dist/integrations/llamaindex.mjs +1 -1
- package/dist/integrations/mcp.d.mts +2 -2
- package/dist/integrations/mcp.mjs +4 -4
- package/dist/{mcp-CAsD7eCj.mjs → mcp-CdMvLQ_9.mjs} +3 -3
- package/dist/{mcp-CAsD7eCj.mjs.map → mcp-CdMvLQ_9.mjs.map} +1 -1
- package/dist/{moonshine-stt-BXoZaHJE.mjs → moonshine-stt-COJeK2Zb.mjs} +6348 -2000
- package/dist/moonshine-stt-COJeK2Zb.mjs.map +1 -0
- package/dist/moonshine-stt-Dp4j5JEB.mjs +4 -0
- package/dist/{one-liner-ppgw4jHH.mjs → one-liner-CMiGNWJ7.mjs} +2 -2
- package/dist/{one-liner-ppgw4jHH.mjs.map → one-liner-CMiGNWJ7.mjs.map} +1 -1
- package/dist/repl-BhaLCJFb.mjs +9 -0
- package/dist/skills/index.d.mts +4 -4
- package/dist/skills/index.mjs +3 -3
- package/dist/{skills-CQa1Gshd.mjs → skills-BXybFWlG.mjs} +2 -2
- package/dist/{skills-CQa1Gshd.mjs.map → skills-BXybFWlG.mjs.map} +1 -1
- package/dist/tune/index.mjs +1 -1
- package/package.json +1 -1
- package/dist/architectures-DmZMEFsA.mjs.map +0 -1
- package/dist/defaults-DfGx4d1m.mjs.map +0 -1
- package/dist/gerbil-CTefAwKp.mjs +0 -4
- package/dist/gerbil-CYVmU8sQ.mjs.map +0 -1
- package/dist/gpu-CrzjQHv2.mjs.map +0 -1
- package/dist/index-h8TDu1qm.d.mts.map +0 -1
- package/dist/moonshine-stt-BXoZaHJE.mjs.map +0 -1
- package/dist/moonshine-stt-DZVnKgPO.mjs +0 -4
- package/dist/repl-C0Ew7_Z-.mjs +0 -9
|
@@ -1,6 +1,6 @@
|
|
|
1
|
-
import { a as resolveDefaultRepo, i as isTTSRepo, n as OUTETTS_ASSETS, r as OUTETTS_PRESET_VOICES, t as DEFAULT_MODELS } from "./defaults-
|
|
2
|
-
import { A as kaniSinTensor, C as computeKaniPositions, D as kaniAttentionLayerIndices, E as generateNanoCodecDecoderGraph, F as CANONICAL_KEYS, I as DTYPE_BYTES, L as GEMMA4_VIS_KEYS, M as parseKaniConfig, N as DEFAULT_GROUP_SIZE, O as kaniCosTensor, S as buildKaniLayerCosSin, T as generateKaniTtsGraph, a as PARLER_DAC_LATENT_DIM, b as KANI_START_OF_HUMAN, c as PARLER_SAMPLE_RATE, d as generateParlerEncoderGraph, f as parseParlerConfig, i as PARLER_DAC_DECODER_DIM, k as kaniLayerAlpha, l as buildT5RelativeBias, o as PARLER_DECODER_RATES, p as revertDelayPattern, r as PARLER_BOS_TOKEN_ID, s as PARLER_EOS_TOKEN_ID, u as generateParlerDecoderGraph, x as audioTokensToCodes, y as KANI_END_OF_HUMAN } from "./architectures-
|
|
3
|
-
import { C as
|
|
1
|
+
import { a as resolveDefaultRepo, i as isTTSRepo, n as OUTETTS_ASSETS, r as OUTETTS_PRESET_VOICES, t as DEFAULT_MODELS } from "./defaults-B0aQZJTM.mjs";
|
|
2
|
+
import { A as kaniSinTensor, C as computeKaniPositions, D as kaniAttentionLayerIndices, E as generateNanoCodecDecoderGraph, F as CANONICAL_KEYS, I as DTYPE_BYTES, L as GEMMA4_VIS_KEYS, M as parseKaniConfig, N as DEFAULT_GROUP_SIZE, O as kaniCosTensor, S as buildKaniLayerCosSin, T as generateKaniTtsGraph, a as PARLER_DAC_LATENT_DIM, b as KANI_START_OF_HUMAN, c as PARLER_SAMPLE_RATE, d as generateParlerEncoderGraph, f as parseParlerConfig, i as PARLER_DAC_DECODER_DIM, k as kaniLayerAlpha, l as buildT5RelativeBias, o as PARLER_DECODER_RATES, p as revertDelayPattern, r as PARLER_BOS_TOKEN_ID, s as PARLER_EOS_TOKEN_ID, u as generateParlerDecoderGraph, x as audioTokensToCodes, y as KANI_END_OF_HUMAN } from "./architectures-DHwj9AQD.mjs";
|
|
3
|
+
import { C as createStorageBuffer, D as initGPU, E as getOrCreatePipeline, O as verifyGPU, S as createBindGroup, T as destroyBuffers, _ as fetchAdapter, a as loadKaniTTS, b as MATMUL_BIAS_F16C_SPEC, c as loadOuteTTS, d as quantizeKaniBackbone, f as remapPrunedToken, g as buildLoRADeltas, i as loadGepardTTS, l as loadParlerTTS, m as Executor, o as loadModel, r as createKeyMapperForArch, u as quantizeBackboneInt4, w as createUniformBuffer, x as clearPipelineCache, y as KERNEL_REGISTRY } from "./moonshine-stt-COJeK2Zb.mjs";
|
|
4
4
|
|
|
5
5
|
//#region src/gpu/architectures/gemma4_vision.ts
|
|
6
6
|
/**
|
|
@@ -767,7 +767,7 @@ const KANI_CODEC_LOOKBACK_FRAMES = 32;
|
|
|
767
767
|
/** Hard cap so a stuck decode cannot loop forever (reference uses 3000). */
|
|
768
768
|
const DEFAULT_MAX_NEW_TOKENS$1 = 3e3;
|
|
769
769
|
/** Build a fresh map of just the constant weights a graph references (see moonshine-stt). */
|
|
770
|
-
function selectGraphWeights$
|
|
770
|
+
function selectGraphWeights$3(graph, weights) {
|
|
771
771
|
const out = /* @__PURE__ */ new Map();
|
|
772
772
|
for (const [name, desc] of Object.entries(graph.tensors)) {
|
|
773
773
|
if (desc.storage !== "constant") continue;
|
|
@@ -891,7 +891,7 @@ var KaniTTS = class KaniTTS {
|
|
|
891
891
|
maxSeqLen,
|
|
892
892
|
kvMode: "f32"
|
|
893
893
|
});
|
|
894
|
-
exec.uploadWeightsMap(selectGraphWeights$
|
|
894
|
+
exec.uploadWeightsMap(selectGraphWeights$3(graph, loaded.backboneWeights));
|
|
895
895
|
exec.initBindGroups();
|
|
896
896
|
return exec;
|
|
897
897
|
}
|
|
@@ -1105,7 +1105,7 @@ var KaniTTS = class KaniTTS {
|
|
|
1105
1105
|
kvMode: "f32"
|
|
1106
1106
|
});
|
|
1107
1107
|
try {
|
|
1108
|
-
exec.uploadWeightsMap(selectGraphWeights$
|
|
1108
|
+
exec.uploadWeightsMap(selectGraphWeights$3(graph, this.loaded.codecWeights));
|
|
1109
1109
|
exec.initBindGroups();
|
|
1110
1110
|
exec.reset();
|
|
1111
1111
|
exec.writeInput("audio_codes", winCodes);
|
|
@@ -1164,66 +1164,1682 @@ function stripRestatement(typed, suggestion, minWords = DEFAULT_RESTATEMENT_MIN_
|
|
|
1164
1164
|
if (normTyped.includes(phrase)) matched = n;
|
|
1165
1165
|
else break;
|
|
1166
1166
|
}
|
|
1167
|
-
if (matched >= minWords) return words.slice(matched).join(" ");
|
|
1168
|
-
return suggestion;
|
|
1169
|
-
}
|
|
1170
|
-
/**
|
|
1171
|
-
* Truncate `text` right before the first point where a span of `minRun` or more
|
|
1172
|
-
* consecutive words repeats a span seen earlier — the classic small-model loop
|
|
1173
|
-
* ("…endless sands …endless sands…"). Returns `text` unchanged (spacing
|
|
1174
|
-
* preserved) when no such repeat exists.
|
|
1175
|
-
*/
|
|
1176
|
-
function truncateInternalRepeat(text, minRun = DEFAULT_INTERNAL_REPEAT_MIN_RUN) {
|
|
1177
|
-
const words = text.split(WHITESPACE_RUN).filter(Boolean);
|
|
1178
|
-
const seen = /* @__PURE__ */ new Set();
|
|
1179
|
-
for (let j = 0; j + minRun <= words.length; j++) {
|
|
1180
|
-
const gram = words.slice(j, j + minRun).join(" ").toLowerCase();
|
|
1181
|
-
if (seen.has(gram)) return words.slice(0, j).join(" ");
|
|
1182
|
-
seen.add(gram);
|
|
1167
|
+
if (matched >= minWords) return words.slice(matched).join(" ");
|
|
1168
|
+
return suggestion;
|
|
1169
|
+
}
|
|
1170
|
+
/**
|
|
1171
|
+
* Truncate `text` right before the first point where a span of `minRun` or more
|
|
1172
|
+
* consecutive words repeats a span seen earlier — the classic small-model loop
|
|
1173
|
+
* ("…endless sands …endless sands…"). Returns `text` unchanged (spacing
|
|
1174
|
+
* preserved) when no such repeat exists.
|
|
1175
|
+
*/
|
|
1176
|
+
function truncateInternalRepeat(text, minRun = DEFAULT_INTERNAL_REPEAT_MIN_RUN) {
|
|
1177
|
+
const words = text.split(WHITESPACE_RUN).filter(Boolean);
|
|
1178
|
+
const seen = /* @__PURE__ */ new Set();
|
|
1179
|
+
for (let j = 0; j + minRun <= words.length; j++) {
|
|
1180
|
+
const gram = words.slice(j, j + minRun).join(" ").toLowerCase();
|
|
1181
|
+
if (seen.has(gram)) return words.slice(0, j).join(" ");
|
|
1182
|
+
seen.add(gram);
|
|
1183
|
+
}
|
|
1184
|
+
return text;
|
|
1185
|
+
}
|
|
1186
|
+
/**
|
|
1187
|
+
* Cap `text` to at most `max` characters, cutting on a word boundary when a
|
|
1188
|
+
* reasonable one exists in the back half of the window.
|
|
1189
|
+
*/
|
|
1190
|
+
function capLength(text, max) {
|
|
1191
|
+
if (text.length <= max) return text;
|
|
1192
|
+
const cut = text.slice(0, max);
|
|
1193
|
+
const lastSpace = cut.lastIndexOf(" ");
|
|
1194
|
+
return (lastSpace > max * .5 ? cut.slice(0, lastSpace) : cut).trimEnd();
|
|
1195
|
+
}
|
|
1196
|
+
/**
|
|
1197
|
+
* Turn a raw model completion into a clean inline ghost continuation:
|
|
1198
|
+
* 1. keep only the first line (when `singleLine`),
|
|
1199
|
+
* 2. strip wrapping quotes,
|
|
1200
|
+
* 3. drop a full verbatim echo of the typed text,
|
|
1201
|
+
* 4. drop a leading character overlap with the typed tail,
|
|
1202
|
+
* 5. drop a leading word-run that restates an earlier typed phrase,
|
|
1203
|
+
* 6. truncate an internal phrase loop,
|
|
1204
|
+
* 7. cap the length,
|
|
1205
|
+
* 8. add a single smart leading space so the ghost joins the caret naturally.
|
|
1206
|
+
*
|
|
1207
|
+
* Returns "" when nothing novel is left — the ghost then simply shows nothing,
|
|
1208
|
+
* which is the correct behavior for a suggestion that only repeats the input.
|
|
1209
|
+
*/
|
|
1210
|
+
function cleanSuggestion(raw, typed, options = {}) {
|
|
1211
|
+
const { singleLine = true, maxChars = DEFAULT_MAX_SUGGESTION_CHARS } = options;
|
|
1212
|
+
let s = singleLine ? raw.replace(AFTER_FIRST_NEWLINE, "") : raw;
|
|
1213
|
+
s = s.replace(WRAPPING_QUOTES_START, "").replace(WRAPPING_QUOTES_END, "");
|
|
1214
|
+
s = s.trim();
|
|
1215
|
+
if (!s) return "";
|
|
1216
|
+
const typedTrim = typed.trim();
|
|
1217
|
+
if (typedTrim && s.startsWith(typedTrim)) s = s.slice(typedTrim.length).trimStart();
|
|
1218
|
+
s = stripLeadingOverlap(typed, s).trimStart();
|
|
1219
|
+
s = stripRestatement(typed, s);
|
|
1220
|
+
s = truncateInternalRepeat(s);
|
|
1221
|
+
s = capLength(s.trim(), maxChars).trim();
|
|
1222
|
+
if (!s) return "";
|
|
1223
|
+
const startsWithPunct = LEADING_PUNCT.test(s);
|
|
1224
|
+
const typedEndsWithSpace = ENDS_WITH_SPACE.test(typed) || typed.length === 0;
|
|
1225
|
+
return startsWithPunct || typedEndsWithSpace ? s : ` ${s}`;
|
|
1226
|
+
}
|
|
1227
|
+
|
|
1228
|
+
//#endregion
|
|
1229
|
+
//#region src/gpu/sampler.ts
|
|
1230
|
+
/** mulberry32 — tiny deterministic PRNG for seeded sampling. */
|
|
1231
|
+
function mulberry32(seed) {
|
|
1232
|
+
let a = seed >>> 0;
|
|
1233
|
+
return () => {
|
|
1234
|
+
a = a + 1831565813 | 0;
|
|
1235
|
+
let t = Math.imul(a ^ a >>> 15, 1 | a);
|
|
1236
|
+
t = t + Math.imul(t ^ t >>> 7, 61 | t) ^ t;
|
|
1237
|
+
return ((t ^ t >>> 14) >>> 0) / 4294967296;
|
|
1238
|
+
};
|
|
1239
|
+
}
|
|
1240
|
+
/** Per-params seeded RNG streams (one stream per options object identity). */
|
|
1241
|
+
const seededStreams = /* @__PURE__ */ new WeakMap();
|
|
1242
|
+
function nextRandom(params) {
|
|
1243
|
+
if (params.seed === void 0) return Math.random();
|
|
1244
|
+
let rng = seededStreams.get(params);
|
|
1245
|
+
if (!rng) {
|
|
1246
|
+
rng = mulberry32(params.seed);
|
|
1247
|
+
seededStreams.set(params, rng);
|
|
1248
|
+
}
|
|
1249
|
+
return rng();
|
|
1250
|
+
}
|
|
1251
|
+
let _heapIndices = null;
|
|
1252
|
+
let _heapValues = null;
|
|
1253
|
+
/**
|
|
1254
|
+
* Sample a token ID from logits.
|
|
1255
|
+
*
|
|
1256
|
+
* Pipeline: repetition penalty → temperature → top-k (min-heap) → softmax → top-p → sample.
|
|
1257
|
+
*/
|
|
1258
|
+
function sampleToken(logits, params = {}, previousTokens) {
|
|
1259
|
+
const temperature = params.temperature ?? .7;
|
|
1260
|
+
const topK = params.topK ?? 50;
|
|
1261
|
+
const topP = params.topP ?? .9;
|
|
1262
|
+
const repetitionPenalty = params.repetitionPenalty ?? 1;
|
|
1263
|
+
if (temperature < 1e-6) return argmax$2(logits);
|
|
1264
|
+
const N = logits.length;
|
|
1265
|
+
const K = Math.min(topK > 0 ? topK : N, N);
|
|
1266
|
+
if (!_heapIndices || _heapIndices.length < K) {
|
|
1267
|
+
_heapIndices = new Uint32Array(K);
|
|
1268
|
+
_heapValues = new Float32Array(K);
|
|
1269
|
+
}
|
|
1270
|
+
const hIdx = _heapIndices;
|
|
1271
|
+
const hVal = _heapValues;
|
|
1272
|
+
let penaltySet = null;
|
|
1273
|
+
if (repetitionPenalty !== 1 && previousTokens?.length) penaltySet = new Set(previousTokens);
|
|
1274
|
+
let heapSize = 0;
|
|
1275
|
+
for (let i = 0; i < N; i++) {
|
|
1276
|
+
let s = logits[i];
|
|
1277
|
+
if (penaltySet?.has(i)) s = s > 0 ? s / repetitionPenalty : s * repetitionPenalty;
|
|
1278
|
+
s /= temperature;
|
|
1279
|
+
if (heapSize < K) {
|
|
1280
|
+
hIdx[heapSize] = i;
|
|
1281
|
+
hVal[heapSize] = s;
|
|
1282
|
+
heapSize++;
|
|
1283
|
+
if (heapSize === K) for (let j = (K >> 1) - 1; j >= 0; j--) siftDown(hIdx, hVal, j, K);
|
|
1284
|
+
} else if (s > hVal[0]) {
|
|
1285
|
+
hIdx[0] = i;
|
|
1286
|
+
hVal[0] = s;
|
|
1287
|
+
siftDown(hIdx, hVal, 0, K);
|
|
1288
|
+
}
|
|
1289
|
+
}
|
|
1290
|
+
for (let i = 1; i < heapSize; i++) {
|
|
1291
|
+
const vi = hVal[i];
|
|
1292
|
+
const ii = hIdx[i];
|
|
1293
|
+
let j = i - 1;
|
|
1294
|
+
while (j >= 0 && hVal[j] < vi) {
|
|
1295
|
+
hVal[j + 1] = hVal[j];
|
|
1296
|
+
hIdx[j + 1] = hIdx[j];
|
|
1297
|
+
j--;
|
|
1298
|
+
}
|
|
1299
|
+
hVal[j + 1] = vi;
|
|
1300
|
+
hIdx[j + 1] = ii;
|
|
1301
|
+
}
|
|
1302
|
+
const maxScore = hVal[0];
|
|
1303
|
+
let sumExp = 0;
|
|
1304
|
+
for (let i = 0; i < heapSize; i++) {
|
|
1305
|
+
const p = Math.exp(hVal[i] - maxScore);
|
|
1306
|
+
hVal[i] = p;
|
|
1307
|
+
sumExp += p;
|
|
1308
|
+
}
|
|
1309
|
+
const invSum = 1 / sumExp;
|
|
1310
|
+
for (let i = 0; i < heapSize; i++) hVal[i] *= invSum;
|
|
1311
|
+
let candidateCount = heapSize;
|
|
1312
|
+
if (topP < 1) {
|
|
1313
|
+
let cumulative$1 = 0;
|
|
1314
|
+
for (let i = 0; i < heapSize; i++) {
|
|
1315
|
+
cumulative$1 += hVal[i];
|
|
1316
|
+
if (cumulative$1 >= topP) {
|
|
1317
|
+
candidateCount = i + 1;
|
|
1318
|
+
break;
|
|
1319
|
+
}
|
|
1320
|
+
}
|
|
1321
|
+
let sum = 0;
|
|
1322
|
+
for (let i = 0; i < candidateCount; i++) sum += hVal[i];
|
|
1323
|
+
const inv = 1 / sum;
|
|
1324
|
+
for (let i = 0; i < candidateCount; i++) hVal[i] *= inv;
|
|
1325
|
+
}
|
|
1326
|
+
const r = nextRandom(params);
|
|
1327
|
+
let cumulative = 0;
|
|
1328
|
+
for (let i = 0; i < candidateCount; i++) {
|
|
1329
|
+
cumulative += hVal[i];
|
|
1330
|
+
if (r <= cumulative) return hIdx[i];
|
|
1331
|
+
}
|
|
1332
|
+
return hIdx[candidateCount - 1];
|
|
1333
|
+
}
|
|
1334
|
+
/**
|
|
1335
|
+
* Return the index of the maximum value (greedy decoding).
|
|
1336
|
+
*/
|
|
1337
|
+
function argmax$2(arr) {
|
|
1338
|
+
let maxIdx = 0;
|
|
1339
|
+
let maxVal = arr[0];
|
|
1340
|
+
for (let i = 1; i < arr.length; i++) if (arr[i] > maxVal) {
|
|
1341
|
+
maxVal = arr[i];
|
|
1342
|
+
maxIdx = i;
|
|
1343
|
+
}
|
|
1344
|
+
return maxIdx;
|
|
1345
|
+
}
|
|
1346
|
+
/** Min-heap sift down on parallel index/value typed arrays. */
|
|
1347
|
+
function siftDown(indices, values, i, n) {
|
|
1348
|
+
while (true) {
|
|
1349
|
+
let smallest = i;
|
|
1350
|
+
const left = 2 * i + 1;
|
|
1351
|
+
const right = 2 * i + 2;
|
|
1352
|
+
if (left < n && values[left] < values[smallest]) smallest = left;
|
|
1353
|
+
if (right < n && values[right] < values[smallest]) smallest = right;
|
|
1354
|
+
if (smallest === i) break;
|
|
1355
|
+
const ti = indices[i];
|
|
1356
|
+
indices[i] = indices[smallest];
|
|
1357
|
+
indices[smallest] = ti;
|
|
1358
|
+
const tv = values[i];
|
|
1359
|
+
values[i] = values[smallest];
|
|
1360
|
+
values[smallest] = tv;
|
|
1361
|
+
i = smallest;
|
|
1362
|
+
}
|
|
1363
|
+
}
|
|
1364
|
+
|
|
1365
|
+
//#endregion
|
|
1366
|
+
//#region src/gpu/batch/scheduler.ts
|
|
1367
|
+
const PIPELINE_DEPTH = 2;
|
|
1368
|
+
var RequestScheduler = class {
|
|
1369
|
+
host;
|
|
1370
|
+
batchSize;
|
|
1371
|
+
maxSeqLen;
|
|
1372
|
+
blockSize;
|
|
1373
|
+
prefillChunk;
|
|
1374
|
+
/** Free KV blocks kept unreserved as safety headroom. */
|
|
1375
|
+
blockHeadroom;
|
|
1376
|
+
nextId = 1;
|
|
1377
|
+
queue = [];
|
|
1378
|
+
/** slot → occupying request (prefilling or decoding). */
|
|
1379
|
+
slots;
|
|
1380
|
+
/** Sum of blocksReserved across live (admitted, unfinished) requests. */
|
|
1381
|
+
reservedBlocks = 0;
|
|
1382
|
+
prefilling = null;
|
|
1383
|
+
inFlight = [];
|
|
1384
|
+
/**
|
|
1385
|
+
* The slot→adapter-id assignment last uploaded to the executor (null =
|
|
1386
|
+
* nothing uploaded this batch session). Steps only re-upload when the
|
|
1387
|
+
* assignment changes (admissions/retirements), and the very first upload is
|
|
1388
|
+
* skipped entirely while every lane is bare base — an adapter-free serving
|
|
1389
|
+
* session never touches the LoRA machinery.
|
|
1390
|
+
*/
|
|
1391
|
+
lastLaneAdapterIds = null;
|
|
1392
|
+
stepCounter = 0;
|
|
1393
|
+
loop = null;
|
|
1394
|
+
closed = false;
|
|
1395
|
+
stats = {
|
|
1396
|
+
submitted: 0,
|
|
1397
|
+
completed: 0,
|
|
1398
|
+
stepsSubmitted: 0,
|
|
1399
|
+
prefillChunks: 0,
|
|
1400
|
+
tokensGenerated: 0,
|
|
1401
|
+
busyMs: 0,
|
|
1402
|
+
peakActiveLanes: 0
|
|
1403
|
+
};
|
|
1404
|
+
constructor(host, options = {}) {
|
|
1405
|
+
this.host = host;
|
|
1406
|
+
this.batchSize = host.executor.batchSize;
|
|
1407
|
+
if (this.batchSize <= 0) throw new Error("RequestScheduler requires GERBIL_BATCH=N (N >= 2) on the Dawn path");
|
|
1408
|
+
this.maxSeqLen = host.executor.maxSequenceLength;
|
|
1409
|
+
this.blockSize = host.executor.kvBlockSizeTokens;
|
|
1410
|
+
this.slots = new Array(this.batchSize).fill(null);
|
|
1411
|
+
const envChunk = typeof process !== "undefined" ? Number(process.env?.GERBIL_PREFILL_CHUNK) : NaN;
|
|
1412
|
+
this.prefillChunk = options.prefillChunkTokens ?? (Number.isFinite(envChunk) && envChunk >= 8 ? envChunk : 128);
|
|
1413
|
+
this.blockHeadroom = 2;
|
|
1414
|
+
}
|
|
1415
|
+
/**
|
|
1416
|
+
* Submit a request; resolves with the full result when it finishes.
|
|
1417
|
+
* Tokens stream through options.onToken as they are decoded.
|
|
1418
|
+
*/
|
|
1419
|
+
submit(prompt, options = {}) {
|
|
1420
|
+
if (this.closed) return Promise.reject(/* @__PURE__ */ new Error("RequestScheduler is closed"));
|
|
1421
|
+
if ((options.sampling?.temperature ?? 0) >= 1e-6) return Promise.reject(/* @__PURE__ */ new Error("RequestScheduler: Phase 5 is greedy-only — pass sampling: { temperature: 0 }"));
|
|
1422
|
+
let adapterReg = null;
|
|
1423
|
+
if (options.adapter != null) {
|
|
1424
|
+
if (!this.host.executor.batchLoraEnabled) return Promise.reject(/* @__PURE__ */ new Error("RequestScheduler: per-request adapters require GERBIL_BATCH_LORA=1 (with GERBIL_BATCH=N)"));
|
|
1425
|
+
try {
|
|
1426
|
+
adapterReg = this.host.resolveAdapter(options.adapter);
|
|
1427
|
+
} catch (err) {
|
|
1428
|
+
return Promise.reject(err instanceof Error ? err : new Error(String(err)));
|
|
1429
|
+
}
|
|
1430
|
+
}
|
|
1431
|
+
const messages = typeof prompt === "string" ? [...options.systemPrompt ? [{
|
|
1432
|
+
role: "system",
|
|
1433
|
+
content: options.systemPrompt
|
|
1434
|
+
}] : [], {
|
|
1435
|
+
role: "user",
|
|
1436
|
+
content: prompt
|
|
1437
|
+
}] : options.systemPrompt ? [{
|
|
1438
|
+
role: "system",
|
|
1439
|
+
content: options.systemPrompt
|
|
1440
|
+
}, ...prompt] : prompt;
|
|
1441
|
+
const promptIds = new Uint32Array(this.host.encodeChat(messages));
|
|
1442
|
+
const roomFor = this.maxSeqLen - promptIds.length - PIPELINE_DEPTH - 1;
|
|
1443
|
+
if (roomFor < 1) return Promise.reject(/* @__PURE__ */ new Error(`RequestScheduler: prompt (${promptIds.length} tokens) leaves no room in maxSeqLen ${this.maxSeqLen}`));
|
|
1444
|
+
const maxTokens = Math.min(options.maxTokens ?? 128, roomFor);
|
|
1445
|
+
return new Promise((resolve, reject) => {
|
|
1446
|
+
const req = {
|
|
1447
|
+
id: this.nextId++,
|
|
1448
|
+
promptIds,
|
|
1449
|
+
maxTokens,
|
|
1450
|
+
stopSequences: options.stopSequences ?? [],
|
|
1451
|
+
adapterName: options.adapter ?? null,
|
|
1452
|
+
adapterId: adapterReg?.id ?? -1,
|
|
1453
|
+
adapterDeltas: adapterReg?.deltas ?? null,
|
|
1454
|
+
onToken: options.onToken,
|
|
1455
|
+
resolve,
|
|
1456
|
+
reject,
|
|
1457
|
+
state: "queued",
|
|
1458
|
+
slot: -1,
|
|
1459
|
+
prefillPos: 0,
|
|
1460
|
+
blocksReserved: 0,
|
|
1461
|
+
generatedIds: [],
|
|
1462
|
+
text: "",
|
|
1463
|
+
finishReason: "max_tokens",
|
|
1464
|
+
submittedAt: performance.now(),
|
|
1465
|
+
admittedAt: 0,
|
|
1466
|
+
firstTokenAt: 0
|
|
1467
|
+
};
|
|
1468
|
+
this.queue.push(req);
|
|
1469
|
+
this.stats.submitted++;
|
|
1470
|
+
this.pump();
|
|
1471
|
+
});
|
|
1472
|
+
}
|
|
1473
|
+
/** Resolves when every submitted request has completed. */
|
|
1474
|
+
async drain() {
|
|
1475
|
+
while (this.loop) await this.loop;
|
|
1476
|
+
}
|
|
1477
|
+
/** Stop accepting requests; resolves when in-flight work drains. */
|
|
1478
|
+
async close() {
|
|
1479
|
+
this.closed = true;
|
|
1480
|
+
await this.drain();
|
|
1481
|
+
}
|
|
1482
|
+
getStats() {
|
|
1483
|
+
return { ...this.stats };
|
|
1484
|
+
}
|
|
1485
|
+
pump() {
|
|
1486
|
+
if (this.loop) return;
|
|
1487
|
+
this.loop = this.run().finally(() => {
|
|
1488
|
+
this.loop = null;
|
|
1489
|
+
if (this.hasWork()) this.pump();
|
|
1490
|
+
});
|
|
1491
|
+
}
|
|
1492
|
+
hasWork() {
|
|
1493
|
+
return this.queue.length > 0 || this.prefilling !== null || this.inFlight.length > 0 || this.slots.some((s) => s !== null);
|
|
1494
|
+
}
|
|
1495
|
+
activeMask() {
|
|
1496
|
+
const mask = new Uint8Array(this.batchSize);
|
|
1497
|
+
const rows = [];
|
|
1498
|
+
for (let b = 0; b < this.batchSize; b++) {
|
|
1499
|
+
const req = this.slots[b];
|
|
1500
|
+
if (req && req.state === "decoding") {
|
|
1501
|
+
mask[b] = 1;
|
|
1502
|
+
rows.push({
|
|
1503
|
+
slot: b,
|
|
1504
|
+
req
|
|
1505
|
+
});
|
|
1506
|
+
}
|
|
1507
|
+
}
|
|
1508
|
+
return {
|
|
1509
|
+
mask,
|
|
1510
|
+
rows
|
|
1511
|
+
};
|
|
1512
|
+
}
|
|
1513
|
+
/** The serving loop. Holds the engine generation lock while live. */
|
|
1514
|
+
async run() {
|
|
1515
|
+
const release = await this.host.acquireLock();
|
|
1516
|
+
const ex = this.host.executor;
|
|
1517
|
+
const loopStart = performance.now();
|
|
1518
|
+
ex.resetBatch();
|
|
1519
|
+
this.stepCounter = 0;
|
|
1520
|
+
this.lastLaneAdapterIds = null;
|
|
1521
|
+
try {
|
|
1522
|
+
while (this.hasWork()) {
|
|
1523
|
+
this.admit();
|
|
1524
|
+
const { mask, rows } = this.activeMask();
|
|
1525
|
+
if (rows.length > 0 && this.inFlight.length < PIPELINE_DEPTH) {
|
|
1526
|
+
const readbackSlot = this.stepCounter % PIPELINE_DEPTH;
|
|
1527
|
+
this.syncLaneAdapters();
|
|
1528
|
+
ex.submitBatchDecodeStepMasked(mask, readbackSlot);
|
|
1529
|
+
this.stepCounter++;
|
|
1530
|
+
this.inFlight.push({
|
|
1531
|
+
readbackSlot,
|
|
1532
|
+
rows
|
|
1533
|
+
});
|
|
1534
|
+
this.stats.stepsSubmitted++;
|
|
1535
|
+
if (rows.length > this.stats.peakActiveLanes) this.stats.peakActiveLanes = rows.length;
|
|
1536
|
+
}
|
|
1537
|
+
if (this.prefilling) await this.prefillChunkStep();
|
|
1538
|
+
const pipelineFull = this.inFlight.length >= PIPELINE_DEPTH;
|
|
1539
|
+
const idleLanes = rows.length === 0 && this.prefilling === null;
|
|
1540
|
+
if (this.inFlight.length > 0 && (pipelineFull || idleLanes)) await this.readStep();
|
|
1541
|
+
}
|
|
1542
|
+
} catch (err) {
|
|
1543
|
+
const error = err instanceof Error ? err : new Error(String(err));
|
|
1544
|
+
if (this.prefilling?.adapterDeltas) try {
|
|
1545
|
+
ex.clearRuntimeLoRA();
|
|
1546
|
+
} catch {}
|
|
1547
|
+
for (const req of [...this.queue, ...this.slots.filter((s) => s !== null)]) if (req.state !== "done") {
|
|
1548
|
+
req.state = "done";
|
|
1549
|
+
req.reject(error);
|
|
1550
|
+
}
|
|
1551
|
+
this.queue = [];
|
|
1552
|
+
this.slots.fill(null);
|
|
1553
|
+
this.prefilling = null;
|
|
1554
|
+
this.inFlight = [];
|
|
1555
|
+
this.reservedBlocks = 0;
|
|
1556
|
+
throw error;
|
|
1557
|
+
} finally {
|
|
1558
|
+
this.stats.busyMs += performance.now() - loopStart;
|
|
1559
|
+
release();
|
|
1560
|
+
}
|
|
1561
|
+
}
|
|
1562
|
+
/**
|
|
1563
|
+
* Upload the slot→adapter assignment for the step about to be submitted,
|
|
1564
|
+
* when it differs from the last uploaded one. Called immediately before
|
|
1565
|
+
* every decode-step submit: `queue.writeBuffer` is queue-ordered, so steps
|
|
1566
|
+
* already in flight keep the assignment they were submitted under, and a
|
|
1567
|
+
* retiring lane's slot can never leak its adapter into the slot's next
|
|
1568
|
+
* occupant — the next occupant's assignment is re-derived from `slots`
|
|
1569
|
+
* before its first decode step. Masked (free/prefilling) lanes are pinned
|
|
1570
|
+
* to bare base; their output is discarded regardless.
|
|
1571
|
+
*/
|
|
1572
|
+
syncLaneAdapters() {
|
|
1573
|
+
if (!this.host.executor.batchLoraEnabled) return;
|
|
1574
|
+
const ids = new Int32Array(this.batchSize).fill(-1);
|
|
1575
|
+
let anyAdapter = false;
|
|
1576
|
+
for (let b = 0; b < this.batchSize; b++) {
|
|
1577
|
+
const req = this.slots[b];
|
|
1578
|
+
if (req && req.state === "decoding" && req.adapterId >= 0) {
|
|
1579
|
+
ids[b] = req.adapterId;
|
|
1580
|
+
anyAdapter = true;
|
|
1581
|
+
}
|
|
1582
|
+
}
|
|
1583
|
+
const last = this.lastLaneAdapterIds;
|
|
1584
|
+
if (last === null) {
|
|
1585
|
+
if (!anyAdapter) return;
|
|
1586
|
+
} else if (ids.every((id, b) => id === last[b])) return;
|
|
1587
|
+
this.host.executor.setBatchLaneAdapters(ids);
|
|
1588
|
+
this.lastLaneAdapterIds = ids;
|
|
1589
|
+
}
|
|
1590
|
+
/**
|
|
1591
|
+
* Admission: move the next queued request into a free slot when (a) a slot
|
|
1592
|
+
* is free, (b) no other prefill is staged (the single-sequence prefill path
|
|
1593
|
+
* stages SSM state in singleton buffers — one at a time), and (c) its
|
|
1594
|
+
* worst-case block need fits the unreserved KV pool with headroom.
|
|
1595
|
+
*/
|
|
1596
|
+
admit() {
|
|
1597
|
+
if (this.prefilling || this.queue.length === 0) return;
|
|
1598
|
+
const slot = this.slots.indexOf(null);
|
|
1599
|
+
if (slot === -1) return;
|
|
1600
|
+
const req = this.queue[0];
|
|
1601
|
+
const worstTokens = req.promptIds.length + req.maxTokens + PIPELINE_DEPTH;
|
|
1602
|
+
const blocksNeeded = Math.ceil(worstTokens / this.blockSize);
|
|
1603
|
+
const totalBlocks = this.host.executor.kvTotalBlocks;
|
|
1604
|
+
if (this.reservedBlocks + blocksNeeded + this.blockHeadroom > totalBlocks) return;
|
|
1605
|
+
this.queue.shift();
|
|
1606
|
+
req.state = "prefilling";
|
|
1607
|
+
req.slot = slot;
|
|
1608
|
+
req.blocksReserved = blocksNeeded;
|
|
1609
|
+
req.admittedAt = performance.now();
|
|
1610
|
+
this.reservedBlocks += blocksNeeded;
|
|
1611
|
+
this.slots[slot] = req;
|
|
1612
|
+
this.prefilling = req;
|
|
1613
|
+
if (req.adapterDeltas) {
|
|
1614
|
+
const { applied } = this.host.executor.applyRuntimeLoRA(req.adapterDeltas);
|
|
1615
|
+
if (applied === 0) throw new Error(`RequestScheduler: adapter "${req.adapterName}" resolved to 0 applicable prefill targets`);
|
|
1616
|
+
}
|
|
1617
|
+
this.host.executor.beginBatchSlot(slot);
|
|
1618
|
+
}
|
|
1619
|
+
/**
|
|
1620
|
+
* Run one bounded prefill chunk for the staged request on the
|
|
1621
|
+
* single-sequence path (into its slot's block-table row via tableBase).
|
|
1622
|
+
* Intermediate chunks are submit-only; the final chunk reads logits, adopts
|
|
1623
|
+
* the singleton SSM state into the slot's pool row, seeds the lane's first
|
|
1624
|
+
* token, and flips the lane to decoding.
|
|
1625
|
+
*/
|
|
1626
|
+
async prefillChunkStep() {
|
|
1627
|
+
const req = this.prefilling;
|
|
1628
|
+
if (!req) return;
|
|
1629
|
+
const ex = this.host.executor;
|
|
1630
|
+
const end = Math.min(req.prefillPos + this.prefillChunk, req.promptIds.length);
|
|
1631
|
+
const isLast = end === req.promptIds.length;
|
|
1632
|
+
const chunk = req.promptIds.subarray(req.prefillPos, end);
|
|
1633
|
+
const { logits } = await ex.forward(chunk, { readLogits: isLast });
|
|
1634
|
+
req.prefillPos = end;
|
|
1635
|
+
this.stats.prefillChunks++;
|
|
1636
|
+
if (!isLast) return;
|
|
1637
|
+
ex.finishBatchSlot(req.slot);
|
|
1638
|
+
if (req.adapterDeltas) ex.clearRuntimeLoRA();
|
|
1639
|
+
this.prefilling = null;
|
|
1640
|
+
const first = this.host.remapToken(sampleToken(logits, { temperature: 0 }));
|
|
1641
|
+
ex.injectBatchToken(req.slot, first);
|
|
1642
|
+
req.state = "decoding";
|
|
1643
|
+
req.firstTokenAt = performance.now();
|
|
1644
|
+
this.consumeToken(req, first);
|
|
1645
|
+
}
|
|
1646
|
+
/** Read the oldest in-flight step and deliver its tokens. */
|
|
1647
|
+
async readStep() {
|
|
1648
|
+
const step = this.inFlight.shift();
|
|
1649
|
+
if (!step) return;
|
|
1650
|
+
const tokens = await this.host.executor.readBatchTokens(step.readbackSlot);
|
|
1651
|
+
for (const { slot, req } of step.rows) {
|
|
1652
|
+
if (this.slots[slot] !== req || req.state !== "decoding") continue;
|
|
1653
|
+
this.consumeToken(req, tokens[slot]);
|
|
1654
|
+
}
|
|
1655
|
+
}
|
|
1656
|
+
/** Append one generated token; finish + retire the lane when terminal. */
|
|
1657
|
+
consumeToken(req, tokenId) {
|
|
1658
|
+
req.generatedIds.push(tokenId);
|
|
1659
|
+
this.stats.tokensGenerated++;
|
|
1660
|
+
if (req.firstTokenAt === 0) req.firstTokenAt = performance.now();
|
|
1661
|
+
const { eosTokenId, eotTokenId } = this.host;
|
|
1662
|
+
if (eosTokenId !== null && tokenId === eosTokenId || eotTokenId !== null && tokenId === eotTokenId) {
|
|
1663
|
+
req.finishReason = "eos";
|
|
1664
|
+
this.finish(req);
|
|
1665
|
+
return;
|
|
1666
|
+
}
|
|
1667
|
+
const piece = this.host.decodeToken(tokenId);
|
|
1668
|
+
req.text += piece;
|
|
1669
|
+
if (req.stopSequences.some((s) => req.text.includes(s))) {
|
|
1670
|
+
for (const s of req.stopSequences) {
|
|
1671
|
+
const idx = req.text.indexOf(s);
|
|
1672
|
+
if (idx !== -1) req.text = req.text.slice(0, idx);
|
|
1673
|
+
}
|
|
1674
|
+
req.finishReason = "stop_sequence";
|
|
1675
|
+
this.finish(req);
|
|
1676
|
+
return;
|
|
1677
|
+
}
|
|
1678
|
+
req.onToken?.(piece, {
|
|
1679
|
+
requestId: req.id,
|
|
1680
|
+
tokenId,
|
|
1681
|
+
generated: req.generatedIds.length
|
|
1682
|
+
});
|
|
1683
|
+
if (req.generatedIds.length >= req.maxTokens) {
|
|
1684
|
+
req.finishReason = "max_tokens";
|
|
1685
|
+
this.finish(req);
|
|
1686
|
+
}
|
|
1687
|
+
}
|
|
1688
|
+
/** Retire the lane immediately: free KV blocks, open the slot, resolve. */
|
|
1689
|
+
finish(req) {
|
|
1690
|
+
req.state = "done";
|
|
1691
|
+
if (req.slot >= 0 && this.slots[req.slot] === req) {
|
|
1692
|
+
this.host.executor.retireBatchSlot(req.slot);
|
|
1693
|
+
this.slots[req.slot] = null;
|
|
1694
|
+
}
|
|
1695
|
+
this.reservedBlocks -= req.blocksReserved;
|
|
1696
|
+
this.stats.completed++;
|
|
1697
|
+
const finishedAt = performance.now();
|
|
1698
|
+
req.resolve({
|
|
1699
|
+
requestId: req.id,
|
|
1700
|
+
text: req.text,
|
|
1701
|
+
tokenIds: [...req.generatedIds],
|
|
1702
|
+
tokensGenerated: req.generatedIds.length,
|
|
1703
|
+
finishReason: req.finishReason,
|
|
1704
|
+
adapter: req.adapterName ?? void 0,
|
|
1705
|
+
submittedAt: req.submittedAt,
|
|
1706
|
+
admittedAt: req.admittedAt,
|
|
1707
|
+
firstTokenAt: req.firstTokenAt,
|
|
1708
|
+
finishedAt,
|
|
1709
|
+
ttftMs: req.firstTokenAt - req.submittedAt,
|
|
1710
|
+
e2eMs: finishedAt - req.submittedAt
|
|
1711
|
+
});
|
|
1712
|
+
}
|
|
1713
|
+
};
|
|
1714
|
+
|
|
1715
|
+
//#endregion
|
|
1716
|
+
//#region src/gpu/batch/swarm.ts
|
|
1717
|
+
/** Default reduce: member-keyed map, disambiguating duplicated members. */
|
|
1718
|
+
function reduceToKeyedMap(results) {
|
|
1719
|
+
const out = {};
|
|
1720
|
+
const seen = /* @__PURE__ */ new Map();
|
|
1721
|
+
for (const { member, result } of results) {
|
|
1722
|
+
const n = (seen.get(member) ?? 0) + 1;
|
|
1723
|
+
seen.set(member, n);
|
|
1724
|
+
out[n === 1 ? member : `${member}#${n}`] = result;
|
|
1725
|
+
}
|
|
1726
|
+
return out;
|
|
1727
|
+
}
|
|
1728
|
+
/**
|
|
1729
|
+
* A swarm handle: registered members + plan/reduce over the engine's
|
|
1730
|
+
* scheduler. Create via {@link WebGPUEngine.createSwarm}.
|
|
1731
|
+
*/
|
|
1732
|
+
var Swarm = class {
|
|
1733
|
+
scheduler;
|
|
1734
|
+
memberSources;
|
|
1735
|
+
plan;
|
|
1736
|
+
reduce;
|
|
1737
|
+
constructor(scheduler, options) {
|
|
1738
|
+
if (Object.keys(options.members).length === 0) throw new Error("Swarm: members must contain at least one member");
|
|
1739
|
+
this.scheduler = scheduler;
|
|
1740
|
+
this.memberSources = { ...options.members };
|
|
1741
|
+
this.plan = options.plan;
|
|
1742
|
+
this.reduce = options.reduce;
|
|
1743
|
+
}
|
|
1744
|
+
/** Member names, declaration order. */
|
|
1745
|
+
get members() {
|
|
1746
|
+
return Object.keys(this.memberSources);
|
|
1747
|
+
}
|
|
1748
|
+
/**
|
|
1749
|
+
* Run one input through the swarm: plan → concurrent fan-out (one batched
|
|
1750
|
+
* pass across the scheduler's lanes) → reduce.
|
|
1751
|
+
*
|
|
1752
|
+
* @param input The task input handed to the plan (default plan: broadcast
|
|
1753
|
+
* it to every member as the prompt).
|
|
1754
|
+
* @param options Per-request options applied to every sub-task (a task's
|
|
1755
|
+
* own `options` win field-by-field).
|
|
1756
|
+
*/
|
|
1757
|
+
async run(input, options = {}) {
|
|
1758
|
+
const tasks = this.plan ? await this.plan(input, this.members) : this.members.map((member) => ({
|
|
1759
|
+
member,
|
|
1760
|
+
prompt: input
|
|
1761
|
+
}));
|
|
1762
|
+
if (tasks.length === 0) throw new Error("Swarm: plan produced no sub-tasks");
|
|
1763
|
+
const results = await Promise.all(tasks.map(async (task) => {
|
|
1764
|
+
const result = await this.submitFor(task.member, task.prompt, {
|
|
1765
|
+
...options,
|
|
1766
|
+
...task.options
|
|
1767
|
+
});
|
|
1768
|
+
return {
|
|
1769
|
+
member: task.member,
|
|
1770
|
+
result
|
|
1771
|
+
};
|
|
1772
|
+
}));
|
|
1773
|
+
if (this.reduce) return this.reduce(results);
|
|
1774
|
+
return reduceToKeyedMap(results);
|
|
1775
|
+
}
|
|
1776
|
+
/**
|
|
1777
|
+
* Mode-A access: decode one prompt with a single member (no plan/reduce).
|
|
1778
|
+
* Still rides the shared scheduler, so it batches with any concurrent work.
|
|
1779
|
+
*/
|
|
1780
|
+
generate(member, prompt, options = {}) {
|
|
1781
|
+
return this.submitFor(member, prompt, options);
|
|
1782
|
+
}
|
|
1783
|
+
submitFor(member, prompt, options) {
|
|
1784
|
+
if (!Object.hasOwn(this.memberSources, member)) throw new Error(`Swarm: unknown member "${member}" (members: ${this.members.join(", ")})`);
|
|
1785
|
+
const adapter = this.memberSources[member] == null ? void 0 : member;
|
|
1786
|
+
return this.scheduler.submit(prompt, {
|
|
1787
|
+
...options,
|
|
1788
|
+
adapter
|
|
1789
|
+
});
|
|
1790
|
+
}
|
|
1791
|
+
};
|
|
1792
|
+
|
|
1793
|
+
//#endregion
|
|
1794
|
+
//#region src/gpu/architectures/gepard_tts.ts
|
|
1795
|
+
const GEPARD_START_OF_TEXT = 248073;
|
|
1796
|
+
const GEPARD_END_OF_TEXT = 248074;
|
|
1797
|
+
const GEPARD_START_OF_SPEECH = 248070;
|
|
1798
|
+
/** FSQ levels per dimension within each codec group (repeats per group). */
|
|
1799
|
+
const GEPARD_FSQ_LEVEL_PATTERN = [
|
|
1800
|
+
8,
|
|
1801
|
+
7,
|
|
1802
|
+
6,
|
|
1803
|
+
6
|
|
1804
|
+
];
|
|
1805
|
+
/** Number of codec groups (codebooks). */
|
|
1806
|
+
const GEPARD_NUM_GROUPS = 8;
|
|
1807
|
+
/** Total per-dimension channels per frame = groups × dims/group. */
|
|
1808
|
+
const GEPARD_NUM_CHANNELS = 32;
|
|
1809
|
+
/** NanoCodec 21.5 fps decoder geometry (codec_config.yaml, mlx-community mirror). */
|
|
1810
|
+
const GEPARD_CODEC_UP_SAMPLE_RATES = [
|
|
1811
|
+
8,
|
|
1812
|
+
8,
|
|
1813
|
+
4,
|
|
1814
|
+
2,
|
|
1815
|
+
2
|
|
1816
|
+
];
|
|
1817
|
+
const GEPARD_CODEC_INPUT_DIM = 32;
|
|
1818
|
+
/** PCM samples per audio frame (product of the up-sample rates). */
|
|
1819
|
+
const GEPARD_HOP = 1024;
|
|
1820
|
+
const GEPARD_SAMPLE_RATE = 22050;
|
|
1821
|
+
/** Per-channel FSQ cardinalities in head order (channel c = group⌊c/4⌋, dim c%4). */
|
|
1822
|
+
function gepardChannelLevels() {
|
|
1823
|
+
const out = [];
|
|
1824
|
+
for (let g = 0; g < GEPARD_NUM_GROUPS; g++) out.push(...GEPARD_FSQ_LEVEL_PATTERN);
|
|
1825
|
+
return out;
|
|
1826
|
+
}
|
|
1827
|
+
/** Matches `level_audio_{i}` head names (channel index = numeric suffix). */
|
|
1828
|
+
const LEVEL_AUDIO_RE = /^level_audio_(\d+)$/;
|
|
1829
|
+
/**
|
|
1830
|
+
* Parse gepard_config.json. Head order is recovered from the numeric suffix of
|
|
1831
|
+
* each `level_audio_{i}` key (NOT the serialized key order), so the wiring is
|
|
1832
|
+
* immune to any JSON round-trip that re-sorts keys.
|
|
1833
|
+
*/
|
|
1834
|
+
function parseGepardConfig(raw) {
|
|
1835
|
+
const backboneConfig = raw.backbone_config;
|
|
1836
|
+
if (!backboneConfig) throw new Error("gepard_config.json has no backbone_config — not a Gepard checkpoint.");
|
|
1837
|
+
const audioHeads = raw.audio_heads ?? {};
|
|
1838
|
+
const byIndex = [];
|
|
1839
|
+
for (const [name, size] of Object.entries(audioHeads)) {
|
|
1840
|
+
const m = LEVEL_AUDIO_RE.exec(name);
|
|
1841
|
+
if (!m) throw new Error(`Unexpected audio head name "${name}" in gepard_config.json.`);
|
|
1842
|
+
byIndex[Number(m[1])] = size;
|
|
1843
|
+
}
|
|
1844
|
+
if (byIndex.length === 0) throw new Error("gepard_config.json has no audio_heads.");
|
|
1845
|
+
for (let i = 0; i < byIndex.length; i++) if (byIndex[i] == null) throw new Error(`gepard_config.json audio_heads is missing level_audio_${i}.`);
|
|
1846
|
+
const special = raw.special_tokens ?? {};
|
|
1847
|
+
const rep = raw.text_repetition ?? {};
|
|
1848
|
+
const codec = raw.codec ?? {};
|
|
1849
|
+
return {
|
|
1850
|
+
backboneConfig,
|
|
1851
|
+
vocabSizes: byIndex,
|
|
1852
|
+
audioEmbedDim: raw.audio_embed_dim ?? 32,
|
|
1853
|
+
startOfText: special.start_of_text ?? GEPARD_START_OF_TEXT,
|
|
1854
|
+
endOfText: special.end_of_text ?? GEPARD_END_OF_TEXT,
|
|
1855
|
+
startOfSpeech: special.start_of_speech ?? GEPARD_START_OF_SPEECH,
|
|
1856
|
+
textRepetition: {
|
|
1857
|
+
enabled: rep.enabled ?? false,
|
|
1858
|
+
targetTextTokens: rep.target_text_tokens ?? 16,
|
|
1859
|
+
applyBelow: rep.apply_below ?? 13,
|
|
1860
|
+
maxRepeats: rep.max_repeats ?? 8
|
|
1861
|
+
},
|
|
1862
|
+
codecId: codec.codec_id ?? "nvidia/nemo-nano-codec-22khz-1.89kbps-21.5fps",
|
|
1863
|
+
fsqLevels: codec.fsq_levels ?? [...GEPARD_FSQ_LEVEL_PATTERN],
|
|
1864
|
+
sampleRate: codec.sample_rate ?? GEPARD_SAMPLE_RATE
|
|
1865
|
+
};
|
|
1866
|
+
}
|
|
1867
|
+
/**
|
|
1868
|
+
* Build the (possibly repeated) text-id layout — mirrors the reference
|
|
1869
|
+
* TextRepeater byte-for-byte (text_repetition.py):
|
|
1870
|
+
* R == 1 → [ SOT, *text, EOT, SOS ]
|
|
1871
|
+
* R > 1 → [ SOT, *text, EOT ] × (R−1) + [ SOT, *text, EOT, SOS ]
|
|
1872
|
+
* Only the final (canonical) copy carries SOS. R is deterministic from the
|
|
1873
|
+
* text-token count: ceil(target / n) capped at maxRepeats, applied only when
|
|
1874
|
+
* n < applyBelow (short prompts derail without the extra text mass).
|
|
1875
|
+
*/
|
|
1876
|
+
function buildGepardTextLayout(textTokenIds, cfg) {
|
|
1877
|
+
const rep = cfg.textRepetition;
|
|
1878
|
+
const n = textTokenIds.length;
|
|
1879
|
+
let repeats = 1;
|
|
1880
|
+
if (rep.enabled && n > 0 && n < rep.applyBelow && n < rep.targetTextTokens) repeats = Math.max(1, Math.min(Math.ceil(rep.targetTextTokens / n), rep.maxRepeats));
|
|
1881
|
+
const block = [
|
|
1882
|
+
cfg.startOfText,
|
|
1883
|
+
...textTokenIds,
|
|
1884
|
+
cfg.endOfText
|
|
1885
|
+
];
|
|
1886
|
+
const out = [];
|
|
1887
|
+
for (let r = 0; r < repeats - 1; r++) out.push(...block);
|
|
1888
|
+
out.push(...block, cfg.startOfSpeech);
|
|
1889
|
+
return out;
|
|
1890
|
+
}
|
|
1891
|
+
/**
|
|
1892
|
+
* Dequantize per-channel FSQ codes [32, T] (channel-major) into the decoder
|
|
1893
|
+
* latent [32, T] f32: value = (code − L//2) / (L//2), matching
|
|
1894
|
+
* UnfoldedCodecModel.decode_from_codes / codec_ops.dequantize_codes.
|
|
1895
|
+
*/
|
|
1896
|
+
function gepardCodesToLatent(codes, numFrames) {
|
|
1897
|
+
const levels = gepardChannelLevels();
|
|
1898
|
+
const latent = new Float32Array(GEPARD_NUM_CHANNELS * numFrames);
|
|
1899
|
+
for (let c = 0; c < GEPARD_NUM_CHANNELS; c++) {
|
|
1900
|
+
const scale = Math.max(1, Math.floor(levels[c] / 2));
|
|
1901
|
+
const base = c * numFrames;
|
|
1902
|
+
for (let t = 0; t < numFrames; t++) latent[base + t] = (codes[base + t] - scale) / scale;
|
|
1903
|
+
}
|
|
1904
|
+
return latent;
|
|
1905
|
+
}
|
|
1906
|
+
function generateGepardBackboneGraph(backboneConfig, dtype, groupSizeOverride, kvDtype) {
|
|
1907
|
+
const hidden_size = backboneConfig.hidden_size;
|
|
1908
|
+
const num_layers = backboneConfig.num_hidden_layers;
|
|
1909
|
+
const num_heads = backboneConfig.num_attention_heads;
|
|
1910
|
+
const num_kv_heads = backboneConfig.num_key_value_heads ?? num_heads;
|
|
1911
|
+
const intermediate_size = backboneConfig.intermediate_size;
|
|
1912
|
+
const context_length = backboneConfig.max_position_embeddings ?? 262144;
|
|
1913
|
+
const rms_norm_eps = backboneConfig.rms_norm_eps ?? 1e-6;
|
|
1914
|
+
const head_dim = backboneConfig.head_dim ?? Math.floor(hidden_size / num_heads);
|
|
1915
|
+
const q_dim = num_heads * head_dim;
|
|
1916
|
+
const kv_dim = num_kv_heads * head_dim;
|
|
1917
|
+
const ropeParams = backboneConfig.rope_parameters;
|
|
1918
|
+
const rope_base = ropeParams?.rope_theta ?? backboneConfig.rope_theta ?? 1e7;
|
|
1919
|
+
const partialRotaryFactor = ropeParams?.partial_rotary_factor ?? backboneConfig.partial_rotary_factor ?? 1;
|
|
1920
|
+
const rope_dim = Math.floor(head_dim * partialRotaryFactor);
|
|
1921
|
+
const layerTypes = backboneConfig.layer_types ?? [];
|
|
1922
|
+
for (let i = 0; i < num_layers; i++) if (layerTypes.length > i && layerTypes[i] !== "full_attention") throw new Error(`Gepard backbone expects all-full-attention layers; layer ${i} is "${layerTypes[i]}".`);
|
|
1923
|
+
const config = {
|
|
1924
|
+
hidden_size,
|
|
1925
|
+
num_layers,
|
|
1926
|
+
num_heads,
|
|
1927
|
+
num_kv_heads,
|
|
1928
|
+
head_dim,
|
|
1929
|
+
intermediate_size,
|
|
1930
|
+
vocab_size: hidden_size,
|
|
1931
|
+
context_length,
|
|
1932
|
+
rms_norm_eps,
|
|
1933
|
+
norm_type: "rmsnorm",
|
|
1934
|
+
rope_base,
|
|
1935
|
+
rope_dim,
|
|
1936
|
+
kv_layout: "LHSd",
|
|
1937
|
+
is_moe: false,
|
|
1938
|
+
has_vision_tower: false
|
|
1939
|
+
};
|
|
1940
|
+
const capabilities = {
|
|
1941
|
+
text: true,
|
|
1942
|
+
vision: false,
|
|
1943
|
+
moe: false
|
|
1944
|
+
};
|
|
1945
|
+
const tensors = {};
|
|
1946
|
+
const nodes = [];
|
|
1947
|
+
const executionOrder = [];
|
|
1948
|
+
const addTensor = (desc) => {
|
|
1949
|
+
tensors[desc.name] = desc;
|
|
1950
|
+
};
|
|
1951
|
+
const addNode = (node) => {
|
|
1952
|
+
nodes.push(node);
|
|
1953
|
+
executionOrder.push(node.id);
|
|
1954
|
+
};
|
|
1955
|
+
const useQ4 = dtype === "q4";
|
|
1956
|
+
const groupSize = groupSizeOverride ?? DEFAULT_GROUP_SIZE;
|
|
1957
|
+
const addLinearOp = (id, activationInput, weightName, weightShape, outputName, K, N) => {
|
|
1958
|
+
if (useQ4) {
|
|
1959
|
+
const wQ = `${weightName}.q`;
|
|
1960
|
+
const wScales = `${weightName}.scales`;
|
|
1961
|
+
const wZeros = `${weightName}.zeros`;
|
|
1962
|
+
const totalElements = K * N;
|
|
1963
|
+
addTensor({
|
|
1964
|
+
name: wQ,
|
|
1965
|
+
shape: [Math.ceil(totalElements / 8)],
|
|
1966
|
+
dtype: "u32",
|
|
1967
|
+
storage: "constant"
|
|
1968
|
+
});
|
|
1969
|
+
addTensor({
|
|
1970
|
+
name: wScales,
|
|
1971
|
+
shape: [Math.ceil(totalElements / groupSize)],
|
|
1972
|
+
dtype: "f32",
|
|
1973
|
+
storage: "constant"
|
|
1974
|
+
});
|
|
1975
|
+
addTensor({
|
|
1976
|
+
name: wZeros,
|
|
1977
|
+
shape: [Math.ceil(totalElements / groupSize)],
|
|
1978
|
+
dtype: "f32",
|
|
1979
|
+
storage: "constant"
|
|
1980
|
+
});
|
|
1981
|
+
addNode({
|
|
1982
|
+
id,
|
|
1983
|
+
opType: "MatMulInt4",
|
|
1984
|
+
inputs: [
|
|
1985
|
+
activationInput,
|
|
1986
|
+
wQ,
|
|
1987
|
+
wScales,
|
|
1988
|
+
wZeros
|
|
1989
|
+
],
|
|
1990
|
+
outputs: [outputName],
|
|
1991
|
+
attributes: {
|
|
1992
|
+
M_tensor: activationInput,
|
|
1993
|
+
K,
|
|
1994
|
+
N,
|
|
1995
|
+
group_size: groupSize
|
|
1996
|
+
}
|
|
1997
|
+
});
|
|
1998
|
+
return;
|
|
1999
|
+
}
|
|
2000
|
+
addTensor({
|
|
2001
|
+
name: weightName,
|
|
2002
|
+
shape: weightShape,
|
|
2003
|
+
dtype: "f32",
|
|
2004
|
+
storage: "constant",
|
|
2005
|
+
safetensorsKey: weightName
|
|
2006
|
+
});
|
|
2007
|
+
addNode({
|
|
2008
|
+
id,
|
|
2009
|
+
opType: "MatMul",
|
|
2010
|
+
inputs: [activationInput, weightName],
|
|
2011
|
+
outputs: [outputName],
|
|
2012
|
+
attributes: {
|
|
2013
|
+
M_tensor: activationInput,
|
|
2014
|
+
K,
|
|
2015
|
+
N
|
|
2016
|
+
}
|
|
2017
|
+
});
|
|
2018
|
+
};
|
|
2019
|
+
addTensor({
|
|
2020
|
+
name: "inputs_embeds",
|
|
2021
|
+
shape: ["T", hidden_size],
|
|
2022
|
+
dtype: "f32",
|
|
2023
|
+
storage: "activation"
|
|
2024
|
+
});
|
|
2025
|
+
const buildAttentionBlock = (i, prefix, norm1Out) => {
|
|
2026
|
+
const qOut = `${prefix}_q`;
|
|
2027
|
+
const gateOut = `${prefix}_attn_gate`;
|
|
2028
|
+
const kOut = `${prefix}_k`;
|
|
2029
|
+
const vOut = `${prefix}_v`;
|
|
2030
|
+
addTensor({
|
|
2031
|
+
name: qOut,
|
|
2032
|
+
shape: ["T", q_dim],
|
|
2033
|
+
dtype: "f32",
|
|
2034
|
+
storage: "activation"
|
|
2035
|
+
});
|
|
2036
|
+
addTensor({
|
|
2037
|
+
name: gateOut,
|
|
2038
|
+
shape: ["T", q_dim],
|
|
2039
|
+
dtype: "f32",
|
|
2040
|
+
storage: "activation"
|
|
2041
|
+
});
|
|
2042
|
+
addTensor({
|
|
2043
|
+
name: kOut,
|
|
2044
|
+
shape: ["T", kv_dim],
|
|
2045
|
+
dtype: "f32",
|
|
2046
|
+
storage: "activation"
|
|
2047
|
+
});
|
|
2048
|
+
addTensor({
|
|
2049
|
+
name: vOut,
|
|
2050
|
+
shape: ["T", kv_dim],
|
|
2051
|
+
dtype: "f32",
|
|
2052
|
+
storage: "activation"
|
|
2053
|
+
});
|
|
2054
|
+
addLinearOp(`${prefix}_q_proj`, norm1Out, CANONICAL_KEYS.qProj(i), [q_dim, hidden_size], qOut, hidden_size, q_dim);
|
|
2055
|
+
addLinearOp(`${prefix}_gate_proj`, norm1Out, CANONICAL_KEYS.attnGate(i), [q_dim, hidden_size], gateOut, hidden_size, q_dim);
|
|
2056
|
+
addLinearOp(`${prefix}_k_proj`, norm1Out, CANONICAL_KEYS.kProj(i), [kv_dim, hidden_size], kOut, hidden_size, kv_dim);
|
|
2057
|
+
addLinearOp(`${prefix}_v_proj`, norm1Out, CANONICAL_KEYS.vProj(i), [kv_dim, hidden_size], vOut, hidden_size, kv_dim);
|
|
2058
|
+
const qNormW = CANONICAL_KEYS.qNorm(i);
|
|
2059
|
+
const kNormW = CANONICAL_KEYS.kNorm(i);
|
|
2060
|
+
addTensor({
|
|
2061
|
+
name: qNormW,
|
|
2062
|
+
shape: [head_dim],
|
|
2063
|
+
dtype: "f32",
|
|
2064
|
+
storage: "constant",
|
|
2065
|
+
safetensorsKey: qNormW
|
|
2066
|
+
});
|
|
2067
|
+
addTensor({
|
|
2068
|
+
name: kNormW,
|
|
2069
|
+
shape: [head_dim],
|
|
2070
|
+
dtype: "f32",
|
|
2071
|
+
storage: "constant",
|
|
2072
|
+
safetensorsKey: kNormW
|
|
2073
|
+
});
|
|
2074
|
+
const qNormed = `${prefix}_q_normed`;
|
|
2075
|
+
const kNormed = `${prefix}_k_normed`;
|
|
2076
|
+
addTensor({
|
|
2077
|
+
name: qNormed,
|
|
2078
|
+
shape: ["T", q_dim],
|
|
2079
|
+
dtype: "f32",
|
|
2080
|
+
storage: "activation"
|
|
2081
|
+
});
|
|
2082
|
+
addTensor({
|
|
2083
|
+
name: kNormed,
|
|
2084
|
+
shape: ["T", kv_dim],
|
|
2085
|
+
dtype: "f32",
|
|
2086
|
+
storage: "activation"
|
|
2087
|
+
});
|
|
2088
|
+
addNode({
|
|
2089
|
+
id: `${prefix}_q_norm`,
|
|
2090
|
+
opType: "RMSNorm",
|
|
2091
|
+
inputs: [qOut, qNormW],
|
|
2092
|
+
outputs: [qNormed],
|
|
2093
|
+
attributes: {
|
|
2094
|
+
hidden_size: head_dim,
|
|
2095
|
+
eps: rms_norm_eps,
|
|
2096
|
+
seq_len_tensor: qOut
|
|
2097
|
+
}
|
|
2098
|
+
});
|
|
2099
|
+
addNode({
|
|
2100
|
+
id: `${prefix}_k_norm`,
|
|
2101
|
+
opType: "RMSNorm",
|
|
2102
|
+
inputs: [kOut, kNormW],
|
|
2103
|
+
outputs: [kNormed],
|
|
2104
|
+
attributes: {
|
|
2105
|
+
hidden_size: head_dim,
|
|
2106
|
+
eps: rms_norm_eps,
|
|
2107
|
+
seq_len_tensor: kOut
|
|
2108
|
+
}
|
|
2109
|
+
});
|
|
2110
|
+
addNode({
|
|
2111
|
+
id: `${prefix}_rope`,
|
|
2112
|
+
opType: "RoPE",
|
|
2113
|
+
inputs: [qNormed, kNormed],
|
|
2114
|
+
outputs: [qNormed, kNormed],
|
|
2115
|
+
attributes: {
|
|
2116
|
+
head_dim,
|
|
2117
|
+
num_q_heads: num_heads,
|
|
2118
|
+
num_kv_heads,
|
|
2119
|
+
rope_base,
|
|
2120
|
+
rope_dim,
|
|
2121
|
+
partial_rotary: true,
|
|
2122
|
+
seq_len_tensor: qNormed
|
|
2123
|
+
}
|
|
2124
|
+
});
|
|
2125
|
+
const kvCacheDtype = kvDtype ?? "f32";
|
|
2126
|
+
const kCache = `${prefix}_k_cache`;
|
|
2127
|
+
const vCache = `${prefix}_v_cache`;
|
|
2128
|
+
addTensor({
|
|
2129
|
+
name: kCache,
|
|
2130
|
+
shape: ["L_max", kv_dim],
|
|
2131
|
+
dtype: kvCacheDtype,
|
|
2132
|
+
storage: "kv_cache"
|
|
2133
|
+
});
|
|
2134
|
+
addTensor({
|
|
2135
|
+
name: vCache,
|
|
2136
|
+
shape: ["L_max", kv_dim],
|
|
2137
|
+
dtype: kvCacheDtype,
|
|
2138
|
+
storage: "kv_cache"
|
|
2139
|
+
});
|
|
2140
|
+
addNode({
|
|
2141
|
+
id: `${prefix}_kv_cache_k`,
|
|
2142
|
+
opType: "KVCacheAppend",
|
|
2143
|
+
inputs: [kNormed],
|
|
2144
|
+
outputs: [kCache],
|
|
2145
|
+
attributes: {
|
|
2146
|
+
width: kv_dim,
|
|
2147
|
+
T_tensor: kNormed
|
|
2148
|
+
}
|
|
2149
|
+
});
|
|
2150
|
+
addNode({
|
|
2151
|
+
id: `${prefix}_kv_cache_v`,
|
|
2152
|
+
opType: "KVCacheAppend",
|
|
2153
|
+
inputs: [vOut],
|
|
2154
|
+
outputs: [vCache],
|
|
2155
|
+
attributes: {
|
|
2156
|
+
width: kv_dim,
|
|
2157
|
+
T_tensor: vOut
|
|
2158
|
+
}
|
|
2159
|
+
});
|
|
2160
|
+
const attnOut = `${prefix}_attn_out`;
|
|
2161
|
+
addTensor({
|
|
2162
|
+
name: attnOut,
|
|
2163
|
+
shape: ["T", q_dim],
|
|
2164
|
+
dtype: "f32",
|
|
2165
|
+
storage: "activation"
|
|
2166
|
+
});
|
|
2167
|
+
addNode({
|
|
2168
|
+
id: `${prefix}_attn`,
|
|
2169
|
+
opType: "Attention",
|
|
2170
|
+
inputs: [
|
|
2171
|
+
qNormed,
|
|
2172
|
+
kCache,
|
|
2173
|
+
vCache
|
|
2174
|
+
],
|
|
2175
|
+
outputs: [attnOut],
|
|
2176
|
+
attributes: {
|
|
2177
|
+
hidden_size: q_dim,
|
|
2178
|
+
num_q_heads: num_heads,
|
|
2179
|
+
num_kv_heads,
|
|
2180
|
+
head_dim,
|
|
2181
|
+
causal: true,
|
|
2182
|
+
layer_index: i
|
|
2183
|
+
}
|
|
2184
|
+
});
|
|
2185
|
+
const gatedOut = `${prefix}_gated`;
|
|
2186
|
+
addTensor({
|
|
2187
|
+
name: gatedOut,
|
|
2188
|
+
shape: ["T", q_dim],
|
|
2189
|
+
dtype: "f32",
|
|
2190
|
+
storage: "activation"
|
|
2191
|
+
});
|
|
2192
|
+
addNode({
|
|
2193
|
+
id: `${prefix}_sigmoid_gate`,
|
|
2194
|
+
opType: "SigmoidGate",
|
|
2195
|
+
inputs: [attnOut, gateOut],
|
|
2196
|
+
outputs: [gatedOut],
|
|
2197
|
+
attributes: { count_tensor: attnOut }
|
|
2198
|
+
});
|
|
2199
|
+
const oProjOut = `${prefix}_o_proj_out`;
|
|
2200
|
+
addTensor({
|
|
2201
|
+
name: oProjOut,
|
|
2202
|
+
shape: ["T", hidden_size],
|
|
2203
|
+
dtype: "f32",
|
|
2204
|
+
storage: "activation"
|
|
2205
|
+
});
|
|
2206
|
+
addLinearOp(`${prefix}_o_proj`, gatedOut, CANONICAL_KEYS.oProj(i), [hidden_size, q_dim], oProjOut, q_dim, hidden_size);
|
|
2207
|
+
return oProjOut;
|
|
2208
|
+
};
|
|
2209
|
+
let prevOutput = "inputs_embeds";
|
|
2210
|
+
let pendingNorm1Out = null;
|
|
2211
|
+
for (let i = 0; i < num_layers; i++) {
|
|
2212
|
+
const prefix = `layer${i}`;
|
|
2213
|
+
let norm1Out;
|
|
2214
|
+
if (pendingNorm1Out) norm1Out = pendingNorm1Out;
|
|
2215
|
+
else {
|
|
2216
|
+
const normWeightName = CANONICAL_KEYS.layerInputNorm(i);
|
|
2217
|
+
addTensor({
|
|
2218
|
+
name: normWeightName,
|
|
2219
|
+
shape: [hidden_size],
|
|
2220
|
+
dtype: "f32",
|
|
2221
|
+
storage: "constant",
|
|
2222
|
+
safetensorsKey: normWeightName
|
|
2223
|
+
});
|
|
2224
|
+
norm1Out = `${prefix}_norm1_out`;
|
|
2225
|
+
addTensor({
|
|
2226
|
+
name: norm1Out,
|
|
2227
|
+
shape: ["T", hidden_size],
|
|
2228
|
+
dtype: "f32",
|
|
2229
|
+
storage: "activation"
|
|
2230
|
+
});
|
|
2231
|
+
addNode({
|
|
2232
|
+
id: `${prefix}_norm1`,
|
|
2233
|
+
opType: "RMSNorm",
|
|
2234
|
+
inputs: [prevOutput, normWeightName],
|
|
2235
|
+
outputs: [norm1Out],
|
|
2236
|
+
attributes: {
|
|
2237
|
+
hidden_size,
|
|
2238
|
+
eps: rms_norm_eps,
|
|
2239
|
+
seq_len_tensor: prevOutput
|
|
2240
|
+
}
|
|
2241
|
+
});
|
|
2242
|
+
}
|
|
2243
|
+
const attnOutput = buildAttentionBlock(i, prefix, norm1Out);
|
|
2244
|
+
const resid1Out = `${prefix}_resid1`;
|
|
2245
|
+
const postNormW = CANONICAL_KEYS.layerPostAttnNorm(i);
|
|
2246
|
+
const norm2Out = `${prefix}_norm2_out`;
|
|
2247
|
+
addTensor({
|
|
2248
|
+
name: resid1Out,
|
|
2249
|
+
shape: ["T", hidden_size],
|
|
2250
|
+
dtype: "f32",
|
|
2251
|
+
storage: "activation"
|
|
2252
|
+
});
|
|
2253
|
+
addTensor({
|
|
2254
|
+
name: postNormW,
|
|
2255
|
+
shape: [hidden_size],
|
|
2256
|
+
dtype: "f32",
|
|
2257
|
+
storage: "constant",
|
|
2258
|
+
safetensorsKey: postNormW
|
|
2259
|
+
});
|
|
2260
|
+
addTensor({
|
|
2261
|
+
name: norm2Out,
|
|
2262
|
+
shape: ["T", hidden_size],
|
|
2263
|
+
dtype: "f32",
|
|
2264
|
+
storage: "activation"
|
|
2265
|
+
});
|
|
2266
|
+
addNode({
|
|
2267
|
+
id: `${prefix}_resid1_norm2`,
|
|
2268
|
+
opType: "ResidualRMSNorm",
|
|
2269
|
+
inputs: [
|
|
2270
|
+
prevOutput,
|
|
2271
|
+
attnOutput,
|
|
2272
|
+
postNormW
|
|
2273
|
+
],
|
|
2274
|
+
outputs: [resid1Out, norm2Out],
|
|
2275
|
+
attributes: {
|
|
2276
|
+
hidden_size,
|
|
2277
|
+
eps: rms_norm_eps,
|
|
2278
|
+
seq_len_tensor: prevOutput
|
|
2279
|
+
}
|
|
2280
|
+
});
|
|
2281
|
+
const gateW = CANONICAL_KEYS.gateProj(i);
|
|
2282
|
+
const upW = CANONICAL_KEYS.upProj(i);
|
|
2283
|
+
const downW = CANONICAL_KEYS.downProj(i);
|
|
2284
|
+
const mlpGateOut = `${prefix}_gate_out`;
|
|
2285
|
+
const upOut = `${prefix}_up_out`;
|
|
2286
|
+
const swigluOut = `${prefix}_swiglu_out`;
|
|
2287
|
+
const mlpOut = `${prefix}_mlp_out`;
|
|
2288
|
+
addTensor({
|
|
2289
|
+
name: mlpGateOut,
|
|
2290
|
+
shape: ["T", intermediate_size],
|
|
2291
|
+
dtype: "f32",
|
|
2292
|
+
storage: "activation"
|
|
2293
|
+
});
|
|
2294
|
+
addTensor({
|
|
2295
|
+
name: upOut,
|
|
2296
|
+
shape: ["T", intermediate_size],
|
|
2297
|
+
dtype: "f32",
|
|
2298
|
+
storage: "activation"
|
|
2299
|
+
});
|
|
2300
|
+
addTensor({
|
|
2301
|
+
name: swigluOut,
|
|
2302
|
+
shape: ["T", intermediate_size],
|
|
2303
|
+
dtype: "f32",
|
|
2304
|
+
storage: "activation"
|
|
2305
|
+
});
|
|
2306
|
+
addTensor({
|
|
2307
|
+
name: mlpOut,
|
|
2308
|
+
shape: ["T", hidden_size],
|
|
2309
|
+
dtype: "f32",
|
|
2310
|
+
storage: "activation"
|
|
2311
|
+
});
|
|
2312
|
+
addLinearOp(`${prefix}_gate`, norm2Out, gateW, [intermediate_size, hidden_size], mlpGateOut, hidden_size, intermediate_size);
|
|
2313
|
+
addLinearOp(`${prefix}_up`, norm2Out, upW, [intermediate_size, hidden_size], upOut, hidden_size, intermediate_size);
|
|
2314
|
+
addNode({
|
|
2315
|
+
id: `${prefix}_swiglu`,
|
|
2316
|
+
opType: "SwiGLU",
|
|
2317
|
+
inputs: [mlpGateOut, upOut],
|
|
2318
|
+
outputs: [swigluOut],
|
|
2319
|
+
attributes: { count_tensor: mlpGateOut }
|
|
2320
|
+
});
|
|
2321
|
+
addLinearOp(`${prefix}_down`, swigluOut, downW, [hidden_size, intermediate_size], mlpOut, intermediate_size, hidden_size);
|
|
2322
|
+
const resid2Out = `${prefix}_resid2`;
|
|
2323
|
+
addTensor({
|
|
2324
|
+
name: resid2Out,
|
|
2325
|
+
shape: ["T", hidden_size],
|
|
2326
|
+
dtype: "f32",
|
|
2327
|
+
storage: "activation"
|
|
2328
|
+
});
|
|
2329
|
+
if (i < num_layers - 1) {
|
|
2330
|
+
const nextNormW = CANONICAL_KEYS.layerInputNorm(i + 1);
|
|
2331
|
+
addTensor({
|
|
2332
|
+
name: nextNormW,
|
|
2333
|
+
shape: [hidden_size],
|
|
2334
|
+
dtype: "f32",
|
|
2335
|
+
storage: "constant",
|
|
2336
|
+
safetensorsKey: nextNormW
|
|
2337
|
+
});
|
|
2338
|
+
const nextNorm1Out = `layer${i + 1}_norm1_out`;
|
|
2339
|
+
addTensor({
|
|
2340
|
+
name: nextNorm1Out,
|
|
2341
|
+
shape: ["T", hidden_size],
|
|
2342
|
+
dtype: "f32",
|
|
2343
|
+
storage: "activation"
|
|
2344
|
+
});
|
|
2345
|
+
addNode({
|
|
2346
|
+
id: `${prefix}_resid2_norm1`,
|
|
2347
|
+
opType: "ResidualRMSNorm",
|
|
2348
|
+
inputs: [
|
|
2349
|
+
resid1Out,
|
|
2350
|
+
mlpOut,
|
|
2351
|
+
nextNormW
|
|
2352
|
+
],
|
|
2353
|
+
outputs: [resid2Out, nextNorm1Out],
|
|
2354
|
+
attributes: {
|
|
2355
|
+
hidden_size,
|
|
2356
|
+
eps: rms_norm_eps,
|
|
2357
|
+
seq_len_tensor: resid1Out
|
|
2358
|
+
}
|
|
2359
|
+
});
|
|
2360
|
+
pendingNorm1Out = nextNorm1Out;
|
|
2361
|
+
} else {
|
|
2362
|
+
addNode({
|
|
2363
|
+
id: `${prefix}_resid2`,
|
|
2364
|
+
opType: "Add",
|
|
2365
|
+
inputs: [resid1Out, mlpOut],
|
|
2366
|
+
outputs: [resid2Out],
|
|
2367
|
+
attributes: {
|
|
2368
|
+
count_tensor: resid1Out,
|
|
2369
|
+
hidden_size
|
|
2370
|
+
}
|
|
2371
|
+
});
|
|
2372
|
+
pendingNorm1Out = null;
|
|
2373
|
+
}
|
|
2374
|
+
prevOutput = resid2Out;
|
|
2375
|
+
}
|
|
2376
|
+
addTensor({
|
|
2377
|
+
name: CANONICAL_KEYS.FINAL_NORM,
|
|
2378
|
+
shape: [hidden_size],
|
|
2379
|
+
dtype: "f32",
|
|
2380
|
+
storage: "constant",
|
|
2381
|
+
safetensorsKey: CANONICAL_KEYS.FINAL_NORM
|
|
2382
|
+
});
|
|
2383
|
+
addTensor({
|
|
2384
|
+
name: "final_norm_out",
|
|
2385
|
+
shape: ["T", hidden_size],
|
|
2386
|
+
dtype: "f32",
|
|
2387
|
+
storage: "activation"
|
|
2388
|
+
});
|
|
2389
|
+
addNode({
|
|
2390
|
+
id: "final_norm",
|
|
2391
|
+
opType: "RMSNorm",
|
|
2392
|
+
inputs: [prevOutput, CANONICAL_KEYS.FINAL_NORM],
|
|
2393
|
+
outputs: ["final_norm_out"],
|
|
2394
|
+
attributes: {
|
|
2395
|
+
hidden_size,
|
|
2396
|
+
eps: rms_norm_eps,
|
|
2397
|
+
seq_len_tensor: prevOutput
|
|
2398
|
+
}
|
|
2399
|
+
});
|
|
2400
|
+
addTensor({
|
|
2401
|
+
name: "logits",
|
|
2402
|
+
shape: [1, hidden_size],
|
|
2403
|
+
dtype: "f32",
|
|
2404
|
+
storage: "activation"
|
|
2405
|
+
});
|
|
2406
|
+
addNode({
|
|
2407
|
+
id: "slice_last_row",
|
|
2408
|
+
opType: "SliceLastRow",
|
|
2409
|
+
inputs: ["final_norm_out"],
|
|
2410
|
+
outputs: ["logits"],
|
|
2411
|
+
attributes: { width: hidden_size }
|
|
2412
|
+
});
|
|
2413
|
+
return {
|
|
2414
|
+
architecture: "GepardForTTS",
|
|
2415
|
+
config,
|
|
2416
|
+
capabilities,
|
|
2417
|
+
tensors,
|
|
2418
|
+
nodes,
|
|
2419
|
+
executionOrder,
|
|
2420
|
+
inputs: ["inputs_embeds"],
|
|
2421
|
+
outputs: ["logits"]
|
|
2422
|
+
};
|
|
2423
|
+
}
|
|
2424
|
+
|
|
2425
|
+
//#endregion
|
|
2426
|
+
//#region src/gpu/gepard-tts.ts
|
|
2427
|
+
/**
|
|
2428
|
+
* Gepard-1.0 — native streaming text-to-speech engine for Gerbil's WebGPU backend.
|
|
2429
|
+
*
|
|
2430
|
+
* Gepard (nineninesix/gepard-1.0, Apache-2.0) is a frame-synchronous multihead
|
|
2431
|
+
* codec-LM TTS. Unlike the token-serial Kani/Oute engines (which emit one codec
|
|
2432
|
+
* token per forward), Gepard emits ONE WHOLE AUDIO FRAME per backbone forward:
|
|
2433
|
+
*
|
|
2434
|
+
* 1. BACKBONE (GPU): stock Qwen3.5 full-attention body (14 layers, hidden 1024)
|
|
2435
|
+
* run through generateGepardBackboneGraph — input is a host-written
|
|
2436
|
+
* `inputs_embeds` row, output is the last position's final-norm hidden state.
|
|
2437
|
+
* 2. OVERLAY (host): 32 parallel classifier heads (cardinalities [8,7,6,6]×8,
|
|
2438
|
+
* 216 logits total) + a binary stop head read the hidden state; the sampled
|
|
2439
|
+
* 32 FSQ codes feed back through per-channel embeddings → concat → GELU MLP
|
|
2440
|
+
* → affine-free LayerNorm → × audio_embed_scale as the next step's input.
|
|
2441
|
+
* 3. CODEC (GPU): NanoCodec 21.5 fps decoder (same validated CausalHiFiGAN
|
|
2442
|
+
* graph family as Kani's codec, rates [8,8,4,2,2], hop 1024) turns the
|
|
2443
|
+
* per-channel codes into 22.05 kHz PCM — decoded incrementally in chunks
|
|
2444
|
+
* with left-context lookback, so audio streams out while the model is
|
|
2445
|
+
* still generating (≈1.5 s of audio per 32-frame chunk).
|
|
2446
|
+
*
|
|
2447
|
+
* Mirrors gepard-inference (runner.py) semantics: adaptive text repetition,
|
|
2448
|
+
* SOS-seeded first frame, stop-head termination, independent per-head sampling
|
|
2449
|
+
* (temperature / top-k / windowed repetition penalty) in fp32. Voice cloning
|
|
2450
|
+
* (the ref_compressor speaker prefix) is NOT wired yet — the default voice ships
|
|
2451
|
+
* first; cloning is a documented follow-up.
|
|
2452
|
+
*/
|
|
2453
|
+
/** Codec decode chunking: 32 frames ≈ 1.49 s of audio per streamed chunk. */
|
|
2454
|
+
const GEPARD_CODEC_CHUNK_FRAMES = 32;
|
|
2455
|
+
/**
|
|
2456
|
+
* Left-context warm-up frames re-decoded (and discarded) per chunk. The decoder
|
|
2457
|
+
* is fully causal with a bounded receptive field; 64 frames (≈65k samples) of
|
|
2458
|
+
* lookback keeps chunked output numerically identical to a monolithic decode
|
|
2459
|
+
* while staying under the WebGPU per-dimension dispatch cap.
|
|
2460
|
+
*/
|
|
2461
|
+
const GEPARD_CODEC_LOOKBACK_FRAMES = 64;
|
|
2462
|
+
/** Hard ceiling on generated frames (the reference runner's max_frames). */
|
|
2463
|
+
const DEFAULT_MAX_FRAMES$1 = 2e3;
|
|
2464
|
+
/** PyTorch nn.LayerNorm default epsilon (audio_embed_proj's affine-free LN). */
|
|
2465
|
+
const LAYERNORM_EPS = 1e-5;
|
|
2466
|
+
/** Build a fresh map of just the constant weights a graph references. */
|
|
2467
|
+
function selectGraphWeights$2(graph, weights) {
|
|
2468
|
+
const out = /* @__PURE__ */ new Map();
|
|
2469
|
+
for (const [name, desc] of Object.entries(graph.tensors)) {
|
|
2470
|
+
if (desc.storage !== "constant") continue;
|
|
2471
|
+
const w = weights.get(name) ?? (desc.safetensorsKey ? weights.get(desc.safetensorsKey) : void 0);
|
|
2472
|
+
if (w) out.set(name, w);
|
|
2473
|
+
}
|
|
2474
|
+
return out;
|
|
2475
|
+
}
|
|
2476
|
+
/** Abramowitz–Stegun 7.1.26 erf approximation (max |err| ≈ 1.5e-7). */
|
|
2477
|
+
function erf(x) {
|
|
2478
|
+
const sign = x < 0 ? -1 : 1;
|
|
2479
|
+
const ax = Math.abs(x);
|
|
2480
|
+
const t = 1 / (1 + .3275911 * ax);
|
|
2481
|
+
return sign * (1 - ((((1.061405429 * t - 1.453152027) * t + 1.421413741) * t - .284496736) * t + .254829592) * t * Math.exp(-ax * ax));
|
|
2482
|
+
}
|
|
2483
|
+
/** Exact (erf-based) GELU, matching PyTorch nn.GELU()'s default. */
|
|
2484
|
+
function gelu(x) {
|
|
2485
|
+
return .5 * x * (1 + erf(x / Math.SQRT2));
|
|
2486
|
+
}
|
|
2487
|
+
/** y = W·x + b for a row-major [out, in] weight. */
|
|
2488
|
+
function linear(w, b, x, outDim, inDim, out) {
|
|
2489
|
+
for (let o = 0; o < outDim; o++) {
|
|
2490
|
+
let acc = b ? b[o] : 0;
|
|
2491
|
+
const base = o * inDim;
|
|
2492
|
+
for (let k = 0; k < inDim; k++) acc += w[base + k] * x[k];
|
|
2493
|
+
out[o] = acc;
|
|
2494
|
+
}
|
|
2495
|
+
}
|
|
2496
|
+
var GepardTTS = class GepardTTS {
|
|
2497
|
+
ctx;
|
|
2498
|
+
loaded;
|
|
2499
|
+
tokenizer;
|
|
2500
|
+
cfg;
|
|
2501
|
+
maxSeqLen;
|
|
2502
|
+
hidden;
|
|
2503
|
+
/** Backbone executor (built once; reused across speak() calls). */
|
|
2504
|
+
backboneExec;
|
|
2505
|
+
/** Cached codec executor for the steady-state window size (chunk + lookback). */
|
|
2506
|
+
codecExecCache = /* @__PURE__ */ new Map();
|
|
2507
|
+
_destroyed = false;
|
|
2508
|
+
/** The selected backbone precision ("q4" or "f32"). */
|
|
2509
|
+
dtype;
|
|
2510
|
+
audioEmbTables;
|
|
2511
|
+
projW0;
|
|
2512
|
+
projB0;
|
|
2513
|
+
projW2;
|
|
2514
|
+
projB2;
|
|
2515
|
+
audioEmbedScale;
|
|
2516
|
+
headW;
|
|
2517
|
+
headB;
|
|
2518
|
+
stopW;
|
|
2519
|
+
stopB;
|
|
2520
|
+
constructor(ctx, loaded, maxSeqLen, dtype) {
|
|
2521
|
+
this.ctx = ctx;
|
|
2522
|
+
this.loaded = loaded;
|
|
2523
|
+
this.tokenizer = loaded.tokenizer;
|
|
2524
|
+
this.cfg = parseGepardConfig(loaded.gepardConfig);
|
|
2525
|
+
this.maxSeqLen = maxSeqLen;
|
|
2526
|
+
this.hidden = this.cfg.backboneConfig.hidden_size;
|
|
2527
|
+
this.dtype = dtype;
|
|
2528
|
+
const overlay = loaded.overlayWeights;
|
|
2529
|
+
const need = (key) => {
|
|
2530
|
+
const w = overlay.get(key);
|
|
2531
|
+
if (!w) throw new Error(`Gepard overlay weight missing: ${key}`);
|
|
2532
|
+
return w.data;
|
|
2533
|
+
};
|
|
2534
|
+
const n = this.cfg.vocabSizes.length;
|
|
2535
|
+
if (n !== GEPARD_NUM_CHANNELS) throw new Error(`Gepard expects ${GEPARD_NUM_CHANNELS} audio heads, config has ${n}.`);
|
|
2536
|
+
this.audioEmbTables = [];
|
|
2537
|
+
this.headW = [];
|
|
2538
|
+
this.headB = [];
|
|
2539
|
+
for (let i = 0; i < n; i++) {
|
|
2540
|
+
this.audioEmbTables.push(need(`audio_embeddings.${i}.weight`));
|
|
2541
|
+
this.headW.push(need(`codebook_heads.${i}.weight`));
|
|
2542
|
+
this.headB.push(need(`codebook_heads.${i}.bias`));
|
|
2543
|
+
}
|
|
2544
|
+
this.projW0 = need("audio_embed_proj.0.weight");
|
|
2545
|
+
this.projB0 = need("audio_embed_proj.0.bias");
|
|
2546
|
+
this.projW2 = need("audio_embed_proj.2.weight");
|
|
2547
|
+
this.projB2 = need("audio_embed_proj.2.bias");
|
|
2548
|
+
this.audioEmbedScale = need("audio_embed_scale")[0];
|
|
2549
|
+
this.stopW = need("stop_head.weight");
|
|
2550
|
+
this.stopB = need("stop_head.bias")[0];
|
|
2551
|
+
const graph = generateGepardBackboneGraph(this.cfg.backboneConfig, this.dtype === "q4" ? "q4" : void 0);
|
|
2552
|
+
if (this.dtype === "q4") quantizeBackboneInt4(graph, loaded.backboneWeights);
|
|
2553
|
+
this.backboneExec = new Executor(ctx, graph, {
|
|
2554
|
+
maxSeqLen,
|
|
2555
|
+
kvMode: "f32"
|
|
2556
|
+
});
|
|
2557
|
+
this.backboneExec.uploadWeightsMap(selectGraphWeights$2(graph, loaded.backboneWeights));
|
|
2558
|
+
this.backboneExec.initBindGroups();
|
|
2559
|
+
}
|
|
2560
|
+
static async create(options = {}) {
|
|
2561
|
+
const ctx = await initGPU();
|
|
2562
|
+
const loaded = await loadGepardTTS({
|
|
2563
|
+
repo: options.repo,
|
|
2564
|
+
codecRepo: options.codecRepo,
|
|
2565
|
+
revision: options.revision,
|
|
2566
|
+
hfToken: options.hfToken,
|
|
2567
|
+
cacheDir: options.cacheDir,
|
|
2568
|
+
onProgress: options.onProgress
|
|
2569
|
+
});
|
|
2570
|
+
const ctxLen = loaded.rawConfig.max_position_embeddings ?? 262144;
|
|
2571
|
+
return new GepardTTS(ctx, loaded, Math.min(options.maxSeqLen ?? 2048, ctxLen), options.dtype ?? "q4");
|
|
2572
|
+
}
|
|
2573
|
+
/** NanoCodec's fixed output sample rate (Hz). */
|
|
2574
|
+
get sampleRate() {
|
|
2575
|
+
return this.cfg.sampleRate;
|
|
2576
|
+
}
|
|
2577
|
+
/** Preset voices. Gepard ships one default voice (cloning is a follow-up). */
|
|
2578
|
+
availableVoices() {
|
|
2579
|
+
return [];
|
|
2580
|
+
}
|
|
2581
|
+
/** Gather text-token rows from the raw BF16 embed table → f32 [T, hidden]. */
|
|
2582
|
+
embedTextTokens(ids) {
|
|
2583
|
+
const h = this.hidden;
|
|
2584
|
+
const table = this.loaded.embedTokensBF16;
|
|
2585
|
+
const out = new Float32Array(ids.length * h);
|
|
2586
|
+
const view = new DataView(out.buffer);
|
|
2587
|
+
for (let t = 0; t < ids.length; t++) {
|
|
2588
|
+
const rowBase = ids[t] * h;
|
|
2589
|
+
const dstBase = t * h;
|
|
2590
|
+
for (let k = 0; k < h; k++) view.setUint32((dstBase + k) * 4, table[rowBase + k] << 16, true);
|
|
2591
|
+
}
|
|
2592
|
+
return out;
|
|
2593
|
+
}
|
|
2594
|
+
/**
|
|
2595
|
+
* Embed one audio frame (32 FSQ codes) through the overlay stack:
|
|
2596
|
+
* per-channel lookup → concat → Linear→GELU→Linear → affine-free LayerNorm →
|
|
2597
|
+
* × audio_embed_scale. Mirrors GepardModel._embed_audio.
|
|
2598
|
+
*/
|
|
2599
|
+
embedAudioFrame(codes) {
|
|
2600
|
+
const h = this.hidden;
|
|
2601
|
+
const dim = this.cfg.audioEmbedDim;
|
|
2602
|
+
const x = new Float32Array(GEPARD_NUM_CHANNELS * dim);
|
|
2603
|
+
for (let c = 0; c < GEPARD_NUM_CHANNELS; c++) {
|
|
2604
|
+
const table = this.audioEmbTables[c];
|
|
2605
|
+
const base = codes[c] * dim;
|
|
2606
|
+
x.set(table.subarray(base, base + dim), c * dim);
|
|
2607
|
+
}
|
|
2608
|
+
const h1 = new Float32Array(h);
|
|
2609
|
+
linear(this.projW0, this.projB0, x, h, x.length, h1);
|
|
2610
|
+
for (let i = 0; i < h; i++) h1[i] = gelu(h1[i]);
|
|
2611
|
+
const h2 = new Float32Array(h);
|
|
2612
|
+
linear(this.projW2, this.projB2, h1, h, h, h2);
|
|
2613
|
+
let mean = 0;
|
|
2614
|
+
for (let i = 0; i < h; i++) mean += h2[i];
|
|
2615
|
+
mean /= h;
|
|
2616
|
+
let variance = 0;
|
|
2617
|
+
for (let i = 0; i < h; i++) {
|
|
2618
|
+
const d = h2[i] - mean;
|
|
2619
|
+
variance += d * d;
|
|
2620
|
+
}
|
|
2621
|
+
variance /= h;
|
|
2622
|
+
const inv = this.audioEmbedScale / Math.sqrt(variance + LAYERNORM_EPS);
|
|
2623
|
+
for (let i = 0; i < h; i++) h2[i] = (h2[i] - mean) * inv;
|
|
2624
|
+
return h2;
|
|
2625
|
+
}
|
|
2626
|
+
/** Stop-head probability for a hidden state. */
|
|
2627
|
+
stopProbability(hidden) {
|
|
2628
|
+
let acc = this.stopB;
|
|
2629
|
+
for (let k = 0; k < this.hidden; k++) acc += this.stopW[k] * hidden[k];
|
|
2630
|
+
return 1 / (1 + Math.exp(-acc));
|
|
2631
|
+
}
|
|
2632
|
+
/**
|
|
2633
|
+
* Sample one audio frame: all 32 heads independently in fp32 —
|
|
2634
|
+
* windowed repetition penalty → temperature → top-k → softmax → multinomial.
|
|
2635
|
+
* Mirrors GepardRunner._sample_frame / _sample_head_logits.
|
|
2636
|
+
*/
|
|
2637
|
+
sampleFrame(hidden, p, recentFrames) {
|
|
2638
|
+
const frame = new Uint32Array(GEPARD_NUM_CHANNELS);
|
|
2639
|
+
for (let c = 0; c < GEPARD_NUM_CHANNELS; c++) {
|
|
2640
|
+
const L = this.cfg.vocabSizes[c];
|
|
2641
|
+
const w = this.headW[c];
|
|
2642
|
+
const b = this.headB[c];
|
|
2643
|
+
const logits = new Float32Array(L);
|
|
2644
|
+
for (let j = 0; j < L; j++) {
|
|
2645
|
+
let acc = b[j];
|
|
2646
|
+
const base = j * this.hidden;
|
|
2647
|
+
for (let k = 0; k < this.hidden; k++) acc += w[base + k] * hidden[k];
|
|
2648
|
+
logits[j] = acc;
|
|
2649
|
+
}
|
|
2650
|
+
if (p.repetitionPenalty !== 1 && recentFrames) for (const prev of recentFrames) {
|
|
2651
|
+
const tok = prev[c];
|
|
2652
|
+
logits[tok] = logits[tok] > 0 ? logits[tok] / p.repetitionPenalty : logits[tok] * p.repetitionPenalty;
|
|
2653
|
+
}
|
|
2654
|
+
if (p.temperature !== 1) for (let j = 0; j < L; j++) logits[j] /= p.temperature;
|
|
2655
|
+
if (p.topK > 0 && p.topK < L) {
|
|
2656
|
+
const threshold = [...logits].sort((a, b2) => b2 - a)[p.topK - 1];
|
|
2657
|
+
for (let j = 0; j < L; j++) if (logits[j] < threshold) logits[j] = Number.NEGATIVE_INFINITY;
|
|
2658
|
+
}
|
|
2659
|
+
let maxV = Number.NEGATIVE_INFINITY;
|
|
2660
|
+
for (let j = 0; j < L; j++) if (logits[j] > maxV) maxV = logits[j];
|
|
2661
|
+
let sum = 0;
|
|
2662
|
+
for (let j = 0; j < L; j++) {
|
|
2663
|
+
const e = Math.exp(logits[j] - maxV);
|
|
2664
|
+
logits[j] = e;
|
|
2665
|
+
sum += e;
|
|
2666
|
+
}
|
|
2667
|
+
let pick$1 = Math.random() * sum;
|
|
2668
|
+
let chosen = L - 1;
|
|
2669
|
+
for (let j = 0; j < L; j++) {
|
|
2670
|
+
pick$1 -= logits[j];
|
|
2671
|
+
if (pick$1 <= 0) {
|
|
2672
|
+
chosen = j;
|
|
2673
|
+
break;
|
|
2674
|
+
}
|
|
2675
|
+
}
|
|
2676
|
+
frame[c] = chosen;
|
|
2677
|
+
}
|
|
2678
|
+
return frame;
|
|
2679
|
+
}
|
|
2680
|
+
/**
|
|
2681
|
+
* Synthesize speech for `text`. Returns 22.05 kHz mono PCM; when
|
|
2682
|
+
* `opts.onChunk` is set, audio streams out chunk-by-chunk during generation.
|
|
2683
|
+
*
|
|
2684
|
+
* Pipeline: adaptive-repetition text layout → prefill (host-gathered text
|
|
2685
|
+
* embeddings) → frame-synchronous AR loop (one frame = 32 codes per forward,
|
|
2686
|
+
* stop head decides termination) → incremental NanoCodec decode → PCM.
|
|
2687
|
+
*/
|
|
2688
|
+
async speak(text, opts = {}) {
|
|
2689
|
+
if (this._destroyed) throw new Error("GepardTTS has been destroyed.");
|
|
2690
|
+
const temperature = opts.temperature ?? .4;
|
|
2691
|
+
const topK = opts.topK ?? 0;
|
|
2692
|
+
const stopThreshold = opts.stopThreshold ?? .5;
|
|
2693
|
+
const repetitionPenalty = opts.repetitionPenalty ?? 1;
|
|
2694
|
+
const repetitionWindow = opts.repetitionWindow ?? 32;
|
|
2695
|
+
const layout = buildGepardTextLayout(this.tokenizer.encode(text), this.cfg);
|
|
2696
|
+
if (layout.length + 2 >= this.maxSeqLen) throw new Error(`Gepard prompt is ${layout.length} tokens >= maxSeqLen ${this.maxSeqLen}.`);
|
|
2697
|
+
const maxFrames = Math.min(opts.maxFrames ?? DEFAULT_MAX_FRAMES$1, this.maxSeqLen - layout.length - 1);
|
|
2698
|
+
const startTime = performance.now();
|
|
2699
|
+
this.backboneExec.reset();
|
|
2700
|
+
this.backboneExec.writeInput("inputs_embeds", this.embedTextTokens(layout));
|
|
2701
|
+
let { logits: hidden } = await this.backboneExec.forward(new Uint32Array(layout.length));
|
|
2702
|
+
const frames = [this.sampleFrame(hidden, {
|
|
2703
|
+
temperature,
|
|
2704
|
+
topK,
|
|
2705
|
+
repetitionPenalty
|
|
2706
|
+
}, null)];
|
|
2707
|
+
const streamer = new GepardChunkStreamer(this, opts.onChunk);
|
|
2708
|
+
const sampleOpts = {
|
|
2709
|
+
temperature,
|
|
2710
|
+
topK,
|
|
2711
|
+
repetitionPenalty
|
|
2712
|
+
};
|
|
2713
|
+
for (let step = 1; step < maxFrames; step++) {
|
|
2714
|
+
const frameEmbed = this.embedAudioFrame(frames[frames.length - 1]);
|
|
2715
|
+
this.backboneExec.writeInput("inputs_embeds", frameEmbed);
|
|
2716
|
+
hidden = (await this.backboneExec.forward(new Uint32Array(1))).logits;
|
|
2717
|
+
if (this.stopProbability(hidden) > stopThreshold) break;
|
|
2718
|
+
const window = repetitionPenalty !== 1 ? frames.slice(repetitionWindow > 0 ? -repetitionWindow : 0) : null;
|
|
2719
|
+
frames.push(this.sampleFrame(hidden, sampleOpts, window));
|
|
2720
|
+
await streamer.maybeEmit(frames, false);
|
|
2721
|
+
}
|
|
2722
|
+
await streamer.maybeEmit(frames, true);
|
|
2723
|
+
const pcm = streamer.assembled(frames.length);
|
|
2724
|
+
const totalTime = (performance.now() - startTime) / 1e3;
|
|
2725
|
+
const audioSeconds = frames.length * GEPARD_HOP / this.cfg.sampleRate;
|
|
2726
|
+
console.log(`[gepard] frames=${frames.length} audio=${audioSeconds.toFixed(2)}s wall=${totalTime.toFixed(2)}s RTF=${(audioSeconds / totalTime).toFixed(2)}x`);
|
|
2727
|
+
return {
|
|
2728
|
+
pcm,
|
|
2729
|
+
sampleRate: this.cfg.sampleRate,
|
|
2730
|
+
frames: frames.length,
|
|
2731
|
+
audioSeconds
|
|
2732
|
+
};
|
|
2733
|
+
}
|
|
2734
|
+
/**
|
|
2735
|
+
* Decode frames [from, to) of the utterance to PCM, re-decoding up to
|
|
2736
|
+
* `GEPARD_CODEC_LOOKBACK_FRAMES` of left context as warm-up (discarded).
|
|
2737
|
+
* The decoder is causal with a bounded receptive field, so chunked output
|
|
2738
|
+
* matches a monolithic decode.
|
|
2739
|
+
*/
|
|
2740
|
+
async decodeFrameRange(frames, from, to) {
|
|
2741
|
+
const ctxStart = Math.max(0, from - GEPARD_CODEC_LOOKBACK_FRAMES);
|
|
2742
|
+
const winFrames = to - ctxStart;
|
|
2743
|
+
const codes = new Uint32Array(GEPARD_NUM_CHANNELS * winFrames);
|
|
2744
|
+
for (let t = 0; t < winFrames; t++) {
|
|
2745
|
+
const frame = frames[ctxStart + t];
|
|
2746
|
+
for (let c = 0; c < GEPARD_NUM_CHANNELS; c++) codes[c * winFrames + t] = frame[c];
|
|
2747
|
+
}
|
|
2748
|
+
const latent = gepardCodesToLatent(codes, winFrames);
|
|
2749
|
+
const win = await this.decodeLatentWindow(latent, winFrames);
|
|
2750
|
+
const drop = (from - ctxStart) * GEPARD_HOP;
|
|
2751
|
+
const keep = (to - from) * GEPARD_HOP;
|
|
2752
|
+
const out = win.slice(drop, drop + keep);
|
|
2753
|
+
for (let i = 0; i < out.length; i++) out[i] = Math.min(1, Math.max(-1, out[i]));
|
|
2754
|
+
return out;
|
|
2755
|
+
}
|
|
2756
|
+
/** Run the codec graph for one latent window → raw PCM (pre-clamp). */
|
|
2757
|
+
async decodeLatentWindow(latent, winFrames) {
|
|
2758
|
+
const steady = winFrames === GEPARD_CODEC_CHUNK_FRAMES + GEPARD_CODEC_LOOKBACK_FRAMES;
|
|
2759
|
+
let exec = steady ? this.codecExecCache.get(winFrames) : void 0;
|
|
2760
|
+
let graphOutputs;
|
|
2761
|
+
if (exec) graphOutputs = exec.graphOutputs;
|
|
2762
|
+
else {
|
|
2763
|
+
const graph = generateNanoCodecDecoderGraph({
|
|
2764
|
+
numFrames: winFrames,
|
|
2765
|
+
geometry: {
|
|
2766
|
+
upSampleRates: GEPARD_CODEC_UP_SAMPLE_RATES,
|
|
2767
|
+
inputDim: GEPARD_CODEC_INPUT_DIM
|
|
2768
|
+
},
|
|
2769
|
+
latentInput: true
|
|
2770
|
+
});
|
|
2771
|
+
exec = new Executor(this.ctx, graph, {
|
|
2772
|
+
maxSeqLen: winFrames,
|
|
2773
|
+
kvMode: "f32"
|
|
2774
|
+
});
|
|
2775
|
+
exec.uploadWeightsMap(selectGraphWeights$2(graph, this.loaded.codecWeights));
|
|
2776
|
+
exec.initBindGroups();
|
|
2777
|
+
graphOutputs = graph.outputs;
|
|
2778
|
+
if (steady) this.codecExecCache.set(winFrames, exec);
|
|
2779
|
+
}
|
|
2780
|
+
try {
|
|
2781
|
+
exec.reset();
|
|
2782
|
+
exec.writeInput("nc_latent", latent);
|
|
2783
|
+
return await exec.runGraphOutput(graphOutputs[0], winFrames * GEPARD_HOP);
|
|
2784
|
+
} finally {
|
|
2785
|
+
if (!steady) exec.destroy();
|
|
2786
|
+
}
|
|
2787
|
+
}
|
|
2788
|
+
destroy() {
|
|
2789
|
+
if (this._destroyed) return;
|
|
2790
|
+
this._destroyed = true;
|
|
2791
|
+
this.backboneExec.destroy();
|
|
2792
|
+
for (const exec of this.codecExecCache.values()) exec.destroy();
|
|
2793
|
+
this.codecExecCache.clear();
|
|
2794
|
+
this.loaded.backboneWeights.clear();
|
|
2795
|
+
this.loaded.overlayWeights.clear();
|
|
2796
|
+
this.loaded.codecWeights.clear();
|
|
1183
2797
|
}
|
|
1184
|
-
|
|
1185
|
-
}
|
|
1186
|
-
/**
|
|
1187
|
-
* Cap `text` to at most `max` characters, cutting on a word boundary when a
|
|
1188
|
-
* reasonable one exists in the back half of the window.
|
|
1189
|
-
*/
|
|
1190
|
-
function capLength(text, max) {
|
|
1191
|
-
if (text.length <= max) return text;
|
|
1192
|
-
const cut = text.slice(0, max);
|
|
1193
|
-
const lastSpace = cut.lastIndexOf(" ");
|
|
1194
|
-
return (lastSpace > max * .5 ? cut.slice(0, lastSpace) : cut).trimEnd();
|
|
1195
|
-
}
|
|
2798
|
+
};
|
|
1196
2799
|
/**
|
|
1197
|
-
*
|
|
1198
|
-
*
|
|
1199
|
-
*
|
|
1200
|
-
* 3. drop a full verbatim echo of the typed text,
|
|
1201
|
-
* 4. drop a leading character overlap with the typed tail,
|
|
1202
|
-
* 5. drop a leading word-run that restates an earlier typed phrase,
|
|
1203
|
-
* 6. truncate an internal phrase loop,
|
|
1204
|
-
* 7. cap the length,
|
|
1205
|
-
* 8. add a single smart leading space so the ghost joins the caret naturally.
|
|
1206
|
-
*
|
|
1207
|
-
* Returns "" when nothing novel is left — the ghost then simply shows nothing,
|
|
1208
|
-
* which is the correct behavior for a suggestion that only repeats the input.
|
|
2800
|
+
* Incremental chunk decoder: every GEPARD_CODEC_CHUNK_FRAMES completed frames
|
|
2801
|
+
* it decodes the new span (with lookback) and emits it via `onChunk`, keeping
|
|
2802
|
+
* all emitted PCM so the final result is assembled without a second decode.
|
|
1209
2803
|
*/
|
|
1210
|
-
|
|
1211
|
-
|
|
1212
|
-
|
|
1213
|
-
|
|
1214
|
-
|
|
1215
|
-
|
|
1216
|
-
|
|
1217
|
-
|
|
1218
|
-
|
|
1219
|
-
|
|
1220
|
-
|
|
1221
|
-
|
|
1222
|
-
|
|
1223
|
-
|
|
1224
|
-
|
|
1225
|
-
|
|
1226
|
-
|
|
2804
|
+
var GepardChunkStreamer = class {
|
|
2805
|
+
tts;
|
|
2806
|
+
onChunk;
|
|
2807
|
+
decodedUpTo = 0;
|
|
2808
|
+
chunks = [];
|
|
2809
|
+
constructor(tts, onChunk) {
|
|
2810
|
+
this.tts = tts;
|
|
2811
|
+
this.onChunk = onChunk;
|
|
2812
|
+
}
|
|
2813
|
+
/** Decode + emit any complete pending chunk (or everything when `final`). */
|
|
2814
|
+
async maybeEmit(frames, final) {
|
|
2815
|
+
if (!(final || this.onChunk)) return;
|
|
2816
|
+
while (frames.length - this.decodedUpTo >= GEPARD_CODEC_CHUNK_FRAMES || final) {
|
|
2817
|
+
const from = this.decodedUpTo;
|
|
2818
|
+
const to = final ? frames.length : Math.min(frames.length, from + GEPARD_CODEC_CHUNK_FRAMES);
|
|
2819
|
+
if (to <= from) break;
|
|
2820
|
+
const pcm = await this.tts.decodeFrameRange(frames, from, to);
|
|
2821
|
+
this.chunks.push(pcm);
|
|
2822
|
+
this.decodedUpTo = to;
|
|
2823
|
+
this.onChunk?.({
|
|
2824
|
+
pcm,
|
|
2825
|
+
sampleRate: this.tts.sampleRate,
|
|
2826
|
+
frameIndex: from,
|
|
2827
|
+
isFinal: final && this.decodedUpTo >= frames.length
|
|
2828
|
+
});
|
|
2829
|
+
if (final && this.decodedUpTo >= frames.length) break;
|
|
2830
|
+
}
|
|
2831
|
+
}
|
|
2832
|
+
/** Concatenate all emitted chunks into the full utterance PCM. */
|
|
2833
|
+
assembled(totalFrames) {
|
|
2834
|
+
const pcm = new Float32Array(totalFrames * GEPARD_HOP);
|
|
2835
|
+
let offset = 0;
|
|
2836
|
+
for (const chunk of this.chunks) {
|
|
2837
|
+
pcm.set(chunk, offset);
|
|
2838
|
+
offset += chunk.length;
|
|
2839
|
+
}
|
|
2840
|
+
return pcm;
|
|
2841
|
+
}
|
|
2842
|
+
};
|
|
1227
2843
|
|
|
1228
2844
|
//#endregion
|
|
1229
2845
|
//#region src/gpu/architectures/outetts.ts
|
|
@@ -2859,7 +4475,7 @@ function selectGraphWeights$1(graph, weights) {
|
|
|
2859
4475
|
return out;
|
|
2860
4476
|
}
|
|
2861
4477
|
/** Index of the maximum logit (greedy / argmax decode). */
|
|
2862
|
-
function argmax$
|
|
4478
|
+
function argmax$1(logits) {
|
|
2863
4479
|
let best = 0;
|
|
2864
4480
|
let bestV = logits[0];
|
|
2865
4481
|
for (let i = 1; i < logits.length; i++) if (logits[i] > bestV) {
|
|
@@ -3188,7 +4804,7 @@ var OuteTTS = class OuteTTS {
|
|
|
3188
4804
|
const generated = [];
|
|
3189
4805
|
let logits = prefillLogits;
|
|
3190
4806
|
for (let step = 0; step < p.maxNewTokens; step++) {
|
|
3191
|
-
const next = p.greedy ? argmax$
|
|
4807
|
+
const next = p.greedy ? argmax$1(logits) : this.sampleToken(logits, {
|
|
3192
4808
|
temperature: p.temperature,
|
|
3193
4809
|
topP: p.topP,
|
|
3194
4810
|
topK: p.topK,
|
|
@@ -3320,7 +4936,7 @@ function selectGraphWeights(graph, weights) {
|
|
|
3320
4936
|
return out;
|
|
3321
4937
|
}
|
|
3322
4938
|
/** Argmax over a logit row. */
|
|
3323
|
-
function argmax
|
|
4939
|
+
function argmax(logits) {
|
|
3324
4940
|
let best = 0;
|
|
3325
4941
|
let bestV = logits[0];
|
|
3326
4942
|
for (let i = 1; i < logits.length; i++) if (logits[i] > bestV) {
|
|
@@ -3671,7 +5287,7 @@ var ParlerTTS = class ParlerTTS {
|
|
|
3671
5287
|
continue;
|
|
3672
5288
|
}
|
|
3673
5289
|
if (cb > frontier) logitRows[cb][PARLER_EOS_TOKEN_ID] = Number.NEGATIVE_INFINITY;
|
|
3674
|
-
const code = opts.greedy ? argmax
|
|
5290
|
+
const code = opts.greedy ? argmax(logitRows[cb]) : sampleLogits(logitRows[cb], opts.temperature ?? 1, opts.topK ?? 0);
|
|
3675
5291
|
nextCodes[cb] = code;
|
|
3676
5292
|
grid[cb * maxSteps + step] = code < PARLER_CODEBOOK_SIZE ? code : 0;
|
|
3677
5293
|
}
|
|
@@ -3742,122 +5358,6 @@ function addPosition(row, posTable, pos, H) {
|
|
|
3742
5358
|
for (let i = 0; i < H; i++) row[i] += posTable[base + i];
|
|
3743
5359
|
}
|
|
3744
5360
|
|
|
3745
|
-
//#endregion
|
|
3746
|
-
//#region src/gpu/sampler.ts
|
|
3747
|
-
let _heapIndices = null;
|
|
3748
|
-
let _heapValues = null;
|
|
3749
|
-
/**
|
|
3750
|
-
* Sample a token ID from logits.
|
|
3751
|
-
*
|
|
3752
|
-
* Pipeline: repetition penalty → temperature → top-k (min-heap) → softmax → top-p → sample.
|
|
3753
|
-
*/
|
|
3754
|
-
function sampleToken(logits, params = {}, previousTokens) {
|
|
3755
|
-
const temperature = params.temperature ?? .7;
|
|
3756
|
-
const topK = params.topK ?? 50;
|
|
3757
|
-
const topP = params.topP ?? .9;
|
|
3758
|
-
const repetitionPenalty = params.repetitionPenalty ?? 1;
|
|
3759
|
-
if (temperature < 1e-6) return argmax(logits);
|
|
3760
|
-
const N = logits.length;
|
|
3761
|
-
const K = Math.min(topK > 0 ? topK : N, N);
|
|
3762
|
-
if (!_heapIndices || _heapIndices.length < K) {
|
|
3763
|
-
_heapIndices = new Uint32Array(K);
|
|
3764
|
-
_heapValues = new Float32Array(K);
|
|
3765
|
-
}
|
|
3766
|
-
const hIdx = _heapIndices;
|
|
3767
|
-
const hVal = _heapValues;
|
|
3768
|
-
let penaltySet = null;
|
|
3769
|
-
if (repetitionPenalty !== 1 && previousTokens?.length) penaltySet = new Set(previousTokens);
|
|
3770
|
-
let heapSize = 0;
|
|
3771
|
-
for (let i = 0; i < N; i++) {
|
|
3772
|
-
let s = logits[i];
|
|
3773
|
-
if (penaltySet?.has(i)) s = s > 0 ? s / repetitionPenalty : s * repetitionPenalty;
|
|
3774
|
-
s /= temperature;
|
|
3775
|
-
if (heapSize < K) {
|
|
3776
|
-
hIdx[heapSize] = i;
|
|
3777
|
-
hVal[heapSize] = s;
|
|
3778
|
-
heapSize++;
|
|
3779
|
-
if (heapSize === K) for (let j = (K >> 1) - 1; j >= 0; j--) siftDown(hIdx, hVal, j, K);
|
|
3780
|
-
} else if (s > hVal[0]) {
|
|
3781
|
-
hIdx[0] = i;
|
|
3782
|
-
hVal[0] = s;
|
|
3783
|
-
siftDown(hIdx, hVal, 0, K);
|
|
3784
|
-
}
|
|
3785
|
-
}
|
|
3786
|
-
for (let i = 1; i < heapSize; i++) {
|
|
3787
|
-
const vi = hVal[i];
|
|
3788
|
-
const ii = hIdx[i];
|
|
3789
|
-
let j = i - 1;
|
|
3790
|
-
while (j >= 0 && hVal[j] < vi) {
|
|
3791
|
-
hVal[j + 1] = hVal[j];
|
|
3792
|
-
hIdx[j + 1] = hIdx[j];
|
|
3793
|
-
j--;
|
|
3794
|
-
}
|
|
3795
|
-
hVal[j + 1] = vi;
|
|
3796
|
-
hIdx[j + 1] = ii;
|
|
3797
|
-
}
|
|
3798
|
-
const maxScore = hVal[0];
|
|
3799
|
-
let sumExp = 0;
|
|
3800
|
-
for (let i = 0; i < heapSize; i++) {
|
|
3801
|
-
const p = Math.exp(hVal[i] - maxScore);
|
|
3802
|
-
hVal[i] = p;
|
|
3803
|
-
sumExp += p;
|
|
3804
|
-
}
|
|
3805
|
-
const invSum = 1 / sumExp;
|
|
3806
|
-
for (let i = 0; i < heapSize; i++) hVal[i] *= invSum;
|
|
3807
|
-
let candidateCount = heapSize;
|
|
3808
|
-
if (topP < 1) {
|
|
3809
|
-
let cumulative$1 = 0;
|
|
3810
|
-
for (let i = 0; i < heapSize; i++) {
|
|
3811
|
-
cumulative$1 += hVal[i];
|
|
3812
|
-
if (cumulative$1 >= topP) {
|
|
3813
|
-
candidateCount = i + 1;
|
|
3814
|
-
break;
|
|
3815
|
-
}
|
|
3816
|
-
}
|
|
3817
|
-
let sum = 0;
|
|
3818
|
-
for (let i = 0; i < candidateCount; i++) sum += hVal[i];
|
|
3819
|
-
const inv = 1 / sum;
|
|
3820
|
-
for (let i = 0; i < candidateCount; i++) hVal[i] *= inv;
|
|
3821
|
-
}
|
|
3822
|
-
const r = Math.random();
|
|
3823
|
-
let cumulative = 0;
|
|
3824
|
-
for (let i = 0; i < candidateCount; i++) {
|
|
3825
|
-
cumulative += hVal[i];
|
|
3826
|
-
if (r <= cumulative) return hIdx[i];
|
|
3827
|
-
}
|
|
3828
|
-
return hIdx[candidateCount - 1];
|
|
3829
|
-
}
|
|
3830
|
-
/**
|
|
3831
|
-
* Return the index of the maximum value (greedy decoding).
|
|
3832
|
-
*/
|
|
3833
|
-
function argmax(arr) {
|
|
3834
|
-
let maxIdx = 0;
|
|
3835
|
-
let maxVal = arr[0];
|
|
3836
|
-
for (let i = 1; i < arr.length; i++) if (arr[i] > maxVal) {
|
|
3837
|
-
maxVal = arr[i];
|
|
3838
|
-
maxIdx = i;
|
|
3839
|
-
}
|
|
3840
|
-
return maxIdx;
|
|
3841
|
-
}
|
|
3842
|
-
/** Min-heap sift down on parallel index/value typed arrays. */
|
|
3843
|
-
function siftDown(indices, values, i, n) {
|
|
3844
|
-
while (true) {
|
|
3845
|
-
let smallest = i;
|
|
3846
|
-
const left = 2 * i + 1;
|
|
3847
|
-
const right = 2 * i + 2;
|
|
3848
|
-
if (left < n && values[left] < values[smallest]) smallest = left;
|
|
3849
|
-
if (right < n && values[right] < values[smallest]) smallest = right;
|
|
3850
|
-
if (smallest === i) break;
|
|
3851
|
-
const ti = indices[i];
|
|
3852
|
-
indices[i] = indices[smallest];
|
|
3853
|
-
indices[smallest] = ti;
|
|
3854
|
-
const tv = values[i];
|
|
3855
|
-
values[i] = values[smallest];
|
|
3856
|
-
values[smallest] = tv;
|
|
3857
|
-
i = smallest;
|
|
3858
|
-
}
|
|
3859
|
-
}
|
|
3860
|
-
|
|
3861
5361
|
//#endregion
|
|
3862
5362
|
//#region src/gpu/vision-executor.ts
|
|
3863
5363
|
const MAP_MODE_READ = 1;
|
|
@@ -5143,8 +6643,14 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
5143
6643
|
_outeVoice = null;
|
|
5144
6644
|
/** Source of the runtime LoRA adapter currently applied on the static base, if any. */
|
|
5145
6645
|
_currentAdapter = null;
|
|
6646
|
+
/** Named adapters registered for per-lane batched decode (GERBIL_BATCH_LORA). */
|
|
6647
|
+
_batchAdapters = /* @__PURE__ */ new Map();
|
|
5146
6648
|
/** Lazily-created Parler-TTS engine (Flan-T5 encoder + decoder LM + dac_44khz). */
|
|
5147
6649
|
_parlerTTS = null;
|
|
6650
|
+
/** Lazily-created Gepard-1.0 engine (Qwen3.5 multihead codec-LM + NanoCodec 21.5 fps). */
|
|
6651
|
+
_gepardTTS = null;
|
|
6652
|
+
/** Lazily-created continuous-batching scheduler (GERBIL_BATCH, Phase 5). */
|
|
6653
|
+
_scheduler = null;
|
|
5148
6654
|
/**
|
|
5149
6655
|
* WebKit group-size probe state. When true, a candidate group size is being
|
|
5150
6656
|
* tried this page-load and must be promoted (or capped) after the FIRST
|
|
@@ -5250,6 +6756,35 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
5250
6756
|
return this._currentAdapter;
|
|
5251
6757
|
}
|
|
5252
6758
|
/**
|
|
6759
|
+
* Register a named LoRA adapter for per-lane batched decode (SGMV phase 1,
|
|
6760
|
+
* `GERBIL_BATCH=N` + `GERBIL_BATCH_LORA=1`). The adapter's factors are
|
|
6761
|
+
* fetched and packed into the executor's shared factor buffers ONCE; any
|
|
6762
|
+
* number of batch lanes can then run it concurrently by passing its name in
|
|
6763
|
+
* `generateBatch`'s per-lane `adapters` option. Registering the same name
|
|
6764
|
+
* twice is a no-op.
|
|
6765
|
+
*/
|
|
6766
|
+
async registerAdapter(name, source) {
|
|
6767
|
+
this.checkDestroyed();
|
|
6768
|
+
if (!this.executor.batchLoraEnabled) throw new Error("registerAdapter requires GERBIL_BATCH=N (N >= 2) and GERBIL_BATCH_LORA=1 on the Dawn path");
|
|
6769
|
+
if (this._batchAdapters.has(name)) return;
|
|
6770
|
+
const adapter = await fetchAdapter(source, {
|
|
6771
|
+
hfToken: this._createOptions.hfToken,
|
|
6772
|
+
revision: this._createOptions.revision
|
|
6773
|
+
});
|
|
6774
|
+
if (!adapter) throw new Error(`No adapter found at ${source}.`);
|
|
6775
|
+
const deltas = buildLoRADeltas(adapter, createKeyMapperForArch(this._architecture));
|
|
6776
|
+
const id = this.executor.registerBatchAdapter(deltas);
|
|
6777
|
+
this._batchAdapters.set(name, {
|
|
6778
|
+
id,
|
|
6779
|
+
deltas
|
|
6780
|
+
});
|
|
6781
|
+
console.log(`[engine] batch adapter "${name}" (${source}) registered as id ${id}.`);
|
|
6782
|
+
}
|
|
6783
|
+
/** Names of the adapters registered for per-lane batched decode. */
|
|
6784
|
+
getRegisteredAdapters() {
|
|
6785
|
+
return [...this._batchAdapters.keys()];
|
|
6786
|
+
}
|
|
6787
|
+
/**
|
|
5253
6788
|
* Write a coarse crash-phase breadcrumb that survives a GPU-process kill / page
|
|
5254
6789
|
* reload. The iPad harness reads `localStorage["gerbil-crash-phase"]` after a
|
|
5255
6790
|
* crash; without these, a describe-time crash only shows the last load phase
|
|
@@ -5580,6 +7115,186 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
5580
7115
|
}
|
|
5581
7116
|
}
|
|
5582
7117
|
/**
|
|
7118
|
+
* Batched lockstep generation (GERBIL_BATCH=N — Phase 2 of the
|
|
7119
|
+
* continuous-batching campaign, Dawn only).
|
|
7120
|
+
*
|
|
7121
|
+
* Prefills each prompt sequentially into its batch slot (KV into the slot's
|
|
7122
|
+
* block-table row, SSM/conv state adopted into the slot's state-pool rows),
|
|
7123
|
+
* then decodes all N sequences in lockstep: one batched dispatch stream per
|
|
7124
|
+
* step produces one token per row. Greedy-only in Phase 2 (per-row GPU
|
|
7125
|
+
* argmax); rows that hit EOS/stop keep decoding in lockstep with their
|
|
7126
|
+
* output discarded until every row has finished — rows never interact, so
|
|
7127
|
+
* per-row streams are token-exact vs single-sequence generation.
|
|
7128
|
+
*
|
|
7129
|
+
* Requires exactly `executor.batchSize` prompts (fixed lockstep batch; the
|
|
7130
|
+
* Phase 3 scheduler will lift this).
|
|
7131
|
+
*/
|
|
7132
|
+
async generateBatch(prompts, options = {}) {
|
|
7133
|
+
this.checkDestroyed();
|
|
7134
|
+
const B = this.executor.batchSize;
|
|
7135
|
+
if (B <= 0) throw new Error("generateBatch requires GERBIL_BATCH=N (N >= 2) on the Dawn path");
|
|
7136
|
+
if (prompts.length !== B) throw new Error(`generateBatch: got ${prompts.length} prompts for a fixed batch of ${B} (set GERBIL_BATCH=${prompts.length})`);
|
|
7137
|
+
if (this._multimodalGraph || this.executor.hasPleSource() || this.executor.hasRuntimeAdapter) throw new Error("generateBatch: multimodal, PLE, and runtime-LoRA models are Phase 2+ work");
|
|
7138
|
+
const { maxTokens = 512, stopSequences = [], sampling = {}, systemPrompt } = options;
|
|
7139
|
+
if ((sampling.temperature ?? .7) >= 1e-6) throw new Error("generateBatch: Phase 2 is greedy-only — pass sampling: { temperature: 0 }");
|
|
7140
|
+
const laneAdapterNames = options.adapters;
|
|
7141
|
+
if (laneAdapterNames && laneAdapterNames.length !== B) throw new Error(`generateBatch: got ${laneAdapterNames.length} lane adapters for a batch of ${B}`);
|
|
7142
|
+
if ((laneAdapterNames?.some((a) => a != null) ?? false) && !this.executor.batchLoraEnabled) throw new Error("generateBatch: per-lane adapters require GERBIL_BATCH_LORA=1");
|
|
7143
|
+
const laneAdapters = new Array(B).fill(null);
|
|
7144
|
+
if (laneAdapterNames) for (let i = 0; i < B; i++) {
|
|
7145
|
+
const name = laneAdapterNames[i];
|
|
7146
|
+
if (name == null) continue;
|
|
7147
|
+
const reg = this._batchAdapters.get(name);
|
|
7148
|
+
if (!reg) throw new Error(`generateBatch: adapter "${name}" not registered (call registerAdapter first)`);
|
|
7149
|
+
laneAdapters[i] = reg;
|
|
7150
|
+
}
|
|
7151
|
+
const release = await this._acquireGenLock();
|
|
7152
|
+
try {
|
|
7153
|
+
const startTime = performance.now();
|
|
7154
|
+
this.executor.resetBatch();
|
|
7155
|
+
const eosId = this.tokenizer.config.eosTokenId;
|
|
7156
|
+
const eotId = this.resolveEndOfTurnId();
|
|
7157
|
+
const keepPlan = this.config.vocabKeepPlan;
|
|
7158
|
+
const remapTok = keepPlan ? (i) => remapPrunedToken(i, keepPlan) : (i) => i;
|
|
7159
|
+
const rows = prompts.map((_p, i) => ({
|
|
7160
|
+
generatedIds: [],
|
|
7161
|
+
text: "",
|
|
7162
|
+
finished: false,
|
|
7163
|
+
finishReason: "max_tokens",
|
|
7164
|
+
limit: options.maxTokensPerRow?.[i] ?? maxTokens
|
|
7165
|
+
}));
|
|
7166
|
+
const firstTokens = new Uint32Array(B);
|
|
7167
|
+
for (let i = 0; i < B; i++) {
|
|
7168
|
+
const p = prompts[i];
|
|
7169
|
+
const messages = typeof p === "string" ? [...systemPrompt ? [{
|
|
7170
|
+
role: "system",
|
|
7171
|
+
content: systemPrompt
|
|
7172
|
+
}] : [], {
|
|
7173
|
+
role: "user",
|
|
7174
|
+
content: p
|
|
7175
|
+
}] : systemPrompt ? [{
|
|
7176
|
+
role: "system",
|
|
7177
|
+
content: systemPrompt
|
|
7178
|
+
}, ...p] : p;
|
|
7179
|
+
const inputIds = this.tokenizer.encodeChat(messages, { addGenerationPrompt: true });
|
|
7180
|
+
const laneReg = laneAdapters[i];
|
|
7181
|
+
if (laneReg) this.executor.applyRuntimeLoRA(laneReg.deltas);
|
|
7182
|
+
try {
|
|
7183
|
+
this.executor.beginBatchSlot(i);
|
|
7184
|
+
const { logits } = await this.executor.forward(new Uint32Array(inputIds));
|
|
7185
|
+
this.executor.finishBatchSlot(i);
|
|
7186
|
+
firstTokens[i] = remapTok(sampleToken(logits, { temperature: 0 }));
|
|
7187
|
+
} finally {
|
|
7188
|
+
if (laneReg) this.executor.clearRuntimeLoRA();
|
|
7189
|
+
}
|
|
7190
|
+
}
|
|
7191
|
+
if (this.executor.batchLoraEnabled) this.executor.setBatchLaneAdapters(laneAdapters.map((a) => a ? a.id : -1));
|
|
7192
|
+
const consume = (row, tok) => {
|
|
7193
|
+
if (row.finished) return;
|
|
7194
|
+
row.generatedIds.push(tok);
|
|
7195
|
+
if (eosId !== null && tok === eosId || eotId !== null && tok === eotId) {
|
|
7196
|
+
row.finished = true;
|
|
7197
|
+
row.finishReason = "eos";
|
|
7198
|
+
return;
|
|
7199
|
+
}
|
|
7200
|
+
row.text += this.tokenizer.decode([tok], true);
|
|
7201
|
+
if (stopSequences.some((s) => row.text.includes(s))) {
|
|
7202
|
+
for (const s of stopSequences) {
|
|
7203
|
+
const idx = row.text.indexOf(s);
|
|
7204
|
+
if (idx !== -1) row.text = row.text.slice(0, idx);
|
|
7205
|
+
}
|
|
7206
|
+
row.finished = true;
|
|
7207
|
+
row.finishReason = "stop_sequence";
|
|
7208
|
+
return;
|
|
7209
|
+
}
|
|
7210
|
+
if (row.generatedIds.length >= row.limit) row.finished = true;
|
|
7211
|
+
};
|
|
7212
|
+
const decodeStart = performance.now();
|
|
7213
|
+
for (let i = 0; i < B; i++) consume(rows[i], firstTokens[i]);
|
|
7214
|
+
const depth = Executor.PIPELINE_DEPTH;
|
|
7215
|
+
const lens = this.executor.batchSlotLengths;
|
|
7216
|
+
let capacity = Number.POSITIVE_INFINITY;
|
|
7217
|
+
for (let i = 0; i < B; i++) capacity = Math.min(capacity, this.maxSeqLen - lens[i] - 1);
|
|
7218
|
+
const maxRowLimit = rows.reduce((a, r) => Math.max(a, r.limit), 0);
|
|
7219
|
+
const stepsNeeded = Math.max(0, Math.min(maxRowLimit - 1, capacity));
|
|
7220
|
+
let submitted = 0;
|
|
7221
|
+
let consumedSteps = 0;
|
|
7222
|
+
while (consumedSteps < stepsNeeded && !rows.every((r) => r.finished)) {
|
|
7223
|
+
while (submitted < stepsNeeded && submitted < consumedSteps + depth) {
|
|
7224
|
+
this.executor.submitBatchDecodeStep(submitted === 0 ? firstTokens : null, submitted % depth);
|
|
7225
|
+
submitted++;
|
|
7226
|
+
}
|
|
7227
|
+
const toks = await this.executor.readBatchTokens(consumedSteps % depth);
|
|
7228
|
+
consumedSteps++;
|
|
7229
|
+
for (let i = 0; i < B; i++) consume(rows[i], toks[i]);
|
|
7230
|
+
}
|
|
7231
|
+
const decodeTime = performance.now() - decodeStart;
|
|
7232
|
+
const totalTime = performance.now() - startTime;
|
|
7233
|
+
return rows.map((r) => ({
|
|
7234
|
+
text: r.text,
|
|
7235
|
+
tokensGenerated: r.generatedIds.length,
|
|
7236
|
+
tokensPerSecond: r.generatedIds.length / (totalTime / 1e3),
|
|
7237
|
+
totalTime,
|
|
7238
|
+
decodeTime,
|
|
7239
|
+
finishReason: r.finishReason,
|
|
7240
|
+
tokenIds: [...r.generatedIds]
|
|
7241
|
+
}));
|
|
7242
|
+
} finally {
|
|
7243
|
+
release();
|
|
7244
|
+
}
|
|
7245
|
+
}
|
|
7246
|
+
/**
|
|
7247
|
+
* Continuous-batching scheduler (GERBIL_BATCH=N — Phase 5 of the batching
|
|
7248
|
+
* campaign, Dawn only). Returns the engine's request scheduler (created on
|
|
7249
|
+
* first call): submit many requests with independent prompts/limits and the
|
|
7250
|
+
* scheduler multiplexes them over the batch lanes with dynamic admission,
|
|
7251
|
+
* immediate lane retirement, and chunked prefill. Greedy-only; per-request
|
|
7252
|
+
* output is token-exact vs single-sequence generate().
|
|
7253
|
+
*/
|
|
7254
|
+
createScheduler() {
|
|
7255
|
+
this.checkDestroyed();
|
|
7256
|
+
if (this._scheduler) return this._scheduler;
|
|
7257
|
+
if (this.executor.batchSize <= 0) throw new Error("createScheduler requires GERBIL_BATCH=N (N >= 2) on the Dawn path");
|
|
7258
|
+
if (this._multimodalGraph || this.executor.hasPleSource() || this.executor.hasRuntimeAdapter) throw new Error("createScheduler: multimodal, PLE, and runtime-LoRA models are unsupported");
|
|
7259
|
+
const keepPlan = this.config.vocabKeepPlan;
|
|
7260
|
+
this._scheduler = new RequestScheduler({
|
|
7261
|
+
executor: this.executor,
|
|
7262
|
+
encodeChat: (messages) => this.tokenizer.encodeChat(messages, { addGenerationPrompt: true }),
|
|
7263
|
+
decodeToken: (id) => this.tokenizer.decode([id], true),
|
|
7264
|
+
eosTokenId: this.tokenizer.config.eosTokenId,
|
|
7265
|
+
eotTokenId: this.resolveEndOfTurnId(),
|
|
7266
|
+
remapToken: keepPlan ? (i) => remapPrunedToken(i, keepPlan) : (i) => i,
|
|
7267
|
+
resolveAdapter: (name) => {
|
|
7268
|
+
const reg = this._batchAdapters.get(name);
|
|
7269
|
+
if (!reg) throw new Error(`scheduler: adapter "${name}" not registered (call registerAdapter first)`);
|
|
7270
|
+
return reg;
|
|
7271
|
+
},
|
|
7272
|
+
acquireLock: () => this._acquireGenLock()
|
|
7273
|
+
});
|
|
7274
|
+
return this._scheduler;
|
|
7275
|
+
}
|
|
7276
|
+
/**
|
|
7277
|
+
* Create a swarm — the MoTA mode-B primitive (docs/research/slm-swarms.md):
|
|
7278
|
+
* named members (LoRA adapters on this engine's shared base, or null for the
|
|
7279
|
+
* bare base), an optional plan that fans an input out into per-member
|
|
7280
|
+
* sub-tasks, and an optional reduce that joins the results. `swarm.run()`
|
|
7281
|
+
* decodes every sub-task CONCURRENTLY through the continuous-batching
|
|
7282
|
+
* scheduler with per-request adapters — one batched pass, so the fan-out
|
|
7283
|
+
* costs roughly the slowest member, not the sum.
|
|
7284
|
+
*
|
|
7285
|
+
* Every non-null member source is registered via {@link registerAdapter}
|
|
7286
|
+
* (downloads in parallel; factors packed into shared GPU buffers once).
|
|
7287
|
+
* Requires `GERBIL_BATCH=N` (N >= 2) + `GERBIL_BATCH_LORA=1` on the Dawn
|
|
7288
|
+
* path; greedy-only like the scheduler underneath.
|
|
7289
|
+
*/
|
|
7290
|
+
async createSwarm(options) {
|
|
7291
|
+
this.checkDestroyed();
|
|
7292
|
+
if (!this.executor.batchLoraEnabled) throw new Error("createSwarm requires GERBIL_BATCH=N (N >= 2) and GERBIL_BATCH_LORA=1 on the Dawn path");
|
|
7293
|
+
const scheduler = this.createScheduler();
|
|
7294
|
+
await Promise.all(Object.entries(options.members).map(([name, source]) => source == null ? Promise.resolve() : this.registerAdapter(name, source)));
|
|
7295
|
+
return new Swarm(scheduler, options);
|
|
7296
|
+
}
|
|
7297
|
+
/**
|
|
5583
7298
|
* Resolve the end-of-turn stop token id, or null if the model has none.
|
|
5584
7299
|
*
|
|
5585
7300
|
* Chat models like Gemma 4 end an assistant turn with a dedicated end-of-turn
|
|
@@ -5665,18 +7380,34 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
5665
7380
|
if (isGreedy && !this.executor.needsMultiEncoder && !mmDecode && !streamsPle) {
|
|
5666
7381
|
const firstToken = remapTok(sampleToken(logits, sampling, [...inputIds, ...generatedIds]));
|
|
5667
7382
|
if (!consumeToken(firstToken)) {
|
|
5668
|
-
const depth = Executor.PIPELINE_DEPTH;
|
|
5669
7383
|
const stepsNeeded = Math.min(maxTokens - 1, this.executor.decodeCapacityRemaining());
|
|
5670
|
-
|
|
5671
|
-
|
|
5672
|
-
|
|
5673
|
-
|
|
5674
|
-
|
|
5675
|
-
|
|
7384
|
+
const windowK = this.executor.decodeWindowK;
|
|
7385
|
+
if (windowK > 0) {
|
|
7386
|
+
let produced = 0;
|
|
7387
|
+
let stopped = false;
|
|
7388
|
+
while (!stopped && produced < stepsNeeded) {
|
|
7389
|
+
const steps = Math.min(windowK, stepsNeeded - produced);
|
|
7390
|
+
this.executor.submitGreedyDecodeWindow(produced === 0 ? firstToken : null, steps, 0);
|
|
7391
|
+
const tokens = await this.executor.readDecodeWindow(0, steps);
|
|
7392
|
+
produced += steps;
|
|
7393
|
+
for (const tok of tokens) if (consumeToken(tok)) {
|
|
7394
|
+
stopped = true;
|
|
7395
|
+
break;
|
|
7396
|
+
}
|
|
7397
|
+
}
|
|
7398
|
+
} else {
|
|
7399
|
+
const depth = Executor.PIPELINE_DEPTH;
|
|
7400
|
+
let submitted = 0;
|
|
7401
|
+
let consumed = 0;
|
|
7402
|
+
while (consumed < stepsNeeded) {
|
|
7403
|
+
while (submitted < stepsNeeded && submitted < consumed + depth) {
|
|
7404
|
+
this.executor.submitGreedyDecodeStep(submitted === 0 ? firstToken : null, submitted % depth);
|
|
7405
|
+
submitted++;
|
|
7406
|
+
}
|
|
7407
|
+
const tok = await this.executor.readDecodeToken(consumed % depth);
|
|
7408
|
+
consumed++;
|
|
7409
|
+
if (consumeToken(tok)) break;
|
|
5676
7410
|
}
|
|
5677
|
-
const tok = await this.executor.readDecodeToken(consumed % depth);
|
|
5678
|
-
consumed++;
|
|
5679
|
-
if (consumeToken(tok)) break;
|
|
5680
7411
|
}
|
|
5681
7412
|
}
|
|
5682
7413
|
} else {
|
|
@@ -5705,7 +7436,8 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
5705
7436
|
tokensGenerated,
|
|
5706
7437
|
tokensPerSecond,
|
|
5707
7438
|
totalTime,
|
|
5708
|
-
finishReason
|
|
7439
|
+
finishReason,
|
|
7440
|
+
tokenIds: [...generatedIds]
|
|
5709
7441
|
};
|
|
5710
7442
|
}
|
|
5711
7443
|
/**
|
|
@@ -5922,6 +7654,7 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
5922
7654
|
const repo = this._createOptions.repo ?? "";
|
|
5923
7655
|
if (/parler/i.test(repo) || /parler/i.test(this._architecture)) return this.speakParler(text, options);
|
|
5924
7656
|
if (/outetts/i.test(repo) || /oute/i.test(this._architecture)) return this.speakOute(text, options);
|
|
7657
|
+
if (/gepard/i.test(repo) || /gepard/i.test(this._architecture)) return this.speakGepard(text, options);
|
|
5925
7658
|
if (!(this._architecture === "KaniTTS2ForCausalLM" || /kani-tts/i.test(repo) || /\btts\b/i.test(repo))) throw new Error(`speak() requires a Kani-TTS or OuteTTS checkpoint (e.g. "${DEFAULT_MODELS.tts}" or "${DEFAULT_MODELS.ttsOute}"), loaded engine is "${this._architecture}" (repo "${repo}").`);
|
|
5926
7659
|
if (!this._kaniTTS) this._kaniTTS = await KaniTTS.create({
|
|
5927
7660
|
repo: this._createOptions.repo,
|
|
@@ -5966,6 +7699,31 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
5966
7699
|
return ParlerTTS.randomVoiceDescription();
|
|
5967
7700
|
}
|
|
5968
7701
|
/**
|
|
7702
|
+
* Gepard-1.0 speak path: lazily build the Gepard engine (stock Qwen3.5
|
|
7703
|
+
* full-attention codec-LM + 32 host-side FSQ classifier heads + NanoCodec
|
|
7704
|
+
* 21.5 fps decoder) and synthesize `text`. Frame-synchronous: when `onChunk`
|
|
7705
|
+
* is set, audio streams out during generation. Single default voice for now
|
|
7706
|
+
* (voice cloning via the ref-compressor is a documented follow-up).
|
|
7707
|
+
*/
|
|
7708
|
+
async speakGepard(text, options) {
|
|
7709
|
+
if (!this._gepardTTS) this._gepardTTS = await GepardTTS.create({
|
|
7710
|
+
repo: this._createOptions.repo,
|
|
7711
|
+
revision: this._createOptions.revision,
|
|
7712
|
+
hfToken: this._createOptions.hfToken,
|
|
7713
|
+
cacheDir: this._createOptions.cacheDir,
|
|
7714
|
+
maxSeqLen: this.maxSeqLen,
|
|
7715
|
+
dtype: options.dtype
|
|
7716
|
+
});
|
|
7717
|
+
return this._gepardTTS.speak(text, {
|
|
7718
|
+
temperature: options.temperature,
|
|
7719
|
+
topK: options.topK,
|
|
7720
|
+
repetitionPenalty: options.repetitionPenalty,
|
|
7721
|
+
maxFrames: options.maxFrames,
|
|
7722
|
+
stopThreshold: options.stopThreshold,
|
|
7723
|
+
onChunk: options.onChunk
|
|
7724
|
+
});
|
|
7725
|
+
}
|
|
7726
|
+
/**
|
|
5969
7727
|
* OuteTTS speak path: lazily build the OuteTTS engine (loading the chosen preset
|
|
5970
7728
|
* voice's speaker JSON + the folded DAC codec weights from the mirror), rebuilding
|
|
5971
7729
|
* if the requested voice changed. Mirrors the Kani public shape (temperature/topP/topK).
|
|
@@ -6214,7 +7972,8 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
6214
7972
|
tokensGenerated,
|
|
6215
7973
|
tokensPerSecond: tokensGenerated / (totalTime / 1e3),
|
|
6216
7974
|
totalTime,
|
|
6217
|
-
finishReason
|
|
7975
|
+
finishReason,
|
|
7976
|
+
tokenIds: [...generatedIds]
|
|
6218
7977
|
};
|
|
6219
7978
|
}
|
|
6220
7979
|
/** Prepare + prefill + decode for a fully-specified multimodal token sequence. */
|
|
@@ -6299,7 +8058,8 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
6299
8058
|
tokensGenerated,
|
|
6300
8059
|
tokensPerSecond: tokensGenerated / (totalTime / 1e3),
|
|
6301
8060
|
totalTime,
|
|
6302
|
-
finishReason
|
|
8061
|
+
finishReason,
|
|
8062
|
+
tokenIds: [...generatedIds]
|
|
6303
8063
|
};
|
|
6304
8064
|
}
|
|
6305
8065
|
/**
|
|
@@ -6746,6 +8506,7 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
6746
8506
|
this._kaniTTS?.destroy();
|
|
6747
8507
|
this._outeTTS?.destroy();
|
|
6748
8508
|
this._parlerTTS?.destroy();
|
|
8509
|
+
this._gepardTTS?.destroy();
|
|
6749
8510
|
this.visionExecutor?.destroy();
|
|
6750
8511
|
if (this._ctx) {
|
|
6751
8512
|
clearPipelineCache(this._ctx.device);
|
|
@@ -6780,5 +8541,5 @@ var WebGPUEngine = class WebGPUEngine {
|
|
|
6780
8541
|
};
|
|
6781
8542
|
|
|
6782
8543
|
//#endregion
|
|
6783
|
-
export {
|
|
6784
|
-
//# sourceMappingURL=gpu-
|
|
8544
|
+
export { generateGepardBackboneGraph as A, generateGemma4VisionGraph as B, audioTokensToDacCodes as C, parseOuteTtsConfig as D, generateOuteTtsBackboneGraph as E, RequestScheduler as F, resolveGemma4VisionInfo as H, KaniTTS as I, generateQwen3_5VisionGraph as L, gepardCodesToLatent as M, parseGepardConfig as N, GepardTTS as O, Swarm as P, dequantizeGemma4VisionProjection as R, loadOuteSpeaker as S, generateDacSpeechDecoderGraph as T, patchGemma4VisionClips as V, smartResize as _, buildGemma4PosEmbeds as a, OuteTTS as b, buildMRoPECosSin as c, buildPositionIds as d, buildRotaryCosSin as f, preprocessImageGemma4 as g, preprocessImage as h, buildGemma4PoolMatrix as i, gepardChannelLevels as j, buildGepardTextLayout as k, buildMRoPEPositionIds as l, mropeFreqDims as m, GEMMA4_IMAGE_PROCESSOR as n, buildGemma4RotaryCosSin as o, buildVisionPositionTensors as p, QWEN3_5_IMAGE_PROCESSOR as r, buildGemma4VisionPositionTensors as s, WebGPUEngine as t, buildPosEmbeds as u, VisionExecutor as v, dacOutputLength as w, buildOutePromptString as x, ParlerTTS as y, dequantizeMLXProjection as z };
|
|
8545
|
+
//# sourceMappingURL=gpu-CU_Mldk0.mjs.map
|