@atlaskit/editor-plugin-autocomplete 0.4.0 → 2.0.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/CHANGELOG.md +12 -0
- package/dist/cjs/pm-plugins/scoring-pipeline.js +3 -3
- package/dist/es2019/pm-plugins/scoring-pipeline.js +3 -3
- package/dist/esm/pm-plugins/scoring-pipeline.js +3 -3
- package/package.json +2 -2
- package/src/autocompletePlugin.tsx +1 -4
- package/src/pm-plugins/scoring-pipeline.ts +201 -204
package/CHANGELOG.md
CHANGED
|
@@ -74,7 +74,7 @@ function scoreStage1(candidate, contextVector, getWordVector, maxTenantFreq) {
|
|
|
74
74
|
var maxPossibleLog = Math.log10(maxTenantFreq + 1);
|
|
75
75
|
|
|
76
76
|
// Diversity Adjustment
|
|
77
|
-
var diversityRaw = (Math.log10(candidate.tenantFreq + 1) * 0.
|
|
77
|
+
var diversityRaw = (Math.log10(candidate.tenantFreq + 1) * 0.5 + Math.log10(candidate.docFreq + 1) * 0.25 + Math.log10(candidate.authorFreq + 1) * 0.25) / maxPossibleLog;
|
|
78
78
|
var sessionMultiplier = candidate.sessionFreq > 0 ? 1 + Math.log10(candidate.sessionFreq + 1) * 2.5 : 1;
|
|
79
79
|
|
|
80
80
|
// Apply multiplier; capped at L1_SESSION_CAP (default 1.2) to prevent excessive over-indexing
|
|
@@ -97,8 +97,8 @@ function scoreStage1(candidate, contextVector, getWordVector, maxTenantFreq) {
|
|
|
97
97
|
|
|
98
98
|
/**
|
|
99
99
|
* GHOST POS DICTIONARY
|
|
100
|
-
* A hardcoded mapping of common structural English words that were stripped
|
|
101
|
-
* from the main domain vocabulary. This allows the grammar filter to understand
|
|
100
|
+
* A hardcoded mapping of common structural English words that were stripped
|
|
101
|
+
* from the main domain vocabulary. This allows the grammar filter to understand
|
|
102
102
|
* context without suggesting these words to the user.
|
|
103
103
|
*/
|
|
104
104
|
var ghostPosTags = _ghost_pos_tags.default;
|
|
@@ -61,7 +61,7 @@ function scoreStage1(candidate, contextVector, getWordVector, maxTenantFreq) {
|
|
|
61
61
|
const maxPossibleLog = Math.log10(maxTenantFreq + 1);
|
|
62
62
|
|
|
63
63
|
// Diversity Adjustment
|
|
64
|
-
const diversityRaw = (Math.log10(candidate.tenantFreq + 1) * 0.
|
|
64
|
+
const diversityRaw = (Math.log10(candidate.tenantFreq + 1) * 0.5 + Math.log10(candidate.docFreq + 1) * 0.25 + Math.log10(candidate.authorFreq + 1) * 0.25) / maxPossibleLog;
|
|
65
65
|
const sessionMultiplier = candidate.sessionFreq > 0 ? 1 + Math.log10(candidate.sessionFreq + 1) * 2.5 : 1;
|
|
66
66
|
|
|
67
67
|
// Apply multiplier; capped at L1_SESSION_CAP (default 1.2) to prevent excessive over-indexing
|
|
@@ -84,8 +84,8 @@ function scoreStage1(candidate, contextVector, getWordVector, maxTenantFreq) {
|
|
|
84
84
|
|
|
85
85
|
/**
|
|
86
86
|
* GHOST POS DICTIONARY
|
|
87
|
-
* A hardcoded mapping of common structural English words that were stripped
|
|
88
|
-
* from the main domain vocabulary. This allows the grammar filter to understand
|
|
87
|
+
* A hardcoded mapping of common structural English words that were stripped
|
|
88
|
+
* from the main domain vocabulary. This allows the grammar filter to understand
|
|
89
89
|
* context without suggesting these words to the user.
|
|
90
90
|
*/
|
|
91
91
|
const ghostPosTags = ghostPosTagsData;
|
|
@@ -70,7 +70,7 @@ function scoreStage1(candidate, contextVector, getWordVector, maxTenantFreq) {
|
|
|
70
70
|
var maxPossibleLog = Math.log10(maxTenantFreq + 1);
|
|
71
71
|
|
|
72
72
|
// Diversity Adjustment
|
|
73
|
-
var diversityRaw = (Math.log10(candidate.tenantFreq + 1) * 0.
|
|
73
|
+
var diversityRaw = (Math.log10(candidate.tenantFreq + 1) * 0.5 + Math.log10(candidate.docFreq + 1) * 0.25 + Math.log10(candidate.authorFreq + 1) * 0.25) / maxPossibleLog;
|
|
74
74
|
var sessionMultiplier = candidate.sessionFreq > 0 ? 1 + Math.log10(candidate.sessionFreq + 1) * 2.5 : 1;
|
|
75
75
|
|
|
76
76
|
// Apply multiplier; capped at L1_SESSION_CAP (default 1.2) to prevent excessive over-indexing
|
|
@@ -93,8 +93,8 @@ function scoreStage1(candidate, contextVector, getWordVector, maxTenantFreq) {
|
|
|
93
93
|
|
|
94
94
|
/**
|
|
95
95
|
* GHOST POS DICTIONARY
|
|
96
|
-
* A hardcoded mapping of common structural English words that were stripped
|
|
97
|
-
* from the main domain vocabulary. This allows the grammar filter to understand
|
|
96
|
+
* A hardcoded mapping of common structural English words that were stripped
|
|
97
|
+
* from the main domain vocabulary. This allows the grammar filter to understand
|
|
98
98
|
* context without suggesting these words to the user.
|
|
99
99
|
*/
|
|
100
100
|
var ghostPosTags = ghostPosTagsData;
|
package/package.json
CHANGED
|
@@ -1,6 +1,6 @@
|
|
|
1
1
|
{
|
|
2
2
|
"name": "@atlaskit/editor-plugin-autocomplete",
|
|
3
|
-
"version": "0.
|
|
3
|
+
"version": "2.0.0",
|
|
4
4
|
"description": "Client-side text autocomplete plugin for @atlaskit/editor-core",
|
|
5
5
|
"author": "Atlassian Pty Ltd",
|
|
6
6
|
"license": "Apache-2.0",
|
|
@@ -33,7 +33,7 @@
|
|
|
33
33
|
"wink-nlp": "^2.4.0"
|
|
34
34
|
},
|
|
35
35
|
"peerDependencies": {
|
|
36
|
-
"@atlaskit/editor-common": "^
|
|
36
|
+
"@atlaskit/editor-common": "^114.0.0",
|
|
37
37
|
"react": "^18.2.0"
|
|
38
38
|
},
|
|
39
39
|
"techstack": {
|
|
@@ -1,8 +1,5 @@
|
|
|
1
1
|
import type { AutocompletePlugin } from './autocompletePluginType';
|
|
2
|
-
import {
|
|
3
|
-
autocompletePluginKey,
|
|
4
|
-
createAutocompletePlugin,
|
|
5
|
-
} from './pm-plugins/autocomplete-plugin';
|
|
2
|
+
import { autocompletePluginKey, createAutocompletePlugin } from './pm-plugins/autocomplete-plugin';
|
|
6
3
|
import type { AutocompletePluginState } from './pm-plugins/autocomplete-plugin';
|
|
7
4
|
|
|
8
5
|
export const autocompletePlugin: AutocompletePlugin = ({ config: options }) => {
|
|
@@ -12,36 +12,36 @@ import grammarTransitionsData from './data/grammar_transitions_10k.json';
|
|
|
12
12
|
// ─── Types ──────────────────────────────────────────────────
|
|
13
13
|
|
|
14
14
|
export interface ScoringCandidate {
|
|
15
|
-
|
|
16
|
-
|
|
17
|
-
|
|
18
|
-
|
|
19
|
-
|
|
15
|
+
authorFreq: number;
|
|
16
|
+
docFreq: number;
|
|
17
|
+
sessionFreq: number;
|
|
18
|
+
tenantFreq: number;
|
|
19
|
+
word: string;
|
|
20
20
|
}
|
|
21
21
|
|
|
22
22
|
export interface ScoredCandidate {
|
|
23
|
-
|
|
24
|
-
|
|
25
|
-
|
|
26
|
-
|
|
27
|
-
|
|
23
|
+
finalScore: number;
|
|
24
|
+
freqScore: number;
|
|
25
|
+
lmScore: number;
|
|
26
|
+
semanticScore: number;
|
|
27
|
+
word: string;
|
|
28
28
|
}
|
|
29
29
|
|
|
30
30
|
/** Metadata returned by the grammar filter for debug logging in the caller. */
|
|
31
31
|
export interface GrammarFilterMeta {
|
|
32
|
-
|
|
33
|
-
|
|
34
|
-
|
|
35
|
-
|
|
36
|
-
|
|
32
|
+
after: number;
|
|
33
|
+
before: number;
|
|
34
|
+
dropped: string[];
|
|
35
|
+
prevTags: string[];
|
|
36
|
+
prevWord: string;
|
|
37
37
|
}
|
|
38
38
|
|
|
39
39
|
interface PosTransitionRule {
|
|
40
|
-
|
|
40
|
+
allowed: string[];
|
|
41
41
|
}
|
|
42
42
|
|
|
43
43
|
interface GrammarTransitions {
|
|
44
|
-
|
|
44
|
+
transitions: Record<string, PosTransitionRule>;
|
|
45
45
|
}
|
|
46
46
|
|
|
47
47
|
// ─── Scoring Constants ──────────────────────────────────────
|
|
@@ -57,7 +57,7 @@ const L1_SESSION_CAP = 1.2;
|
|
|
57
57
|
// ─── Grammar Data (loaded once on import) ───────────────────
|
|
58
58
|
|
|
59
59
|
const posTags: Map<string, string[]> = new Map(
|
|
60
|
-
|
|
60
|
+
Object.entries(posTagsData as Record<string, string[]>),
|
|
61
61
|
);
|
|
62
62
|
|
|
63
63
|
const grammarTransitions = grammarTransitionsData as GrammarTransitions;
|
|
@@ -68,227 +68,224 @@ const grammarTransitions = grammarTransitionsData as GrammarTransitions;
|
|
|
68
68
|
* never re-iterates the transition rules per call.
|
|
69
69
|
*/
|
|
70
70
|
const precomputedAllowedByPos: Map<string, Set<string>> = new Map(
|
|
71
|
-
Object.entries(grammarTransitions.transitions).map(([pos, rule]) => [
|
|
72
|
-
pos,
|
|
73
|
-
new Set(rule.allowed),
|
|
74
|
-
]),
|
|
71
|
+
Object.entries(grammarTransitions.transitions).map(([pos, rule]) => [pos, new Set(rule.allowed)]),
|
|
75
72
|
);
|
|
76
73
|
|
|
77
74
|
// ─── Math ───────────────────────────────────────────────────
|
|
78
75
|
|
|
79
76
|
function cosineSimilarity(a: Float32Array, b: Float32Array): number {
|
|
80
|
-
|
|
81
|
-
|
|
82
|
-
|
|
83
|
-
|
|
84
|
-
|
|
85
|
-
|
|
86
|
-
|
|
87
|
-
|
|
88
|
-
|
|
89
|
-
|
|
90
|
-
|
|
91
|
-
|
|
92
|
-
|
|
93
|
-
|
|
77
|
+
let dot = 0;
|
|
78
|
+
let normA = 0;
|
|
79
|
+
let normB = 0;
|
|
80
|
+
for (let i = 0; i < a.length; i++) {
|
|
81
|
+
dot += a[i] * b[i];
|
|
82
|
+
normA += a[i] * a[i];
|
|
83
|
+
normB += b[i] * b[i];
|
|
84
|
+
}
|
|
85
|
+
const dNormA = Math.sqrt(normA);
|
|
86
|
+
const dNormB = Math.sqrt(normB);
|
|
87
|
+
if (dNormA === 0 || dNormB === 0) {
|
|
88
|
+
return NEUTRAL_SCORE;
|
|
89
|
+
}
|
|
90
|
+
return (1 + dot / (dNormA * dNormB)) / 2;
|
|
94
91
|
}
|
|
95
92
|
|
|
96
93
|
// ─── Stage 1: Semantic + Frequency ──────────────────────────
|
|
97
94
|
|
|
98
95
|
function scoreStage1(
|
|
99
|
-
|
|
100
|
-
|
|
101
|
-
|
|
102
|
-
|
|
96
|
+
candidate: ScoringCandidate,
|
|
97
|
+
contextVector: Float32Array | null,
|
|
98
|
+
getWordVector: (word: string) => Float32Array | null,
|
|
99
|
+
maxTenantFreq: number,
|
|
103
100
|
): { freqScore: number; semanticScore: number; stage1Score: number } {
|
|
104
|
-
|
|
105
101
|
// 1. Calculate Base Global Score (Normalized Log)
|
|
106
|
-
|
|
107
|
-
|
|
108
|
-
|
|
109
|
-
|
|
110
|
-
|
|
111
|
-
|
|
112
|
-
|
|
113
|
-
|
|
114
|
-
|
|
115
|
-
|
|
116
|
-
|
|
117
|
-
|
|
118
|
-
|
|
119
|
-
|
|
120
|
-
|
|
121
|
-
|
|
122
|
-
|
|
123
|
-
|
|
124
|
-
|
|
125
|
-
|
|
126
|
-
|
|
127
|
-
|
|
128
|
-
|
|
129
|
-
|
|
130
|
-
|
|
131
|
-
|
|
132
|
-
|
|
133
|
-
};
|
|
102
|
+
const maxPossibleLog = Math.log10(maxTenantFreq + 1);
|
|
103
|
+
|
|
104
|
+
// Diversity Adjustment
|
|
105
|
+
const diversityRaw =
|
|
106
|
+
(Math.log10(candidate.tenantFreq + 1) * 0.5 +
|
|
107
|
+
Math.log10(candidate.docFreq + 1) * 0.25 +
|
|
108
|
+
Math.log10(candidate.authorFreq + 1) * 0.25) /
|
|
109
|
+
maxPossibleLog;
|
|
110
|
+
|
|
111
|
+
const sessionMultiplier =
|
|
112
|
+
candidate.sessionFreq > 0 ? 1 + Math.log10(candidate.sessionFreq + 1) * 2.5 : 1;
|
|
113
|
+
|
|
114
|
+
// Apply multiplier; capped at L1_SESSION_CAP (default 1.2) to prevent excessive over-indexing
|
|
115
|
+
const freqScore = Math.min(diversityRaw * sessionMultiplier, L1_SESSION_CAP);
|
|
116
|
+
|
|
117
|
+
// 3. Semantic Scoring
|
|
118
|
+
let semanticScore = NEUTRAL_SCORE;
|
|
119
|
+
if (contextVector) {
|
|
120
|
+
const wordVec = getWordVector(candidate.word);
|
|
121
|
+
semanticScore = wordVec ? cosineSimilarity(contextVector, wordVec) : NEUTRAL_SCORE;
|
|
122
|
+
}
|
|
123
|
+
|
|
124
|
+
return {
|
|
125
|
+
semanticScore,
|
|
126
|
+
freqScore,
|
|
127
|
+
stage1Score: ALPHA * semanticScore + BETA * freqScore,
|
|
128
|
+
};
|
|
134
129
|
}
|
|
135
130
|
|
|
136
131
|
// ─── Grammar Filter ─────────────────────────────────────────
|
|
137
132
|
|
|
138
133
|
/**
|
|
139
134
|
* GHOST POS DICTIONARY
|
|
140
|
-
* A hardcoded mapping of common structural English words that were stripped
|
|
141
|
-
* from the main domain vocabulary. This allows the grammar filter to understand
|
|
135
|
+
* A hardcoded mapping of common structural English words that were stripped
|
|
136
|
+
* from the main domain vocabulary. This allows the grammar filter to understand
|
|
142
137
|
* context without suggesting these words to the user.
|
|
143
138
|
*/
|
|
144
139
|
const ghostPosTags: Record<string, string[]> = ghostPosTagsData as Record<string, string[]>;
|
|
145
140
|
|
|
146
|
-
type FilterEntry = {
|
|
141
|
+
type FilterEntry = {
|
|
142
|
+
candidate: ScoringCandidate;
|
|
143
|
+
freqScore: number;
|
|
144
|
+
semanticScore: number;
|
|
145
|
+
stage1Score: number;
|
|
146
|
+
};
|
|
147
147
|
|
|
148
148
|
function applyGrammarFilter(
|
|
149
|
-
|
|
150
|
-
|
|
149
|
+
candidates: FilterEntry[],
|
|
150
|
+
previousWord: string,
|
|
151
151
|
): { filtered: FilterEntry[]; grammarMeta: GrammarFilterMeta | null } {
|
|
152
|
-
|
|
153
|
-
|
|
154
|
-
|
|
155
|
-
|
|
156
|
-
|
|
157
|
-
|
|
158
|
-
|
|
159
|
-
|
|
160
|
-
|
|
161
|
-
|
|
162
|
-
|
|
163
|
-
|
|
164
|
-
|
|
165
|
-
|
|
166
|
-
|
|
167
|
-
|
|
168
|
-
|
|
169
|
-
|
|
170
|
-
|
|
171
|
-
|
|
172
|
-
|
|
173
|
-
|
|
174
|
-
|
|
175
|
-
|
|
176
|
-
|
|
177
|
-
|
|
178
|
-
|
|
179
|
-
|
|
180
|
-
|
|
181
|
-
|
|
182
|
-
|
|
183
|
-
|
|
184
|
-
|
|
185
|
-
|
|
186
|
-
|
|
187
|
-
|
|
188
|
-
|
|
189
|
-
|
|
190
|
-
|
|
191
|
-
|
|
192
|
-
|
|
193
|
-
|
|
194
|
-
|
|
195
|
-
|
|
196
|
-
|
|
197
|
-
|
|
198
|
-
|
|
199
|
-
|
|
200
|
-
|
|
201
|
-
|
|
202
|
-
|
|
203
|
-
|
|
152
|
+
if (!previousWord) return { filtered: candidates, grammarMeta: null };
|
|
153
|
+
|
|
154
|
+
const lowerPrev = previousWord.toLowerCase();
|
|
155
|
+
const prevTags = ghostPosTags[lowerPrev] || posTags.get(lowerPrev);
|
|
156
|
+
|
|
157
|
+
if (!prevTags || prevTags.length === 0) {
|
|
158
|
+
return { filtered: candidates, grammarMeta: null };
|
|
159
|
+
}
|
|
160
|
+
|
|
161
|
+
let allowedNextTags: Set<string>;
|
|
162
|
+
if (prevTags.length === 1) {
|
|
163
|
+
// Common case: single POS tag — reuse the precomputed Set directly (no allocation)
|
|
164
|
+
allowedNextTags = precomputedAllowedByPos.get(prevTags[0]) ?? new Set();
|
|
165
|
+
} else {
|
|
166
|
+
allowedNextTags = new Set<string>();
|
|
167
|
+
for (const pt of prevTags) {
|
|
168
|
+
const allowed = precomputedAllowedByPos.get(pt);
|
|
169
|
+
if (allowed) allowed.forEach((tag) => allowedNextTags.add(tag));
|
|
170
|
+
}
|
|
171
|
+
}
|
|
172
|
+
|
|
173
|
+
const filtered: FilterEntry[] = [];
|
|
174
|
+
const dropped: string[] = [];
|
|
175
|
+
|
|
176
|
+
for (const entry of candidates) {
|
|
177
|
+
const candidateTags = posTags.get(entry.candidate.word.toLowerCase());
|
|
178
|
+
|
|
179
|
+
// If candidate has no tags (unknown word), let it pass to be safe
|
|
180
|
+
if (!candidateTags || candidateTags.length === 0) {
|
|
181
|
+
filtered.push(entry);
|
|
182
|
+
continue;
|
|
183
|
+
}
|
|
184
|
+
|
|
185
|
+
if (candidateTags.some((ct) => allowedNextTags.has(ct))) {
|
|
186
|
+
filtered.push(entry);
|
|
187
|
+
} else {
|
|
188
|
+
dropped.push(entry.candidate.word);
|
|
189
|
+
}
|
|
190
|
+
}
|
|
191
|
+
|
|
192
|
+
const finalFiltered = filtered.length > 0 ? filtered : candidates;
|
|
193
|
+
|
|
194
|
+
return {
|
|
195
|
+
filtered: finalFiltered,
|
|
196
|
+
grammarMeta: {
|
|
197
|
+
prevWord: lowerPrev,
|
|
198
|
+
prevTags,
|
|
199
|
+
before: candidates.length,
|
|
200
|
+
after: finalFiltered.length,
|
|
201
|
+
dropped: filtered.length > 0 ? dropped : [],
|
|
202
|
+
},
|
|
203
|
+
};
|
|
204
204
|
}
|
|
205
205
|
|
|
206
206
|
// ─── Stage 2: LM Re-ranking ────────────────────────────────
|
|
207
|
-
function getLmScore(
|
|
208
|
-
|
|
209
|
-
|
|
210
|
-
|
|
211
|
-
|
|
212
|
-
|
|
213
|
-
|
|
214
|
-
|
|
215
|
-
|
|
216
|
-
return val;
|
|
217
|
-
}
|
|
218
|
-
return 0;
|
|
207
|
+
function getLmScore(word: string, lmLogits: Record<string, number> | null): number {
|
|
208
|
+
if (!lmLogits) return 0;
|
|
209
|
+
|
|
210
|
+
// Look up the word directly! No more tokens.
|
|
211
|
+
const val = lmLogits[word.toLowerCase()];
|
|
212
|
+
if (typeof val === 'number') {
|
|
213
|
+
return val;
|
|
214
|
+
}
|
|
215
|
+
return 0;
|
|
219
216
|
}
|
|
220
217
|
|
|
221
218
|
// ─── Public API ─────────────────────────────────────────────
|
|
222
219
|
|
|
223
220
|
export interface RankCandidatesResult {
|
|
224
|
-
|
|
225
|
-
|
|
221
|
+
candidates: ScoredCandidate[];
|
|
222
|
+
grammarMeta: GrammarFilterMeta | null;
|
|
226
223
|
}
|
|
227
224
|
|
|
228
225
|
export function rankCandidates(
|
|
229
|
-
|
|
230
|
-
|
|
231
|
-
|
|
232
|
-
|
|
233
|
-
|
|
234
|
-
|
|
226
|
+
candidates: ScoringCandidate[],
|
|
227
|
+
contextVector: Float32Array | null,
|
|
228
|
+
getWordVector: (word: string) => Float32Array | null,
|
|
229
|
+
lmLogits: Record<string, number> | null,
|
|
230
|
+
maxTenantFreq: number,
|
|
231
|
+
previousWord: string,
|
|
235
232
|
): RankCandidatesResult {
|
|
236
|
-
|
|
237
|
-
|
|
238
|
-
|
|
239
|
-
|
|
240
|
-
|
|
241
|
-
|
|
242
|
-
|
|
243
|
-
|
|
244
|
-
|
|
245
|
-
|
|
246
|
-
|
|
247
|
-
const stage1Survivors = stage1Results.filter(entry => entry.stage1Score >= MIN_STAGE1_SCORE);
|
|
248
|
-
|
|
249
|
-
|
|
250
|
-
|
|
251
|
-
|
|
252
|
-
|
|
253
|
-
|
|
254
|
-
|
|
255
|
-
|
|
256
|
-
|
|
257
|
-
|
|
258
|
-
|
|
259
|
-
|
|
260
|
-
|
|
261
|
-
|
|
262
|
-
|
|
263
|
-
|
|
264
|
-
|
|
265
|
-
|
|
266
|
-
|
|
267
|
-
|
|
268
|
-
|
|
269
|
-
|
|
270
|
-
|
|
271
|
-
|
|
272
|
-
|
|
273
|
-
|
|
274
|
-
|
|
275
|
-
|
|
276
|
-
|
|
277
|
-
|
|
278
|
-
|
|
279
|
-
|
|
280
|
-
|
|
281
|
-
|
|
282
|
-
|
|
283
|
-
|
|
284
|
-
|
|
285
|
-
|
|
233
|
+
// Stage 1
|
|
234
|
+
const stage1Results = candidates.map((candidate) => {
|
|
235
|
+
const { semanticScore, freqScore, stage1Score } = scoreStage1(
|
|
236
|
+
candidate,
|
|
237
|
+
contextVector,
|
|
238
|
+
getWordVector,
|
|
239
|
+
maxTenantFreq,
|
|
240
|
+
);
|
|
241
|
+
return { candidate, semanticScore, freqScore, stage1Score };
|
|
242
|
+
});
|
|
243
|
+
|
|
244
|
+
const stage1Survivors = stage1Results.filter((entry) => entry.stage1Score >= MIN_STAGE1_SCORE);
|
|
245
|
+
|
|
246
|
+
// Grammar Filter
|
|
247
|
+
const { filtered, grammarMeta } = applyGrammarFilter(stage1Survivors, previousWord);
|
|
248
|
+
|
|
249
|
+
// Stage 2 + final assembly
|
|
250
|
+
let lmMax = 0;
|
|
251
|
+
if (lmLogits && Object.keys(lmLogits).length > 0) {
|
|
252
|
+
const values = Object.values(lmLogits);
|
|
253
|
+
lmMax = Math.max(...values);
|
|
254
|
+
}
|
|
255
|
+
|
|
256
|
+
const scored: ScoredCandidate[] = filtered.map((entry) => {
|
|
257
|
+
let lmScore = 0;
|
|
258
|
+
let finalScore = entry.stage1Score;
|
|
259
|
+
|
|
260
|
+
if (lmLogits && Object.keys(lmLogits).length > 0) {
|
|
261
|
+
const rawLm = getLmScore(entry.candidate.word, lmLogits);
|
|
262
|
+
|
|
263
|
+
if (rawLm !== 0) {
|
|
264
|
+
// The word was in the top_k! Score it normally.
|
|
265
|
+
const logitDiff = Math.log(rawLm) - Math.log(lmMax);
|
|
266
|
+
lmScore = Math.exp(logitDiff);
|
|
267
|
+
} else {
|
|
268
|
+
lmScore = 0.05;
|
|
269
|
+
}
|
|
270
|
+
|
|
271
|
+
finalScore = STAGE1_WEIGHT * entry.stage1Score + STAGE2_WEIGHT * lmScore;
|
|
272
|
+
}
|
|
273
|
+
|
|
274
|
+
return {
|
|
275
|
+
word: entry.candidate.word,
|
|
276
|
+
freqScore: entry.freqScore,
|
|
277
|
+
semanticScore: entry.semanticScore,
|
|
278
|
+
lmScore,
|
|
279
|
+
finalScore,
|
|
280
|
+
};
|
|
281
|
+
});
|
|
282
|
+
|
|
286
283
|
scored.sort((a, b) => {
|
|
287
|
-
|
|
288
|
-
|
|
289
|
-
|
|
290
|
-
|
|
291
|
-
|
|
292
|
-
|
|
293
|
-
|
|
294
|
-
}
|
|
284
|
+
if (b.finalScore !== a.finalScore) {
|
|
285
|
+
return b.finalScore - a.finalScore;
|
|
286
|
+
}
|
|
287
|
+
return a.word.length - b.word.length;
|
|
288
|
+
});
|
|
289
|
+
|
|
290
|
+
return { candidates: scored, grammarMeta };
|
|
291
|
+
}
|