@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 +37 -11
- package/dist/index.cjs.map +1 -1
- package/dist/index.d.cts +14 -1
- package/dist/index.d.ts +14 -1
- package/dist/index.js +36 -11
- package/dist/index.js.map +1 -1
- package/package.json +2 -2
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
|
|
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
|
-
|
|
116
|
-
|
|
117
|
-
|
|
118
|
-
maskSum
|
|
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]
|
|
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
|
|
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
|
|
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,
|