@atlaskit/editor-plugin-autocomplete 3.0.0 → 3.2.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.
@@ -1,18 +1,27 @@
1
1
  /**
2
2
  * Local Slow Lane Client: On-device inference via @mlc-ai/web-llm.
3
3
  *
4
- * Drop-in replacement for the network-based slow-lane-client. Instead of
5
- * calling a backend API, this client uses MLC WebLLM to run a small language
6
- * model (SmolLM 135M) directly in the browser via WebGPU.
4
+ * Drop-in replacement for the network-based slow-lane-client. Instead of calling
5
+ * a backend API, this client runs two models in the browser via WebGPU, in a
6
+ * single MLCEngine, to reproduce the BE encoder's outputs on-device:
7
+ *
8
+ * - Causal LM (SmolLM2-135M-Instruct): one decode step per word boundary. A
9
+ * registered LogitProcessor captures the raw next-token logits, which
10
+ * `computeBePayload` turns into a whole-word `lm_logits` payload — a faithful
11
+ * port of the BE `CausalLMEncoder._get_top_k_probs` (masked softmax over the
12
+ * vocab's first-tokens, prefix expansion, L2 reservation, log-space pooling).
13
+ * - Semantic embedder (Snowflake Arctic Embed S): produces the real 384-d
14
+ * `semantic_vector`. Inputs are wrapped as passages (see `wrapForArctic`) so
15
+ * the runtime vector lands in the same space as the precomputed word bin.
7
16
  *
8
17
  * ── Why main thread (no Web Worker)? ─────────────────────────────────────
9
- * SmolLM 135M is small enough (~270 MB weights, 350-400 MB VRAM) that
10
- * WebGPU inference on the main thread is production-viable:
18
+ * The models are small enough (~640 MB combined VRAM) that WebGPU inference on
19
+ * the main thread is viable:
11
20
  *
12
21
  * - WebGPU GPU compute is inherently async (doesn't block the main thread)
13
- * - CPU overhead (tokenization + post-processing) is only 5-10 ms
14
- * - Single forward pass latency is 50-150 ms — well within autocomplete
15
- * expectations (~250 ms between word boundaries)
22
+ * - CPU overhead (BE-parity post-processing) is a few ms
23
+ * - Per-inference latency is well within autocomplete expectations
24
+ * (~250 ms between word boundaries)
16
25
  *
17
26
  * This avoids all the complexity of Web Workers:
18
27
  * - No CSP workarounds (blob URLs, inline scripts)
@@ -26,7 +35,7 @@
26
35
  * asynchronously after each updateContext() call.
27
36
  */
28
37
 
29
- import type { MLCEngine, InitProgressReport, AppConfig } from '@mlc-ai/web-llm';
38
+ import type { MLCEngine, InitProgressReport, AppConfig, LogitProcessor } from '@mlc-ai/web-llm';
30
39
 
31
40
  import { isAutocompleteDebugEnabled } from './debug-mode';
32
41
  import { isWordBoundary } from './slow-lane-client';
@@ -85,20 +94,437 @@ export interface LocalSlowLaneClient {
85
94
 
86
95
  const DEFAULT_DEBOUNCE_MS = 300;
87
96
 
88
- export const LOCAL_MLC_MODEL_ID = 'SmolLM2-135M-Instruct-q0f16-MLC';
97
+ export const LOCAL_MLC_CAUSAL_MODEL_ID = 'SmolLM2-135M-Instruct-q0f16-MLC';
98
+
99
+ /**
100
+ * MLC ID for the semantic embedder (Snowflake Arctic Embed S, batch=4 variant).
101
+ *
102
+ * The `-b4` suffix selects the prebuilt variant compiled for a max batch size of
103
+ * 4 (≈239 MB VRAM) rather than `-b32` (≈1023 MB VRAM). Autocomplete embeds one
104
+ * context at a time, so `-b4` is the right fit. This model IS in
105
+ * `prebuiltAppConfig.model_list` of web-llm 0.2.82 — no `customModelConfig` needed.
106
+ */
107
+ export const LOCAL_MLC_EMBEDDING_MODEL_ID = 'snowflake-arctic-embed-s-q0f32-MLC-b4';
108
+
109
+ /**
110
+ * Wrap raw context text with BERT special tokens before embedding.
111
+ *
112
+ * web-llm's `EmbeddingPipeline` does NOT auto-prepend `[CLS]` / append `[SEP]`
113
+ * (the official MLC embeddings example wraps manually). The Python
114
+ * `sentence_transformers` side that generated the word-vector bin adds these
115
+ * inside `model.encode()`, so we must mirror it here for the runtime context
116
+ * vector to land in the same region of Arctic's embedding space as the bin.
117
+ *
118
+ * No query prefix is applied: the semantic step is sentence-to-sentence (`s2s`)
119
+ * similarity ("which words are conceptually similar to this context?"), not
120
+ * sentence-to-passage (`s2p`) retrieval. Arctic's query prefix would misframe
121
+ * the relationship. Encode both sides as passages. See implementation.md §4.3.
122
+ */
123
+ export const wrapForArctic = (text: string): string => `[CLS] ${text} [SEP]`;
124
+
125
+ /**
126
+ * BE-parity constants — must match `CausalLMEncoder` defaults in the Python
127
+ * sidecar (`cc-smarts/python-sidecar/src/causal_lm_encoder.py`) and
128
+ * `SlowLaneEngine` (`typeahead_context_encoding.py`) so local payloads behave
129
+ * identically to the server-client setup.
130
+ */
131
+ export const BE_PARITY = {
132
+ /** Final payload size cap (BE: `top_k_words`). */
133
+ TOP_K_WORDS: 2000,
134
+ /** L2 (domain) words admitted unconditionally before pooling (BE: `reserved_l2_slots`). */
135
+ RESERVED_L2_SLOTS: 500,
136
+ /** Log-space additive bias favouring L2 over L3 in the pool (BE: `l2_bias`). */
137
+ L2_BIAS: 1.0,
138
+ /** Drop words below this probability from the final payload (BE: `> 0.00001`). */
139
+ MIN_PROB: 0.00001,
140
+ /**
141
+ * Word-level approximation of the BE causal LM token limit.
142
+ *
143
+ * BE: `CausalLMEncoder.max_context_tokens = 100` (BPE tokens, left-truncated).
144
+ * FE: no tokenizer available, so we approximate with word count. English text
145
+ * averages ~1.3–1.5 BPE tokens/word, meaning 100 words ≈ 130–150 tokens.
146
+ * Using 100 words keeps the approximation simple and errs on the side of
147
+ * sending slightly more context than the BE sees — acceptable for a PoC.
148
+ */
149
+ MAX_CONTEXT_TOKENS: 100,
150
+ /**
151
+ * Word-level rolling window for the semantic embedder.
152
+ *
153
+ * BE: `SlowLaneEngine.max_context_words = 100` (applied in
154
+ * `typeahead_context_encoding.py` before calling `SemanticEncoder.encode`).
155
+ * Truncated identically here so the runtime Arctic vector lands in the same
156
+ * region of the embedding space as the precomputed word-vector bin.
157
+ */
158
+ MAX_CONTEXT_WORDS: 100,
159
+ } as const;
160
+
161
+ /**
162
+ * Return the last `n` whitespace-separated words of `text`, joined by spaces.
163
+ * Mirrors the BE rolling-window truncation applied before both encoders.
164
+ */
165
+ const truncateToLastNWords = (text: string, n: number): string => {
166
+ const words = text.trim().split(/\s+/u);
167
+ return words.length <= n ? text : words.slice(-n).join(' ');
168
+ };
169
+
170
+ // ─── Logit capture ─────────────────────────────────────────────────────────
89
171
 
90
- /** HF root for the default weights (includes `tensor-cache.json` for WebLLM 0.2+). */
91
- export const LOCAL_MLC_HF_MODEL_REPO =
92
- 'https://huggingface.co/mlc-ai/SmolLM2-135M-Instruct-q0f16-MLC';
172
+ /**
173
+ * A LogitProcessor that captures the raw next-token logits and passes them
174
+ * through unmodified.
175
+ *
176
+ * web-llm invokes `processLogits` on the CPU after the model's forward pass and
177
+ * before sampling, handing us the full `Float32Array(vocab_size)` at the current
178
+ * decode position. We copy it off web-llm's shared buffer (which it may reuse
179
+ * across calls) and return the original untouched so sampling is unaffected.
180
+ *
181
+ * This is the raw-logit access the BE-parity algorithm needs (masked softmax +
182
+ * prefix expansion, consumed in a later step). Registered for the causal LM
183
+ * only — the embedder never decodes tokens, so it produces no logits.
184
+ */
185
+ class CapturingLogitProcessor implements LogitProcessor {
186
+ captured: Float32Array | null = null;
93
187
 
94
- export const LOCAL_MLC_MODEL_LIB_WASM_NAME = 'SmolLM2-135M-Instruct-q0f16-ctx4k_cs1k-webgpu.wasm';
188
+ processLogits = (logits: Float32Array): Float32Array => {
189
+ // Copy off web-llm's shared buffer — it may reuse `logits` across calls.
190
+ this.captured = new Float32Array(logits);
191
+ return logits;
192
+ };
193
+
194
+ processSampledToken = (): void => {
195
+ // No-op — we don't track sampled tokens.
196
+ };
197
+
198
+ resetState = (): void => {
199
+ this.captured = null;
200
+ };
201
+ }
202
+
203
+ // ─── BE-parity data + algorithm ──────────────────────────────────────────────
95
204
 
96
205
  /**
97
- * Original target repo (add-basics fine-tune). **Not compatible with WebLLM 0.2.x** (no `tensor-cache.json`).
98
- * @see module doc above
206
+ * Prefix-expansion map: first-token id → words whose space-prefixed SmolLM2
207
+ * encoding starts with that token. Generated offline by
208
+ * `scripts/gen_first_token_to_words.py`, which mirrors the BE's in-memory map
209
+ * (`CausalLMEncoder._ensure_loaded`).
210
+ *
211
+ * Populated lazily by `loadBePayloadData()` from a dynamically-imported JSON so
212
+ * the (large) payload is only fetched when the local client is actually
213
+ * initialised — keeping it out of the editor's main chunk for the vast majority
214
+ * of users (who run with `useLocalModel` off).
99
215
  */
100
- export const HUGGINGFACE_TB_SMOLLM_ADD_BASICS_REPO =
101
- 'https://huggingface.co/HuggingFaceTB/smollm-135M-instruct-add-basics-q0f16-MLC';
216
+ let firstTokenToWords: Map<number, string[]> = new Map();
217
+
218
+ /**
219
+ * L2 (Atlassian-domain) word set, derived from the keys of `vocabulary_10k.json`.
220
+ * Used by `computeBePayload` for tier-aware ranking: any word in the prefix map
221
+ * that is not in this set is treated as L3 (general English), matching the BE.
222
+ * Populated lazily alongside `firstTokenToWords` — see `loadBePayloadData()`.
223
+ */
224
+ let l2Words: Set<string> = new Set();
225
+
226
+ /**
227
+ * Array of token IDs that appear as a first token for at least one vocabulary
228
+ * word. Derived from `firstTokenToWords` when the data loads so `computeBePayload`
229
+ * does not re-allocate this array on every word-boundary call.
230
+ */
231
+ let prefixMapTokenIds: number[] = [];
232
+
233
+ /** De-dupes concurrent loads and lets repeated calls await the same payload. */
234
+ let bePayloadDataPromise: Promise<void> | undefined;
235
+
236
+ /**
237
+ * Unwrap a dynamically imported JSON module to the parsed JSON value, working
238
+ * across the two interop modes AFM's bundler chain emits:
239
+ *
240
+ * 1. **`.default`-wrapped namespace** — classic webpack (and Jest) hang the
241
+ * JSON value under the `default` export.
242
+ * 2. **Named-exports namespace** — webpack 5 / atlaspack with JSON
243
+ * named-exports (or native ESM JSON modules) expose each top-level key as
244
+ * a named export and shadow `default`, so `mod.default` can be `undefined`
245
+ * (or some unrelated value) even though `mod` itself holds the data.
246
+ *
247
+ * The caller MUST declare the underlying JSON shape via `shape` because, in
248
+ * named-exports mode, a dense array `["a","b"]` and a sparse numeric-keyed
249
+ * object `{"5":"a","12":"b"}` are emitted identically (`{"0":..}` / `{"5":..}`);
250
+ * no runtime heuristic can tell them apart, so only the caller knows which:
251
+ *
252
+ * - `'object'` — the JSON is a `{...}` (including sparse maps keyed by integer
253
+ * IDs). The named exports are rebuilt into a plain object so `Object.entries`
254
+ * yields the real keys, not synthetic array indices.
255
+ * - `'array'` — the JSON is a `[...]`, reconstructed from the `0..n-1` indices.
256
+ *
257
+ * :param mod: The raw module object returned by `await import('./*.json')`.
258
+ * :param shape: `'object'` if the source JSON is `{...}`, `'array'` if `[...]`.
259
+ * :returns: The parsed JSON value, or `null` if neither interop mode applies.
260
+ */
261
+ const unwrapJsonModule = <T,>(mod: unknown, shape: 'object' | 'array'): T | null => {
262
+ if (mod == null || typeof mod !== 'object') {
263
+ return null;
264
+ }
265
+ const namespace = mod as Record<string, unknown> & { default?: unknown };
266
+
267
+ // Compute the named-export own-keys (strip synthetic markers).
268
+ const ownKeys = Object.keys(namespace).filter(
269
+ (k) => k !== 'default' && k !== '__esModule',
270
+ );
271
+
272
+ // PREFER named exports when present — they always reflect the JSON's real
273
+ // top-level keys / indices, regardless of what `default` happens to be.
274
+ // Under JSON named-exports mode `default` is not necessarily the parsed
275
+ // value (e.g. for `{"service": 0, ...}` it can be the number `0`, with the
276
+ // real data in the named exports), so taking `default` first would corrupt it.
277
+ if (ownKeys.length > 0) {
278
+ if (shape === 'array') {
279
+ // JSON arrays are dense; reconstruct from `0..length-1` indices.
280
+ const len = ownKeys.length;
281
+ const arr = new Array(len);
282
+ for (let i = 0; i < len; i++) {
283
+ arr[i] = namespace[String(i)];
284
+ }
285
+ return arr as T;
286
+ }
287
+ // shape === 'object'. Rebuild a plain object from the (stripped) own
288
+ // keys so callers can `Object.entries()` it without iterating over
289
+ // `default` / `__esModule`, and to detach from the module-namespace
290
+ // object (which is sealed/non-extensible on some bundler outputs).
291
+ const obj: Record<string, unknown> = {};
292
+ for (const k of ownKeys) {
293
+ obj[k] = namespace[k];
294
+ }
295
+ return obj as T;
296
+ }
297
+
298
+ // Fallback: no named exports — classic webpack JSON-module interop where
299
+ // the whole parsed JSON value is hung under `default`. Trust it.
300
+ if ('default' in namespace && namespace.default != null) {
301
+ return namespace.default as T;
302
+ }
303
+
304
+ return null;
305
+ };
306
+
307
+ /**
308
+ * Lazily load and build the BE-parity lookup tables from their JSON payloads.
309
+ * The dynamic imports are split into their own async chunks so neither file is
310
+ * bundled into the editor's main chunk unless local inference is initialised.
311
+ *
312
+ * :returns:
313
+ * A promise that resolves once `firstTokenToWords`, `l2Words` and
314
+ * `prefixMapTokenIds` are populated.
315
+ */
316
+ const loadBePayloadData = (): Promise<void> => {
317
+ if (!bePayloadDataPromise) {
318
+ bePayloadDataPromise = (async () => {
319
+ const [firstTokenToWordsModule, vocabularyModule] = await Promise.all([
320
+ import(
321
+ /* webpackChunkName: "@atlaskit-internal_editor-plugin-autocomplete-first-token-to-words" */ './data/first_token_to_words.json'
322
+ ),
323
+ import(
324
+ /* webpackChunkName: "@atlaskit-internal_editor-plugin-autocomplete-vocabulary-10k" */ './data/vocabulary_10k.json'
325
+ ),
326
+ ]);
327
+
328
+ const firstTokenToWordsData = unwrapJsonModule<Record<string, string[]>>(
329
+ firstTokenToWordsModule,
330
+ 'object',
331
+ );
332
+ const vocabularyData = unwrapJsonModule<{ words: Record<string, unknown> }>(
333
+ vocabularyModule,
334
+ 'object',
335
+ );
336
+
337
+ if (firstTokenToWordsData == null || vocabularyData?.words == null) {
338
+ // Hard-fail with a precise message so the catch() in initEngine logs
339
+ // exactly which import couldn't be unwrapped, rather than the generic
340
+ // V8 "Cannot convert undefined or null to object" we hit before the
341
+ // helper was added.
342
+ throw new Error(
343
+ `[LocalSlowLane] JSON module could not be unwrapped — ` +
344
+ `firstTokenToWordsData=${firstTokenToWordsData == null ? 'null/undefined' : 'defined'}, ` +
345
+ `vocabularyData=${vocabularyData == null ? 'null/undefined' : vocabularyData.words == null ? 'defined but missing .words' : 'defined'}`,
346
+ );
347
+ }
348
+
349
+ firstTokenToWords = new Map(
350
+ Object.entries(firstTokenToWordsData).map(([tokenId, words]) => [
351
+ Number(tokenId),
352
+ words,
353
+ ]),
354
+ );
355
+ l2Words = new Set(Object.keys(vocabularyData.words));
356
+ prefixMapTokenIds = Array.from(firstTokenToWords.keys());
357
+
358
+ if (isAutocompleteDebugEnabled()) {
359
+ // eslint-disable-next-line no-console
360
+ console.log(
361
+ '%c[LocalSlowLane] %c✅ BE-parity payload data loaded:',
362
+ 'color: #9c27b0; font-weight: bold;',
363
+ 'color: #4caf50; font-weight: bold;',
364
+ {
365
+ firstTokenToWordsEntries: firstTokenToWords.size,
366
+ l2WordsCount: l2Words.size,
367
+ prefixMapTokenIdsLength: prefixMapTokenIds.length,
368
+ },
369
+ );
370
+ }
371
+ })().catch((e) => {
372
+ // Don't cache a rejected promise — a transient import failure would
373
+ // otherwise prevent the local model from ever initialising again this
374
+ // session. Reset so the next init attempt retries.
375
+ bePayloadDataPromise = undefined;
376
+ throw e;
377
+ });
378
+ }
379
+ return bePayloadDataPromise;
380
+ };
381
+
382
+ /**
383
+ * Convert a raw next-token logit vector into a whole-word probability payload,
384
+ * faithfully porting the BE `CausalLMEncoder._get_top_k_probs`
385
+ * (`cc-smarts/python-sidecar/src/causal_lm_encoder.py`).
386
+ *
387
+ * Steps: (1) numerically-stable masked softmax over only the token ids present
388
+ * in the prefix-expansion map; (2) spread each token's probability to every
389
+ * whole word sharing that first token, taking the max; (3) reserve the top L2
390
+ * words unconditionally; (4) rank the remainder in a log-space pool with an
391
+ * additive L2 bias; (5) emit raw probabilities for the survivors, lowercased
392
+ * and trimmed at `MIN_PROB`.
393
+ *
394
+ * :params:
395
+ * rawLogits: Full-vocabulary logits from the LM's single decode step
396
+ * prefixMap: Map of first-token id to the words starting with that token
397
+ * domainWords: Set of L2 (domain) words, for tier-aware ranking
398
+ * :returns:
399
+ * A record of lowercase word to probability — the BE `lm_logits` payload
400
+ */
401
+ export const computeBePayload = (
402
+ rawLogits: Float32Array,
403
+ prefixMap: Map<number, string[]>,
404
+ domainWords: Set<string>,
405
+ /**
406
+ * Pre-derived token-ID array for the softmax mask. Defaults to the
407
+ * module-level `prefixMapTokenIds` (zero allocation in production). Pass
408
+ * `Array.from(prefixMap.keys())` in tests that supply a custom prefixMap so
409
+ * the softmax mask stays consistent with the iteration in Step 2.
410
+ */
411
+ validTokenIds: number[] = prefixMapTokenIds,
412
+ ): Record<string, number> => {
413
+
414
+ // 1. Numerically-stable masked softmax over validTokenIds only.
415
+ let maxLogit = -Infinity;
416
+ for (const id of validTokenIds) {
417
+ const v = rawLogits[id];
418
+ if (v > maxLogit) {
419
+ maxLogit = v;
420
+ }
421
+ }
422
+ let sumExp = 0;
423
+ const expByToken = new Map<number, number>();
424
+ for (const id of validTokenIds) {
425
+ const e = Math.exp(rawLogits[id] - maxLogit);
426
+ expByToken.set(id, e);
427
+ sumExp += e;
428
+ }
429
+
430
+ // 2. Prefix expansion with max aggregation (probabilities sum to 1 over the
431
+ // masked subset, so divide each token's exp by sumExp on the fly).
432
+ const wordProbs = new Map<string, number>();
433
+ for (const [id, words] of prefixMap) {
434
+ const p = sumExp > 0 ? (expByToken.get(id) ?? 0) / sumExp : 0;
435
+ for (const w of words) {
436
+ const prev = wordProbs.get(w) ?? 0;
437
+ if (p > prev) {
438
+ wordProbs.set(w, p);
439
+ }
440
+ }
441
+ }
442
+
443
+ // 3. Split into L2 / L3 and reserve the top L2 slots unconditionally.
444
+ const l2Matches: Array<[string, number]> = [];
445
+ const l3Matches: Array<[string, number]> = [];
446
+ for (const [w, p] of wordProbs) {
447
+ if (domainWords.has(w)) {
448
+ l2Matches.push([w, p]);
449
+ } else {
450
+ l3Matches.push([w, p]);
451
+ }
452
+ }
453
+ l2Matches.sort((a, b) => b[1] - a[1]);
454
+ const reserved = l2Matches.slice(0, BE_PARITY.RESERVED_L2_SLOTS);
455
+
456
+ // 4. Pool the leftovers in log space; the L2 bias only affects ranking here.
457
+ // Words in l2Matches are unique and the array is sorted descending, so the
458
+ // non-reserved entries are exactly the tail after the reserved prefix — slice
459
+ // it directly rather than allocating a Set and scanning every entry on this
460
+ // hot path (runs ~every word boundary while typing).
461
+ const pool: Array<[string, number]> = [];
462
+ for (const [w, p] of l2Matches.slice(BE_PARITY.RESERVED_L2_SLOTS)) {
463
+ pool.push([w, Math.log(Math.max(p, 1e-10)) + BE_PARITY.L2_BIAS]);
464
+ }
465
+ for (const [w, p] of l3Matches) {
466
+ pool.push([w, Math.log(Math.max(p, 1e-10))]);
467
+ }
468
+ pool.sort((a, b) => b[1] - a[1]);
469
+ const remainingSlots = Math.max(0, BE_PARITY.TOP_K_WORDS - reserved.length);
470
+ const poolWinners = pool.slice(0, remainingSlots);
471
+
472
+ // 5. Assemble payload: store RAW probabilities (the bias was ranking-only),
473
+ // lowercase keys, trimmed at MIN_PROB. Reserved first, then pool winners.
474
+ // Reserved entries are written first; pool-winner writes must NOT clobber a
475
+ // reserved entry whose normalised key collides (two source words can
476
+ // `.trim().toLowerCase()` to the same key — e.g. "Function" vs "function ").
477
+ // Without the existence guard, a low-probability pool winner would silently
478
+ // overwrite the (higher-probability) reserved entry, degrading top-K
479
+ // quality in a way that's invisible from the debug summary.
480
+ const result: Record<string, number> = {};
481
+ const addEntry = (word: string, prob: number, allowOverwrite: boolean): void => {
482
+ if (prob <= BE_PARITY.MIN_PROB) {
483
+ return;
484
+ }
485
+ const key = word.trim().toLowerCase();
486
+ if (!allowOverwrite && key in result) {
487
+ return;
488
+ }
489
+ result[key] = prob;
490
+ };
491
+ for (const [w, p] of reserved) {
492
+ addEntry(w, p, true);
493
+ }
494
+ for (const [w] of poolWinners) {
495
+ addEntry(w, wordProbs.get(w) ?? 0, false);
496
+ }
497
+
498
+ if (isAutocompleteDebugEnabled()) {
499
+ const topReserved = reserved
500
+ .slice(0, 5)
501
+ .map(([w, p]) => `${w}:${(p * 100).toFixed(2)}%`)
502
+ .join(', ');
503
+ const topPool = poolWinners
504
+ .slice(0, 5)
505
+ .map(([w]) => `${w}:${((wordProbs.get(w) ?? 0) * 100).toFixed(2)}%`)
506
+ .join(', ');
507
+ // eslint-disable-next-line no-console
508
+ console.log(
509
+ '%c[computeBePayload] %c%d valid tokens → %d words expanded | L2: %d / L3: %d | reserved: %d | pool winners: %d | final: %d words\n maxLogit(masked): %s | sumExp: %s\n top reserved L2: %s\n top pool: %s',
510
+ 'color: #9c27b0; font-weight: bold;',
511
+ 'color: inherit;',
512
+ validTokenIds.length,
513
+ wordProbs.size,
514
+ l2Matches.length,
515
+ l3Matches.length,
516
+ reserved.length,
517
+ poolWinners.length,
518
+ Object.keys(result).length,
519
+ maxLogit.toFixed(3),
520
+ sumExp.toFixed(1),
521
+ topReserved || '(none)',
522
+ topPool || '(none)',
523
+ );
524
+ }
525
+
526
+ return result;
527
+ };
102
528
 
103
529
  // ─── Factory ─────────────────────────────────────────────────────────────────
104
530
 
@@ -127,7 +553,7 @@ export const createLocalSlowLaneClient = (
127
553
  debounceMs = DEFAULT_DEBOUNCE_MS,
128
554
  onUpdate,
129
555
  onStatus,
130
- modelId = LOCAL_MLC_MODEL_ID,
556
+ modelId = LOCAL_MLC_CAUSAL_MODEL_ID,
131
557
  customModelConfig,
132
558
  } = config;
133
559
 
@@ -143,6 +569,9 @@ export const createLocalSlowLaneClient = (
143
569
  let initFailed = false;
144
570
  let engine: MLCEngine | null = null;
145
571
  let engineInitPromise: Promise<void> | null = null;
572
+ // Captures raw next-token logits from the LM's single decode step. Registered
573
+ // with the engine below; `lmLogitsCapture.captured` is consumed in a later step.
574
+ const lmLogitsCapture = new CapturingLogitProcessor();
146
575
 
147
576
  const unloadEngine = (engineToUnload: MLCEngine): void => {
148
577
  engineToUnload.unload().catch((error: unknown) => {
@@ -178,20 +607,25 @@ export const createLocalSlowLaneClient = (
178
607
  if (isAutocompleteDebugEnabled()) {
179
608
  // eslint-disable-next-line no-console
180
609
  console.log(
181
- `%c[LocalSlowLane] %c🚀 Initialising MLC engine with model: ${modelId}`,
610
+ `%c[LocalSlowLane] %c🚀 Initialising MLC engine with models: ${modelId} (LM) + ${LOCAL_MLC_EMBEDDING_MODEL_ID} (embedder)`,
182
611
  'color: #9c27b0; font-weight: bold;',
183
612
  'color: inherit;',
184
613
  );
185
614
  }
186
- onStatus?.(`Initialising model: ${modelId}…`);
615
+ onStatus?.(`Initialising models: ${modelId} + ${LOCAL_MLC_EMBEDDING_MODEL_ID}…`);
187
616
 
188
617
  if (!('gpu' in navigator)) {
189
618
  throw new Error('WebGPU not supported');
190
619
  }
191
620
 
192
- const { CreateMLCEngine, prebuiltAppConfig } = await import(
193
- /* webpackChunkName: "@atlaskit-internal_editor-plugin-autocomplete-mlc-web-llm" */ '@mlc-ai/web-llm'
194
- );
621
+ // Fetch the web-llm runtime and the BE-parity lookup tables in parallel;
622
+ // both are dynamically imported so they stay out of the main editor chunk.
623
+ const [{ MLCEngine: MLCEngineCtor, prebuiltAppConfig }] = await Promise.all([
624
+ import(
625
+ /* webpackChunkName: "@atlaskit-internal_editor-plugin-autocomplete-mlc-web-llm" */ '@mlc-ai/web-llm'
626
+ ),
627
+ loadBePayloadData(),
628
+ ]);
195
629
 
196
630
  const customModelRecord: WebLlmModelRecord | undefined = customModelConfig
197
631
  ? {
@@ -222,27 +656,49 @@ export const createLocalSlowLaneClient = (
222
656
  ],
223
657
  };
224
658
 
225
- engine = await CreateMLCEngine(modelId, {
659
+ // Construct the engine with the logit-capture processor registered for
660
+ // the causal LM only (the embedder never decodes tokens), then load
661
+ // both the LM and the embedder into the same engine (multi-model).
662
+ const newEngine = new MLCEngineCtor({
226
663
  appConfig,
227
664
  initProgressCallback,
665
+ logitProcessorRegistry: new Map([[modelId, lmLogitsCapture]]),
228
666
  });
229
667
 
668
+ await newEngine.reload([modelId, LOCAL_MLC_EMBEDDING_MODEL_ID]);
669
+
230
670
  if (destroyed) {
231
671
  // destroy() was called while we were loading — clean up
232
- unloadEngine(engine);
233
- engine = null;
672
+ unloadEngine(newEngine);
234
673
  return;
235
674
  }
236
675
 
676
+ engine = newEngine;
237
677
  ready = true;
238
678
 
239
679
  if (isAutocompleteDebugEnabled()) {
240
680
  // eslint-disable-next-line no-console
241
681
  console.log(
242
- '%c[LocalSlowLane] %c✅ MLC engine loaded and ready',
682
+ '%c[LocalSlowLane] %c✅ Both models loaded and ready',
243
683
  'color: #9c27b0; font-weight: bold;',
244
684
  'color: #4caf50;',
245
685
  );
686
+ // One-time identity summary so you can confirm which models are active
687
+ // without digging through the init-progress scroll.
688
+ // eslint-disable-next-line no-console
689
+ console.log(
690
+ '%c[LocalSlowLane] %c🧠 Causal LM →',
691
+ 'color: #9c27b0; font-weight: bold;',
692
+ 'color: #2196f3; font-weight: bold;',
693
+ modelId,
694
+ );
695
+ // eslint-disable-next-line no-console
696
+ console.log(
697
+ '%c[LocalSlowLane] %c🔢 Embedder →',
698
+ 'color: #9c27b0; font-weight: bold;',
699
+ 'color: #009688; font-weight: bold;',
700
+ LOCAL_MLC_EMBEDDING_MODEL_ID,
701
+ );
246
702
  }
247
703
  onStatus?.('Model loaded and ready.');
248
704
  } catch (err) {
@@ -271,83 +727,112 @@ export const createLocalSlowLaneClient = (
271
727
  // ── Inference ──────────────────────────────────────────────────────────
272
728
 
273
729
  /**
274
- * Run a single forward pass to extract next-token logit probabilities.
730
+ * Run a single forward pass to produce the BE-parity slow-lane outputs.
275
731
  *
276
- * We use the chat completions API with `max_tokens: 1` and `logprobs: true`
277
- * to get the model's next-token distribution without generating text.
278
- * This is the cheapest possible inference call — a single forward pass.
732
+ * Two calls run in parallel on the shared engine:
733
+ * - `completions.create({ max_tokens: 1 })` runs the causal LM for exactly
734
+ * one decode step. We ignore the generated text; the LogitProcessor
735
+ * captures the raw next-token logits during that step, which we turn into
736
+ * a whole-word payload via `computeBePayload`.
737
+ * - `embeddings.create(...)` runs the Arctic embedder to produce the real
738
+ * 384-d semantic vector (passage-encoded; see `wrapForArctic`).
279
739
  */
280
740
  const runInference = async (text: string, requestId: number): Promise<void> => {
281
741
  if (!engine || destroyed) {
282
742
  return;
283
743
  }
284
744
 
285
- try {
286
- // Use chat completion with logprobs to get next-token distribution
287
- const response = await engine.chat.completions.create({
288
- messages: [
289
- {
290
- role: 'user',
291
- content: text,
292
- },
293
- ],
294
- max_tokens: 1,
295
- logprobs: true,
296
- top_logprobs: 5,
297
- temperature: 0,
298
- });
745
+ // Clear the capture buffer so we read only this pass's logits. The engine
746
+ // serialises per-model requests and updateContext is debounced, so the
747
+ // latest request's decode step is the last to populate `captured` before
748
+ // we read it below; stale requests bail on the latestRequestId guard.
749
+ lmLogitsCapture.resetState();
299
750
 
300
- // Discard stale results
301
- if (requestId < latestRequestId || destroyed) {
302
- return;
303
- }
751
+ // Apply BE-parity rolling-window truncation before both encoders.
752
+ // BE semantic: last max_context_words words (typeahead_context_encoding.py:36)
753
+ // BE causal LM: last max_context_tokens BPE tokens (causal_lm_encoder.py:194–198),
754
+ // approximated here with word count (no tokenizer available on FE).
755
+ const lmText = truncateToLastNWords(text, BE_PARITY.MAX_CONTEXT_TOKENS);
756
+ const semanticText = truncateToLastNWords(text, BE_PARITY.MAX_CONTEXT_WORDS);
757
+ const arcticInput = wrapForArctic(semanticText);
304
758
 
305
- // ── Extract LM logits ───────────────────────────────────────
306
- const lmLogits: Record<string, number> = {};
759
+ if (isAutocompleteDebugEnabled()) {
760
+ // eslint-disable-next-line no-console
761
+ console.log(
762
+ `%c[LocalSlowLane] %c🔢 Arctic input (${arcticInput.length} chars, ${semanticText.split(/\s+/u).length} words): "${arcticInput.length > 100 ? `${arcticInput.slice(0, 100)}…` : arcticInput}"`,
763
+ 'color: #9c27b0; font-weight: bold;',
764
+ 'color: #009688;',
765
+ );
766
+ // eslint-disable-next-line no-console
767
+ console.log(
768
+ `%c[LocalSlowLane] %c🧠 LM input (${lmText.length} chars, ${lmText.split(/\s+/u).length} words): "${lmText.length > 100 ? `${lmText.slice(0, 100)}…` : lmText}"`,
769
+ 'color: #9c27b0; font-weight: bold;',
770
+ 'color: #2196f3;',
771
+ );
772
+ }
307
773
 
308
- const logprobsContent = response.choices?.[0]?.logprobs?.content;
309
- if (logprobsContent && logprobsContent.length > 0) {
310
- const tokenLogprobs = logprobsContent[0];
774
+ try {
775
+ const tStart = performance.now();
776
+ let tLmDone = 0;
777
+ let tEmbDone = 0;
778
+
779
+ const [, embeddingResponse] = await Promise.all([
780
+ engine.completions
781
+ .create({
782
+ model: modelId,
783
+ prompt: lmText,
784
+ max_tokens: 1,
785
+ temperature: 0,
786
+ logprobs: false,
787
+ })
788
+ .then((r) => {
789
+ tLmDone = performance.now();
790
+ return r;
791
+ }),
792
+ engine.embeddings
793
+ .create({
794
+ model: LOCAL_MLC_EMBEDDING_MODEL_ID,
795
+ input: arcticInput,
796
+ })
797
+ .then((r) => {
798
+ tEmbDone = performance.now();
799
+ return r;
800
+ }),
801
+ ]);
311
802
 
312
- // Add the top token
313
- if (tokenLogprobs.token) {
314
- const token = tokenLogprobs.token.trim().toLowerCase();
315
- // @ts-ignore TS1501: Unicode regex flag requires a newer TS target than the declaration build uses.
316
- if (token.length > 0 && /^[a-z]/iu.test(token)) {
317
- lmLogits[token] = Math.exp(tokenLogprobs.logprob);
318
- }
319
- }
803
+ if (isAutocompleteDebugEnabled()) {
804
+ // eslint-disable-next-line no-console
805
+ console.log(
806
+ `%c[LocalSlowLane] %c⏱ LM: ${(tLmDone - tStart).toFixed(0)}ms | Embedder: ${(tEmbDone - tStart).toFixed(0)}ms | Total: ${(Math.max(tLmDone, tEmbDone) - tStart).toFixed(0)}ms`,
807
+ 'color: #9c27b0; font-weight: bold;',
808
+ 'color: #ff9800;',
809
+ );
810
+ }
320
811
 
321
- // Add alternative tokens from top_logprobs
322
- if (tokenLogprobs.top_logprobs) {
323
- for (const alt of tokenLogprobs.top_logprobs) {
324
- const token = alt.token.trim().toLowerCase();
325
- // @ts-ignore TS1501: Unicode regex flag requires a newer TS target than the declaration build uses.
326
- if (token.length > 0 && /^[a-z]/iu.test(token)) {
327
- lmLogits[token] = Math.exp(alt.logprob);
328
- }
329
- }
330
- }
812
+ // Discard stale results
813
+ if (requestId < latestRequestId || destroyed) {
814
+ return;
331
815
  }
332
816
 
333
- storedLmLogits = Object.keys(lmLogits).length > 0 ? lmLogits : null;
334
-
335
- // ── Semantic vector ─────────────────────────────────────────
336
- // SmolLM is a generative model, not an embedding model, so we
337
- // don't get a true semantic vector. We generate a lightweight
338
- // pseudo-embedding from the logit distribution for compatibility
339
- // with the existing scoring pipeline.
340
- //
341
- // For a production implementation, you would use a dedicated
342
- // embedding model (e.g. via web-llm's embeddings API with an
343
- // embedding-specific model).
344
- if (storedLmLogits) {
345
- const logitValues = Object.values(storedLmLogits);
346
- storedContextVector = new Float32Array(logitValues);
817
+ // ── LM logits: whole-word BE-parity payload ──────────────────
818
+ const rawLogits = lmLogitsCapture.captured;
819
+ if (rawLogits) {
820
+ const payload = computeBePayload(rawLogits, firstTokenToWords, l2Words);
821
+ storedLmLogits = Object.keys(payload).length > 0 ? payload : null;
347
822
  } else {
348
- storedContextVector = null;
823
+ storedLmLogits = null;
349
824
  }
350
825
 
826
+ // ── Semantic vector: real 384-d Arctic embedding ─────────────
827
+ // Guard against base64-encoded responses (encoding_format: 'base64' would
828
+ // yield a string, and new Float32Array(string) silently produces an empty
829
+ // array, corrupting downstream cosine-similarity scoring).
830
+ const embedding = embeddingResponse.data?.[0]?.embedding;
831
+ storedContextVector =
832
+ Array.isArray(embedding) && embedding.length > 0
833
+ ? new Float32Array(embedding as number[])
834
+ : null;
835
+
351
836
  if (isAutocompleteDebugEnabled()) {
352
837
  // eslint-disable-next-line no-console
353
838
  console.groupCollapsed(
@@ -355,16 +840,23 @@ export const createLocalSlowLaneClient = (
355
840
  'color: #9c27b0; font-weight: bold;',
356
841
  'color: inherit;',
357
842
  );
358
- // eslint-disable-next-line no-console
359
- console.log(
360
- storedContextVector
361
- ? `✅ pseudo-vector: ${storedContextVector.length} dims`
362
- : '❌ No vector',
363
- );
843
+ if (storedContextVector) {
844
+ let sumSq = 0;
845
+ for (let i = 0; i < storedContextVector.length; i++) {
846
+ sumSq += storedContextVector[i] * storedContextVector[i];
847
+ }
848
+ // eslint-disable-next-line no-console
849
+ console.log(
850
+ `✅ semantic vector: ${storedContextVector.length} dims (L2 norm ${Math.sqrt(sumSq).toFixed(3)})`,
851
+ );
852
+ } else {
853
+ // eslint-disable-next-line no-console
854
+ console.log('❌ No vector');
855
+ }
364
856
  // eslint-disable-next-line no-console
365
857
  console.log(
366
858
  storedLmLogits
367
- ? `✅ lm_logits: ${Object.keys(storedLmLogits).length} tokens`
859
+ ? `✅ lm_logits: ${Object.keys(storedLmLogits).length} words`
368
860
  : '❌ No lm_logits',
369
861
  );
370
862
  if (storedLmLogits) {