@amemhq/core 1.0.0 → 1.0.1

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/dist/index.cjs CHANGED
@@ -49,6 +49,7 @@ __export(index_exports, {
49
49
  generateReviewBatch: () => generateReviewBatch,
50
50
  getEmbeddingDim: () => getEmbeddingDim,
51
51
  getEmbeddingModel: () => getEmbeddingModel,
52
+ getEmbeddingPooling: () => getEmbeddingPooling,
52
53
  getNote: () => getNote,
53
54
  invalidateNote: () => invalidateNote,
54
55
  isModelLoaded: () => isModelLoaded,
@@ -88,6 +89,25 @@ var DEFAULT_EMBEDDING_MODEL = "Xenova/paraphrase-multilingual-MiniLM-L12-v2";
88
89
  function getEmbeddingModel() {
89
90
  return process.env.AMEM_EMBED_MODEL?.trim() || DEFAULT_EMBEDDING_MODEL;
90
91
  }
92
+ var CLS_POOLED_MODELS = /* @__PURE__ */ new Set([
93
+ "bge-m3",
94
+ "bge-base-zh-v1.5",
95
+ "bge-small-zh-v1.5",
96
+ "bge-base-en-v1.5",
97
+ "bge-small-en-v1.5",
98
+ "bge-large-en-v1.5",
99
+ "gte-multilingual-base",
100
+ "gte-modernbert-base",
101
+ "gte-large-en-v1.5",
102
+ "snowflake-arctic-embed-m",
103
+ "snowflake-arctic-embed-l"
104
+ ]);
105
+ function getEmbeddingPooling() {
106
+ const explicit = process.env.AMEM_EMBED_POOLING?.trim().toLowerCase();
107
+ if (explicit === "mean" || explicit === "cls") return explicit;
108
+ const basename2 = getEmbeddingModel().split("/").pop()?.toLowerCase() ?? "";
109
+ return CLS_POOLED_MODELS.has(basename2) ? "cls" : "mean";
110
+ }
91
111
  async function getExtractor() {
92
112
  const wanted = getEmbeddingModel();
93
113
  if (extractor && loadedModelName === wanted) return extractor;
@@ -108,21 +128,25 @@ async function getEmbeddingDim() {
108
128
  cachedDim = probe.length;
109
129
  return cachedDim;
110
130
  }
111
- function meanPoolingNormalize(output, attentionMask) {
131
+ function poolNormalize(output, attentionMask, mode) {
112
132
  const seqLen = output.length;
113
133
  const dim = output[0].length;
114
134
  const pooled = new Array(dim).fill(0);
115
- let maskSum = 0;
116
- for (let i = 0; i < seqLen; i++) {
117
- const m = attentionMask[i];
118
- maskSum += m;
135
+ if (mode === "cls") {
136
+ for (let j = 0; j < dim; j++) pooled[j] = output[0][j];
137
+ } else {
138
+ let maskSum = 0;
139
+ for (let i = 0; i < seqLen; i++) {
140
+ const m = attentionMask[i];
141
+ maskSum += m;
142
+ for (let j = 0; j < dim; j++) {
143
+ pooled[j] += output[i][j] * m;
144
+ }
145
+ }
119
146
  for (let j = 0; j < dim; j++) {
120
- pooled[j] += output[i][j] * m;
147
+ pooled[j] /= Math.max(maskSum, 1e-9);
121
148
  }
122
149
  }
123
- for (let j = 0; j < dim; j++) {
124
- pooled[j] /= Math.max(maskSum, 1e-9);
125
- }
126
150
  let norm = 0;
127
151
  for (const v of pooled) norm += v * v;
128
152
  norm = Math.sqrt(norm);
@@ -130,7 +154,8 @@ function meanPoolingNormalize(output, attentionMask) {
130
154
  }
131
155
  async function encode(text) {
132
156
  const ext = await getExtractor();
133
- const result = await ext(text, { pooling: "mean", normalize: true });
157
+ const pooling = getEmbeddingPooling();
158
+ const result = await ext(text, { pooling, normalize: true });
134
159
  if (result && result.data) {
135
160
  return Array.from(result.data);
136
161
  }
@@ -146,7 +171,7 @@ async function encode(text) {
146
171
  }
147
172
  raw.push(row);
148
173
  }
149
- return meanPoolingNormalize(raw, new Array(seqLen).fill(1));
174
+ return poolNormalize(raw, new Array(seqLen).fill(1), pooling);
150
175
  }
151
176
  throw new Error("Unexpected embedding output shape");
152
177
  }
@@ -2349,6 +2374,7 @@ function isPlausibleUpdateTarget(newEmbedding, targetEmbedding, minSimilarity) {
2349
2374
  generateReviewBatch,
2350
2375
  getEmbeddingDim,
2351
2376
  getEmbeddingModel,
2377
+ getEmbeddingPooling,
2352
2378
  getNote,
2353
2379
  invalidateNote,
2354
2380
  isModelLoaded,