fix(AddonContext): getEmbedding returned the last token's embedding for every requested index; add multi-position getEmbeddings - #665
Conversation
…ng for every requested index; add multi-position `getEmbeddings` `getEmbedding(inputTokensLength)` ignored `inputTokensLength` on the `POOLING_TYPE_NONE` path and always read output index `-1`, so every hidden-state readout of an earlier token silently returned the LAST token's embedding. Introduced in withcatai#645 (v3.22.0), which changed `llama_get_embeddings_ith(ctx, inputTokensLength - 1)` to `llama_get_embeddings_ith(ctx, -1)`. For the ranking chunking feature that call site intended (`getEmbedding(input.length, 1)`) both resolve to the same token, but the change broke any other use of the parameter: with `-1` hard-coded, `getEmbedding(n)` for any `n` returns the same vector. Verified on bge-small-en-v1.5: reading "banana" vs "cherry" from one decode returned identical vectors on v3.22.1. Fix: index `-inputTokensLength`, i.e. interpret the parameter as "tokens from the end of the last decoded batch" (1 = last token), which matches `LlamaEmbeddingContext`/`LlamaRankingContext` usage (their batch ends with the sequence's last token) and restores per-token access. Also adds `getEmbeddings(positions: Uint32Array, maxVectorSize?)`: read many token states from one decode in a single call, where each position is an output index among the tokens `addToBatch` marked via `tokenLogitIndexes`. Prefill-only consumers (classification / reward heads / embedding probes that score specific token positions) need this to avoid re-decoding the same batch once per readout position — with the fixed `getEmbedding` that's the only way left to read positions other than the batch tail. Co-Authored-By: Claude Code <noreply@anthropic.com>
ae3237a to
0668053
Compare
|
Friendly ping @giladgd — this fixes a silent regression from #645 (the getEmbedding index change): on POOLING_TYPE_NONE contexts, getEmbedding(n) returns the last token's embedding for every n, so hidden-state readouts of earlier tokens silently get wrong data. Verified against a stock v3.22.1 prebuilt binary and a from-source build with the fix (repro + numbers in the PR description). Happy to adjust anything if needed! |
|
@MaskerPRC I don't think I understand what's the issue you're describing here. Is there a test that you can add that only uses the public |
Description of change
getEmbedding(inputTokensLength)on thePOOLING_TYPE_NONEpath ignoredinputTokensLengthand always read output index-1, so every hidden-state readout of an earlier token silently returned the LAST token's embedding.Introduced in #645 (v3.22.0), which changed
llama_get_embeddings_ith(ctx, inputTokensLength - 1)tollama_get_embeddings_ith(ctx, -1). For the ranking-chunking call site that change was intended for (getEmbedding(input.length, 1)), both indexes resolve to the same token — but it silently broke the parameter for every other use: with-1hard-coded,getEmbedding(n)returns the same vector for anyn.Reproduced on
bge-small-en-v1.5-q8_0(v3.22.1): decoding "apple banana cherry" once and reading "banana" vs "cherry" throughgetEmbeddingreturns identical vectors. Downstream effect seen in practice: a prefill-only classifier that scores specific token positions (a decision-model readout head) dropped from ~87% to ~47% accuracy against a reference implementation, purely from reading the wrong token's state.Fix: index
-inputTokensLength, i.e. interpret the parameter as "tokens from the end of the last decoded batch" (1 = last token). This matches howLlamaEmbeddingContextandLlamaRankingContextuse it (their batch ends with the sequence's last token) and restores per-token access.Also adds
getEmbeddings(positions: Uint32Array, maxVectorSize?): read many token states from one decoded batch in a single call. Each position is an output index — the i-th entry among the tokens the lastaddToBatchmarked viatokenLogitIndexes. Motivation: prefill-only consumers (classification / reward / readout heads that score several token positions per forward pass) currently have to re-decode the same batch once per readout position, which is the dominant cost of such workloads (measured ~2x on a 4-row batch, more with more readout positions).Verification
test/modelDependent/bge/embeddingByTokenIndex.test.ts):getEmbeddingsrows match direct single readouts; out-of-range and empty positions throw.npx eslintandnpm run test:typescriptpass.Pull-Request Checklist
masterbranchnpm run formatto apply eslint formattingnpm run testpasses with this change (eslint + tsc pass; the two new model-dependent tests pass locally — they need local model files, same as the existing bge tests)Fixes #0000(no existing issue found; happy to link one if it exists)