@huggingface/transformers-structured-output 4.3.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/LICENSE +202 -0
- package/README.md +56 -0
- package/dist/index.cjs +2425 -0
- package/dist/index.cjs.map +7 -0
- package/dist/index.js +2402 -0
- package/dist/index.js.map +7 -0
- package/package.json +65 -0
- package/src/StructuredOutputProcessor.ts +103 -0
- package/src/engine/constraint.ts +331 -0
- package/src/engine/index.ts +2 -0
- package/src/engine/json.ts +1853 -0
- package/src/engine/regex.ts +419 -0
- package/src/engine/tokenizer.ts +205 -0
- package/src/engine/types.ts +12 -0
- package/src/index.ts +2 -0
- package/src/utils/mask.ts +30 -0
- package/types/StructuredOutputProcessor.d.ts +22 -0
- package/types/StructuredOutputProcessor.d.ts.map +1 -0
- package/types/engine/constraint.d.ts +30 -0
- package/types/engine/constraint.d.ts.map +1 -0
- package/types/engine/index.d.ts +3 -0
- package/types/engine/index.d.ts.map +1 -0
- package/types/engine/json.d.ts +70 -0
- package/types/engine/json.d.ts.map +1 -0
- package/types/engine/regex.d.ts +3 -0
- package/types/engine/regex.d.ts.map +1 -0
- package/types/engine/tokenizer.d.ts +8 -0
- package/types/engine/tokenizer.d.ts.map +1 -0
- package/types/engine/types.d.ts +16 -0
- package/types/engine/types.d.ts.map +1 -0
- package/types/index.d.ts +3 -0
- package/types/index.d.ts.map +1 -0
- package/types/utils/mask.d.ts +3 -0
- package/types/utils/mask.d.ts.map +1 -0
|
@@ -0,0 +1,419 @@
|
|
|
1
|
+
import type { ConstraintState } from './types';
|
|
2
|
+
|
|
3
|
+
type Expression =
|
|
4
|
+
| { kind: 'empty' }
|
|
5
|
+
| { kind: 'set'; bytes: Uint32Array }
|
|
6
|
+
| { kind: 'choice'; choices: Expression[] }
|
|
7
|
+
| { kind: 'sequence'; parts: Expression[] }
|
|
8
|
+
| { kind: 'star'; child: Expression };
|
|
9
|
+
|
|
10
|
+
const EMPTY: Expression = { kind: 'empty' };
|
|
11
|
+
const encoder = new TextEncoder();
|
|
12
|
+
|
|
13
|
+
export function compileRegex(source: string): ConstraintState<number> {
|
|
14
|
+
const machine = new RegexMachine(new RegexParser(source).parse());
|
|
15
|
+
return {
|
|
16
|
+
initial: machine.initial,
|
|
17
|
+
transition: (state, byte) => machine.transition(state, byte),
|
|
18
|
+
viable: (state) => state >= 0,
|
|
19
|
+
accepting: (state) => machine.accepting(state),
|
|
20
|
+
maskKey: (state) => machine.stateKey(state),
|
|
21
|
+
};
|
|
22
|
+
}
|
|
23
|
+
|
|
24
|
+
class RegexParser {
|
|
25
|
+
private index = 0;
|
|
26
|
+
|
|
27
|
+
constructor(private readonly source: string) {}
|
|
28
|
+
|
|
29
|
+
parse(): Expression {
|
|
30
|
+
if (this.source.startsWith('^')) this.index++;
|
|
31
|
+
const expression = this.alternation();
|
|
32
|
+
if (this.peek() === '$' && this.index === this.source.length - 1) this.index++;
|
|
33
|
+
if (this.index !== this.source.length) this.fail(`unexpected ${JSON.stringify(this.peek())}`);
|
|
34
|
+
return expression;
|
|
35
|
+
}
|
|
36
|
+
|
|
37
|
+
private alternation(): Expression {
|
|
38
|
+
const choices = [this.concatenation()];
|
|
39
|
+
while (this.peek() === '|') {
|
|
40
|
+
this.index++;
|
|
41
|
+
choices.push(this.concatenation());
|
|
42
|
+
}
|
|
43
|
+
return choice(choices);
|
|
44
|
+
}
|
|
45
|
+
|
|
46
|
+
private concatenation(): Expression {
|
|
47
|
+
const parts: Expression[] = [];
|
|
48
|
+
while (this.index < this.source.length && this.peek() !== ')' && this.peek() !== '|') {
|
|
49
|
+
if (this.peek() === '$' && this.index === this.source.length - 1) break;
|
|
50
|
+
parts.push(this.quantified());
|
|
51
|
+
}
|
|
52
|
+
return sequence(parts);
|
|
53
|
+
}
|
|
54
|
+
|
|
55
|
+
private quantified(): Expression {
|
|
56
|
+
const atom = this.atom();
|
|
57
|
+
const quantifier = this.peek();
|
|
58
|
+
if (!['*', '+', '?', '{'].includes(quantifier ?? '')) return atom;
|
|
59
|
+
let minimum: number;
|
|
60
|
+
let maximum: number | null;
|
|
61
|
+
if (quantifier === '*') {
|
|
62
|
+
this.index++;
|
|
63
|
+
[minimum, maximum] = [0, null];
|
|
64
|
+
} else if (quantifier === '+') {
|
|
65
|
+
this.index++;
|
|
66
|
+
[minimum, maximum] = [1, null];
|
|
67
|
+
} else if (quantifier === '?') {
|
|
68
|
+
this.index++;
|
|
69
|
+
[minimum, maximum] = [0, 1];
|
|
70
|
+
} else {
|
|
71
|
+
[minimum, maximum] = this.bounds();
|
|
72
|
+
}
|
|
73
|
+
if (this.peek() === '?' || this.peek() === '+') this.fail('lazy and possessive quantifiers are unsupported');
|
|
74
|
+
const parts = Array.from({ length: minimum }, () => atom);
|
|
75
|
+
if (maximum === null) parts.push(star(atom));
|
|
76
|
+
else for (let count = minimum; count < maximum; ++count) parts.push(choice([EMPTY, atom]));
|
|
77
|
+
return sequence(parts);
|
|
78
|
+
}
|
|
79
|
+
|
|
80
|
+
private atom(): Expression {
|
|
81
|
+
const character = this.peek();
|
|
82
|
+
if (character === undefined) this.fail('expected an expression');
|
|
83
|
+
if (character === '(') {
|
|
84
|
+
this.index++;
|
|
85
|
+
if (this.source.startsWith('?:', this.index)) this.index += 2;
|
|
86
|
+
else if (this.peek() === '?') this.fail('only non-capturing special groups are supported');
|
|
87
|
+
const result = this.alternation();
|
|
88
|
+
if (this.peek() !== ')') this.fail('unterminated group');
|
|
89
|
+
this.index++;
|
|
90
|
+
return result;
|
|
91
|
+
}
|
|
92
|
+
if (character === '[') return this.characterClass();
|
|
93
|
+
if (character === '.') {
|
|
94
|
+
this.index++;
|
|
95
|
+
return byteSet(range(0, 255));
|
|
96
|
+
}
|
|
97
|
+
if (character === '\\') return this.escape(false).expression;
|
|
98
|
+
if ('*+?{})'.includes(character)) this.fail(`unexpected ${JSON.stringify(character)}`);
|
|
99
|
+
this.index += character.length;
|
|
100
|
+
return literal(character);
|
|
101
|
+
}
|
|
102
|
+
|
|
103
|
+
private characterClass(): Expression {
|
|
104
|
+
this.index++;
|
|
105
|
+
const negated = this.peek() === '^';
|
|
106
|
+
if (negated) this.index++;
|
|
107
|
+
const bytes = new Uint32Array(8);
|
|
108
|
+
let hasValue = false;
|
|
109
|
+
while (this.peek() !== ']') {
|
|
110
|
+
if (this.peek() === undefined) this.fail('unterminated character class');
|
|
111
|
+
const first = this.classValue();
|
|
112
|
+
if (this.peek() === '-' && this.source[this.index + 1] !== ']') {
|
|
113
|
+
this.index++;
|
|
114
|
+
const last = this.classValue();
|
|
115
|
+
if (first.single === undefined || last.single === undefined || first.single > last.single) {
|
|
116
|
+
this.fail('invalid character class range');
|
|
117
|
+
}
|
|
118
|
+
addRange(bytes, first.single, last.single);
|
|
119
|
+
} else {
|
|
120
|
+
union(bytes, first.bytes);
|
|
121
|
+
}
|
|
122
|
+
hasValue = true;
|
|
123
|
+
}
|
|
124
|
+
this.index++;
|
|
125
|
+
if (!hasValue) this.fail('empty character class');
|
|
126
|
+
if (negated) for (let word = 0; word < bytes.length; ++word) bytes[word] = ~bytes[word];
|
|
127
|
+
return byteSet(bytes);
|
|
128
|
+
}
|
|
129
|
+
|
|
130
|
+
private classValue(): { bytes: Uint32Array; single?: number } {
|
|
131
|
+
if (this.peek() === '\\') {
|
|
132
|
+
const escaped = this.escape(true);
|
|
133
|
+
if (escaped.bytes === undefined) this.fail('multi-byte escapes are unsupported in character classes');
|
|
134
|
+
return { bytes: escaped.bytes, single: escaped.single };
|
|
135
|
+
}
|
|
136
|
+
const character = this.peek()!;
|
|
137
|
+
this.index += character.length;
|
|
138
|
+
const encoded = encoder.encode(character);
|
|
139
|
+
if (encoded.length !== 1) this.fail('non-ASCII character classes are unsupported');
|
|
140
|
+
return { bytes: singleton(encoded[0]), single: encoded[0] };
|
|
141
|
+
}
|
|
142
|
+
|
|
143
|
+
private escape(inClass: boolean): { expression: Expression; bytes?: Uint32Array; single?: number } {
|
|
144
|
+
this.index++;
|
|
145
|
+
const code = this.peek();
|
|
146
|
+
if (code === undefined) this.fail('trailing escape');
|
|
147
|
+
this.index++;
|
|
148
|
+
if ('dDsSwW'.includes(code)) {
|
|
149
|
+
const bytes = shorthand(code.toLowerCase());
|
|
150
|
+
if (code === code.toUpperCase()) for (let word = 0; word < bytes.length; ++word) bytes[word] = ~bytes[word];
|
|
151
|
+
return { expression: byteSet(bytes), bytes };
|
|
152
|
+
}
|
|
153
|
+
if (code === 'b' && !inClass) this.fail('word boundaries are unsupported');
|
|
154
|
+
let value: number;
|
|
155
|
+
if (code === 'x') value = this.hex(2);
|
|
156
|
+
else if (code === 'u') value = this.hex(4);
|
|
157
|
+
else
|
|
158
|
+
value =
|
|
159
|
+
({ n: 10, r: 13, t: 9, f: 12, v: 11, b: 8 } as Record<string, number>)[code] ?? code.codePointAt(0)!;
|
|
160
|
+
const text = String.fromCodePoint(value);
|
|
161
|
+
const encoded = encoder.encode(text);
|
|
162
|
+
const bytes = encoded.length === 1 ? singleton(encoded[0]) : undefined;
|
|
163
|
+
return { expression: literal(text), bytes, single: encoded.length === 1 ? encoded[0] : undefined };
|
|
164
|
+
}
|
|
165
|
+
|
|
166
|
+
private bounds(): [number, number | null] {
|
|
167
|
+
this.index++;
|
|
168
|
+
const minimum = this.decimal();
|
|
169
|
+
let maximum: number | null = minimum;
|
|
170
|
+
if (this.peek() === ',') {
|
|
171
|
+
this.index++;
|
|
172
|
+
maximum = this.peek() === '}' ? null : this.decimal();
|
|
173
|
+
}
|
|
174
|
+
if (this.peek() !== '}') this.fail('unterminated repetition');
|
|
175
|
+
this.index++;
|
|
176
|
+
if (minimum > 1000 || (maximum !== null && (maximum < minimum || maximum > 1000))) {
|
|
177
|
+
this.fail('invalid or excessive repetition');
|
|
178
|
+
}
|
|
179
|
+
return [minimum, maximum];
|
|
180
|
+
}
|
|
181
|
+
|
|
182
|
+
private decimal(): number {
|
|
183
|
+
const start = this.index;
|
|
184
|
+
while (/\d/.test(this.peek() ?? '')) this.index++;
|
|
185
|
+
if (start === this.index) this.fail('expected a repetition count');
|
|
186
|
+
return Number(this.source.slice(start, this.index));
|
|
187
|
+
}
|
|
188
|
+
|
|
189
|
+
private hex(length: number): number {
|
|
190
|
+
const value = this.source.slice(this.index, this.index + length);
|
|
191
|
+
if (!new RegExp(`^[\\da-f]{${length}}$`, 'i').test(value)) this.fail('invalid hexadecimal escape');
|
|
192
|
+
this.index += length;
|
|
193
|
+
return Number.parseInt(value, 16);
|
|
194
|
+
}
|
|
195
|
+
|
|
196
|
+
private peek(): string | undefined {
|
|
197
|
+
return this.source[this.index];
|
|
198
|
+
}
|
|
199
|
+
|
|
200
|
+
private fail(message: string): never {
|
|
201
|
+
throw new SyntaxError(`Invalid regex at index ${this.index}: ${message}.`);
|
|
202
|
+
}
|
|
203
|
+
}
|
|
204
|
+
|
|
205
|
+
const OP_SET = 0;
|
|
206
|
+
const OP_SPLIT = 1;
|
|
207
|
+
const OP_JUMP = 2;
|
|
208
|
+
const OP_MATCH = 3;
|
|
209
|
+
const UNKNOWN = -2;
|
|
210
|
+
|
|
211
|
+
type Fragment = { start: number; outs: number[] };
|
|
212
|
+
|
|
213
|
+
class NfaBuilder {
|
|
214
|
+
readonly ops: number[] = [];
|
|
215
|
+
readonly out1: number[] = [];
|
|
216
|
+
readonly out2: number[] = [];
|
|
217
|
+
readonly sets: Uint32Array[] = [];
|
|
218
|
+
|
|
219
|
+
compile(expression: Expression): Fragment {
|
|
220
|
+
switch (expression.kind) {
|
|
221
|
+
case 'empty': {
|
|
222
|
+
const state = this.emit(OP_JUMP);
|
|
223
|
+
return { start: state, outs: [state << 1] };
|
|
224
|
+
}
|
|
225
|
+
case 'set': {
|
|
226
|
+
const state = this.emit(OP_SET, -1, -1, expression.bytes);
|
|
227
|
+
return { start: state, outs: [state << 1] };
|
|
228
|
+
}
|
|
229
|
+
case 'sequence': {
|
|
230
|
+
let result = this.compile(expression.parts[0]);
|
|
231
|
+
for (let index = 1; index < expression.parts.length; ++index) {
|
|
232
|
+
const next = this.compile(expression.parts[index]);
|
|
233
|
+
this.patch(result.outs, next.start);
|
|
234
|
+
result = { start: result.start, outs: next.outs };
|
|
235
|
+
}
|
|
236
|
+
return result;
|
|
237
|
+
}
|
|
238
|
+
case 'choice': {
|
|
239
|
+
let result = this.compile(expression.choices[0]);
|
|
240
|
+
for (let index = 1; index < expression.choices.length; ++index) {
|
|
241
|
+
const right = this.compile(expression.choices[index]);
|
|
242
|
+
result = {
|
|
243
|
+
start: this.emit(OP_SPLIT, result.start, right.start),
|
|
244
|
+
outs: [...result.outs, ...right.outs],
|
|
245
|
+
};
|
|
246
|
+
}
|
|
247
|
+
return result;
|
|
248
|
+
}
|
|
249
|
+
case 'star': {
|
|
250
|
+
const child = this.compile(expression.child);
|
|
251
|
+
const split = this.emit(OP_SPLIT, child.start);
|
|
252
|
+
this.patch(child.outs, split);
|
|
253
|
+
return { start: split, outs: [(split << 1) | 1] };
|
|
254
|
+
}
|
|
255
|
+
}
|
|
256
|
+
}
|
|
257
|
+
|
|
258
|
+
emit(op: number, first = -1, second = -1, set?: Uint32Array): number {
|
|
259
|
+
const state = this.ops.length;
|
|
260
|
+
this.ops.push(op);
|
|
261
|
+
this.out1.push(first);
|
|
262
|
+
this.out2.push(second);
|
|
263
|
+
this.sets.push(set ?? new Uint32Array(0));
|
|
264
|
+
return state;
|
|
265
|
+
}
|
|
266
|
+
|
|
267
|
+
patch(outs: number[], target: number): void {
|
|
268
|
+
for (const output of outs) {
|
|
269
|
+
if (output & 1) this.out2[output >>> 1] = target;
|
|
270
|
+
else this.out1[output >>> 1] = target;
|
|
271
|
+
}
|
|
272
|
+
}
|
|
273
|
+
}
|
|
274
|
+
|
|
275
|
+
class RegexMachine {
|
|
276
|
+
readonly initial: number;
|
|
277
|
+
private readonly builder = new NfaBuilder();
|
|
278
|
+
private readonly states: number[][] = [];
|
|
279
|
+
private readonly acceptingStates: boolean[] = [];
|
|
280
|
+
private readonly transitionTables: Int32Array[] = [];
|
|
281
|
+
private readonly stateIds = new Map<string, number>();
|
|
282
|
+
private readonly stateKeys: string[] = [];
|
|
283
|
+
|
|
284
|
+
constructor(expression: Expression) {
|
|
285
|
+
const fragment = this.builder.compile(expression);
|
|
286
|
+
const match = this.builder.emit(OP_MATCH);
|
|
287
|
+
this.builder.patch(fragment.outs, match);
|
|
288
|
+
this.initial = this.intern(this.closure([fragment.start]));
|
|
289
|
+
}
|
|
290
|
+
|
|
291
|
+
transition(state: number, byte: number): number {
|
|
292
|
+
if (state < 0) return -1;
|
|
293
|
+
const table = this.transitionTables[state];
|
|
294
|
+
const cached = table[byte];
|
|
295
|
+
if (cached !== UNKNOWN) return cached;
|
|
296
|
+
const seeds: number[] = [];
|
|
297
|
+
for (const pc of this.states[state]) {
|
|
298
|
+
if (this.builder.ops[pc] !== OP_SET) continue;
|
|
299
|
+
const set = this.builder.sets[pc];
|
|
300
|
+
if (set[byte >>> 5] & (1 << (byte & 31))) seeds.push(this.builder.out1[pc]);
|
|
301
|
+
}
|
|
302
|
+
const next = seeds.length === 0 ? -1 : this.intern(this.closure(seeds));
|
|
303
|
+
table[byte] = next;
|
|
304
|
+
return next;
|
|
305
|
+
}
|
|
306
|
+
|
|
307
|
+
accepting(state: number): boolean {
|
|
308
|
+
return state >= 0 && this.acceptingStates[state];
|
|
309
|
+
}
|
|
310
|
+
|
|
311
|
+
// The interned NFA state set is intrinsic to the regex (unlike the interned
|
|
312
|
+
// ids, which depend on discovery order), so it is a stable mask-cache key
|
|
313
|
+
// across constraint instances compiled from the same source.
|
|
314
|
+
stateKey(state: number): string | undefined {
|
|
315
|
+
return state >= 0 ? this.stateKeys[state] : undefined;
|
|
316
|
+
}
|
|
317
|
+
|
|
318
|
+
private closure(seeds: number[]): number[] {
|
|
319
|
+
const result: number[] = [];
|
|
320
|
+
const stack = [...seeds];
|
|
321
|
+
const seen = new Set<number>();
|
|
322
|
+
while (stack.length > 0) {
|
|
323
|
+
const state = stack.pop()!;
|
|
324
|
+
if (state < 0 || seen.has(state)) continue;
|
|
325
|
+
seen.add(state);
|
|
326
|
+
const op = this.builder.ops[state];
|
|
327
|
+
if (op === OP_SPLIT) {
|
|
328
|
+
stack.push(this.builder.out1[state], this.builder.out2[state]);
|
|
329
|
+
} else if (op === OP_JUMP) {
|
|
330
|
+
stack.push(this.builder.out1[state]);
|
|
331
|
+
} else {
|
|
332
|
+
result.push(state);
|
|
333
|
+
}
|
|
334
|
+
}
|
|
335
|
+
result.sort((left, right) => left - right);
|
|
336
|
+
return result;
|
|
337
|
+
}
|
|
338
|
+
|
|
339
|
+
private intern(active: number[]): number {
|
|
340
|
+
const key = active.join(',');
|
|
341
|
+
const existing = this.stateIds.get(key);
|
|
342
|
+
if (existing !== undefined) return existing;
|
|
343
|
+
if (this.states.length >= 4096) throw new Error('Regex produced too many runtime states.');
|
|
344
|
+
const id = this.states.length;
|
|
345
|
+
const transitions = new Int32Array(256);
|
|
346
|
+
transitions.fill(UNKNOWN);
|
|
347
|
+
this.states.push(active);
|
|
348
|
+
this.acceptingStates.push(active.some((state) => this.builder.ops[state] === OP_MATCH));
|
|
349
|
+
this.transitionTables.push(transitions);
|
|
350
|
+
this.stateIds.set(key, id);
|
|
351
|
+
this.stateKeys.push(key);
|
|
352
|
+
return id;
|
|
353
|
+
}
|
|
354
|
+
}
|
|
355
|
+
|
|
356
|
+
function choice(items: Expression[]): Expression {
|
|
357
|
+
const flattened = items.flatMap((item) => (item.kind === 'choice' ? item.choices : [item]));
|
|
358
|
+
if (flattened.length === 1) return flattened[0];
|
|
359
|
+
return { kind: 'choice', choices: flattened };
|
|
360
|
+
}
|
|
361
|
+
|
|
362
|
+
function sequence(items: Expression[]): Expression {
|
|
363
|
+
const flattened = items
|
|
364
|
+
.flatMap((item) => (item.kind === 'sequence' ? item.parts : [item]))
|
|
365
|
+
.filter((item) => item !== EMPTY);
|
|
366
|
+
if (flattened.length === 0) return EMPTY;
|
|
367
|
+
if (flattened.length === 1) return flattened[0];
|
|
368
|
+
return { kind: 'sequence', parts: flattened };
|
|
369
|
+
}
|
|
370
|
+
|
|
371
|
+
function star(child: Expression): Expression {
|
|
372
|
+
if (child === EMPTY) return EMPTY;
|
|
373
|
+
if (child.kind === 'star') return child;
|
|
374
|
+
return { kind: 'star', child };
|
|
375
|
+
}
|
|
376
|
+
|
|
377
|
+
function literal(value: string): Expression {
|
|
378
|
+
return sequence([...encoder.encode(value)].map((byte) => byteSet(singleton(byte))));
|
|
379
|
+
}
|
|
380
|
+
|
|
381
|
+
function byteSet(bytes: Uint32Array): Expression {
|
|
382
|
+
return { kind: 'set', bytes };
|
|
383
|
+
}
|
|
384
|
+
|
|
385
|
+
function shorthand(code: string): Uint32Array {
|
|
386
|
+
const bytes = new Uint32Array(8);
|
|
387
|
+
if (code === 'd' || code === 'w') addRange(bytes, 48, 57);
|
|
388
|
+
if (code === 'w') {
|
|
389
|
+
addRange(bytes, 65, 90);
|
|
390
|
+
addRange(bytes, 97, 122);
|
|
391
|
+
add(bytes, 95);
|
|
392
|
+
}
|
|
393
|
+
if (code === 's') for (const byte of [9, 10, 11, 12, 13, 32]) add(bytes, byte);
|
|
394
|
+
return bytes;
|
|
395
|
+
}
|
|
396
|
+
|
|
397
|
+
function singleton(byte: number): Uint32Array {
|
|
398
|
+
const bytes = new Uint32Array(8);
|
|
399
|
+
add(bytes, byte);
|
|
400
|
+
return bytes;
|
|
401
|
+
}
|
|
402
|
+
|
|
403
|
+
function range(first: number, last: number): Uint32Array {
|
|
404
|
+
const bytes = new Uint32Array(8);
|
|
405
|
+
addRange(bytes, first, last);
|
|
406
|
+
return bytes;
|
|
407
|
+
}
|
|
408
|
+
|
|
409
|
+
function addRange(bytes: Uint32Array, first: number, last: number): void {
|
|
410
|
+
for (let byte = first; byte <= last; ++byte) add(bytes, byte);
|
|
411
|
+
}
|
|
412
|
+
|
|
413
|
+
function add(bytes: Uint32Array, byte: number): void {
|
|
414
|
+
bytes[byte >>> 5] |= 1 << (byte & 31);
|
|
415
|
+
}
|
|
416
|
+
|
|
417
|
+
function union(target: Uint32Array, source: Uint32Array): void {
|
|
418
|
+
for (let word = 0; word < target.length; ++word) target[word] |= source[word];
|
|
419
|
+
}
|
|
@@ -0,0 +1,205 @@
|
|
|
1
|
+
import type { TokenizerSource } from './types';
|
|
2
|
+
|
|
3
|
+
type RecordLike = Record<string, unknown>;
|
|
4
|
+
|
|
5
|
+
export type TokenizerData = {
|
|
6
|
+
tokens: Uint8Array[];
|
|
7
|
+
eosTokenId: number;
|
|
8
|
+
specialTokenIds: Set<number>;
|
|
9
|
+
};
|
|
10
|
+
|
|
11
|
+
const encoder = new TextEncoder();
|
|
12
|
+
let byteLevelMap: Map<string, number> | undefined;
|
|
13
|
+
|
|
14
|
+
export function extractTokenizer(tokenizer: TokenizerSource): TokenizerData {
|
|
15
|
+
const source = asRecord(tokenizer, 'tokenizer');
|
|
16
|
+
if (Array.isArray(source.tokens)) {
|
|
17
|
+
return normalizeDirectTokenizer(source);
|
|
18
|
+
}
|
|
19
|
+
|
|
20
|
+
const tokenizerJson = getTokenizerJson(source);
|
|
21
|
+
const vocabulary = getVocabulary(source, tokenizerJson);
|
|
22
|
+
if (vocabulary === undefined) {
|
|
23
|
+
throw new TypeError('Could not extract the tokenizer vocabulary.');
|
|
24
|
+
}
|
|
25
|
+
|
|
26
|
+
const vocabularyTokens = Object.keys(vocabulary);
|
|
27
|
+
let size = 0;
|
|
28
|
+
for (const token of vocabularyTokens) {
|
|
29
|
+
const id = Number(vocabulary[token]);
|
|
30
|
+
if (!Number.isInteger(id) || id < 0) throw new TypeError(`Tokenizer has an invalid token ID for ${token}.`);
|
|
31
|
+
if (id + 1 > size) size = id + 1;
|
|
32
|
+
}
|
|
33
|
+
const tokenBytes = tokenBytesConverter(source, tokenizerJson);
|
|
34
|
+
const tokens = new Array<Uint8Array | undefined>(size);
|
|
35
|
+
for (const token of vocabularyTokens) {
|
|
36
|
+
const id = Number(vocabulary[token]);
|
|
37
|
+
tokens[id] = tokenBytes(token, id);
|
|
38
|
+
}
|
|
39
|
+
|
|
40
|
+
const addedTokens = field(tokenizerJson, 'added_tokens');
|
|
41
|
+
if (Array.isArray(addedTokens)) {
|
|
42
|
+
for (const added of addedTokens) {
|
|
43
|
+
if (!isRecord(added) || !Number.isInteger(added.id)) continue;
|
|
44
|
+
while (tokens.length <= Number(added.id)) tokens.push(undefined);
|
|
45
|
+
const id = Number(added.id);
|
|
46
|
+
tokens[id] = tokenBytes(typeof added.content === 'string' ? added.content : '', id);
|
|
47
|
+
}
|
|
48
|
+
}
|
|
49
|
+
for (let id = 0; id < tokens.length; ++id) {
|
|
50
|
+
if (tokens[id] === undefined) throw new Error(`Tokenizer vocabulary is missing token ID ${id}.`);
|
|
51
|
+
}
|
|
52
|
+
|
|
53
|
+
const eosTokenId = tokenId(source, tokenizerJson, ['eos_token_id', 'eosTokenId', 'eos_token', 'eosToken']);
|
|
54
|
+
if (eosTokenId === undefined) throw new TypeError('Tokenizer does not expose an EOS token ID.');
|
|
55
|
+
const specialTokenIds = new Set<number>([eosTokenId]);
|
|
56
|
+
for (const value of [
|
|
57
|
+
source.special_token_ids,
|
|
58
|
+
source.specialTokenIds,
|
|
59
|
+
source.all_special_ids,
|
|
60
|
+
source.allSpecialIds,
|
|
61
|
+
]) {
|
|
62
|
+
if (Array.isArray(value)) for (const id of value) if (Number.isInteger(id)) specialTokenIds.add(Number(id));
|
|
63
|
+
}
|
|
64
|
+
if (Array.isArray(addedTokens)) {
|
|
65
|
+
for (const added of addedTokens) {
|
|
66
|
+
if (isRecord(added) && added.special === true && Number.isInteger(added.id))
|
|
67
|
+
specialTokenIds.add(Number(added.id));
|
|
68
|
+
}
|
|
69
|
+
}
|
|
70
|
+
return { tokens: tokens as Uint8Array[], eosTokenId, specialTokenIds };
|
|
71
|
+
}
|
|
72
|
+
|
|
73
|
+
function normalizeDirectTokenizer(source: RecordLike): TokenizerData {
|
|
74
|
+
const configuredTokens = source.tokens as unknown[];
|
|
75
|
+
const tokens = configuredTokens.map((token, id) => {
|
|
76
|
+
if (!(token instanceof Uint8Array) && !Array.isArray(token)) {
|
|
77
|
+
throw new TypeError(`Tokenizer token ${id} must be a byte array.`);
|
|
78
|
+
}
|
|
79
|
+
const values = Array.from(token as ArrayLike<number>);
|
|
80
|
+
if (values.some((byte) => !Number.isInteger(byte) || byte < 0 || byte > 255))
|
|
81
|
+
throw new TypeError(`Tokenizer token ${id} is invalid.`);
|
|
82
|
+
return Uint8Array.from(values);
|
|
83
|
+
});
|
|
84
|
+
const eosTokenId = Number(source.eosTokenId ?? source.eos_token_id);
|
|
85
|
+
if (!Number.isInteger(eosTokenId) || eosTokenId < 0 || eosTokenId >= tokens.length) {
|
|
86
|
+
throw new TypeError('A valid eos_token_id is required with tokenizer tokens.');
|
|
87
|
+
}
|
|
88
|
+
const configured = source.specialTokenIds ?? source.special_token_ids;
|
|
89
|
+
const specialTokenIds = new Set<number>([eosTokenId]);
|
|
90
|
+
if (Array.isArray(configured)) for (const id of configured) specialTokenIds.add(Number(id));
|
|
91
|
+
return { tokens, eosTokenId, specialTokenIds };
|
|
92
|
+
}
|
|
93
|
+
|
|
94
|
+
function getTokenizerJson(source: RecordLike): unknown {
|
|
95
|
+
const value = source._tokenizerJSON ?? source.tokenizerJSON ?? source.tokenizer_json;
|
|
96
|
+
return typeof value === 'string' ? JSON.parse(value) : value;
|
|
97
|
+
}
|
|
98
|
+
|
|
99
|
+
function getVocabulary(source: RecordLike, tokenizerJson: unknown): RecordLike | undefined {
|
|
100
|
+
const modelVocabulary = field(field(tokenizerJson, 'model'), 'vocab');
|
|
101
|
+
if (isRecord(modelVocabulary)) return modelVocabulary;
|
|
102
|
+
for (const name of ['get_vocab', 'getVocab']) {
|
|
103
|
+
const method = source[name];
|
|
104
|
+
if (typeof method !== 'function') continue;
|
|
105
|
+
const value = method.call(source, true);
|
|
106
|
+
if (value instanceof Map) return Object.fromEntries(value);
|
|
107
|
+
if (isRecord(value)) return value;
|
|
108
|
+
}
|
|
109
|
+
return isRecord(source.vocab) ? source.vocab : undefined;
|
|
110
|
+
}
|
|
111
|
+
|
|
112
|
+
function tokenBytesConverter(source: RecordLike, tokenizerJson: unknown): (token: string, id: number) => Uint8Array {
|
|
113
|
+
const decoder = field(tokenizerJson, 'decoder');
|
|
114
|
+
if (field(decoder, 'type') === 'ByteLevel') return byteLevelTokenBytes;
|
|
115
|
+
|
|
116
|
+
const decode = source.decode;
|
|
117
|
+
if (typeof decode === 'function') {
|
|
118
|
+
const byteFallback = hasComponent(decoder, 'ByteFallback');
|
|
119
|
+
return (token, id) => {
|
|
120
|
+
const fallback = byteFallback ? /^<0x([\da-f]{2})>$/i.exec(token) : null;
|
|
121
|
+
if (fallback) return Uint8Array.of(Number.parseInt(fallback[1], 16));
|
|
122
|
+
const decoded = decode.call(source, [id], {
|
|
123
|
+
skip_special_tokens: false,
|
|
124
|
+
clean_up_tokenization_spaces: false,
|
|
125
|
+
});
|
|
126
|
+
if (typeof decoded !== 'string') throw new TypeError(`Tokenizer.decode([${id}]) must return a string.`);
|
|
127
|
+
return encoder.encode(decoded);
|
|
128
|
+
};
|
|
129
|
+
}
|
|
130
|
+
|
|
131
|
+
if (decoder !== undefined && decoder !== null && field(decoder, 'type') !== 'ByteLevel') {
|
|
132
|
+
throw new TypeError('Tokenizer decoder semantics require a tokenizer with a decode() method.');
|
|
133
|
+
}
|
|
134
|
+
if (hasComponent(decoder, 'ByteLevel') || hasComponent(field(tokenizerJson, 'pre_tokenizer'), 'ByteLevel')) {
|
|
135
|
+
return byteLevelTokenBytes;
|
|
136
|
+
}
|
|
137
|
+
const modelType = field(field(tokenizerJson, 'model'), 'type');
|
|
138
|
+
const sentencePiece = modelType === 'Unigram' || modelType === 'SentencePiece';
|
|
139
|
+
const prefix = field(field(tokenizerJson, 'model'), 'continuing_subword_prefix');
|
|
140
|
+
return (token) => {
|
|
141
|
+
const fallback = /^<0x([\da-f]{2})>$/i.exec(token);
|
|
142
|
+
if (fallback) return Uint8Array.of(Number.parseInt(fallback[1], 16));
|
|
143
|
+
if (sentencePiece || token.includes('▁')) return encoder.encode(token.replaceAll('▁', ' '));
|
|
144
|
+
return encoder.encode(
|
|
145
|
+
typeof prefix === 'string' && token.startsWith(prefix) ? token.slice(prefix.length) : token,
|
|
146
|
+
);
|
|
147
|
+
};
|
|
148
|
+
}
|
|
149
|
+
|
|
150
|
+
function byteLevelTokenBytes(token: string): Uint8Array {
|
|
151
|
+
const fallback = /^<0x([\da-f]{2})>$/i.exec(token);
|
|
152
|
+
if (fallback) return Uint8Array.of(Number.parseInt(fallback[1], 16));
|
|
153
|
+
const map = getByteLevelMap();
|
|
154
|
+
const bytes: number[] = [];
|
|
155
|
+
for (const character of token) {
|
|
156
|
+
const byte = map.get(character);
|
|
157
|
+
if (byte === undefined) bytes.push(...encoder.encode(character));
|
|
158
|
+
else bytes.push(byte);
|
|
159
|
+
}
|
|
160
|
+
return Uint8Array.from(bytes);
|
|
161
|
+
}
|
|
162
|
+
|
|
163
|
+
function getByteLevelMap(): Map<string, number> {
|
|
164
|
+
if (byteLevelMap !== undefined) return byteLevelMap;
|
|
165
|
+
const visible = new Set<number>();
|
|
166
|
+
for (let code = 33; code <= 126; ++code) visible.add(code);
|
|
167
|
+
for (let code = 161; code <= 172; ++code) visible.add(code);
|
|
168
|
+
for (let code = 174; code <= 255; ++code) visible.add(code);
|
|
169
|
+
let extra = 0;
|
|
170
|
+
byteLevelMap = new Map();
|
|
171
|
+
for (let byte = 0; byte < 256; ++byte) {
|
|
172
|
+
byteLevelMap.set(String.fromCharCode(visible.has(byte) ? byte : 256 + extra++), byte);
|
|
173
|
+
}
|
|
174
|
+
return byteLevelMap;
|
|
175
|
+
}
|
|
176
|
+
|
|
177
|
+
function tokenId(source: RecordLike, tokenizerJson: unknown, keys: string[]): number | undefined {
|
|
178
|
+
const vocabulary = getVocabulary(source, tokenizerJson);
|
|
179
|
+
for (const key of keys) {
|
|
180
|
+
const value = source[key] ?? field(tokenizerJson, key);
|
|
181
|
+
if (Number.isInteger(value)) return Number(value);
|
|
182
|
+
if (typeof value === 'string' && Number.isInteger(vocabulary?.[value])) return Number(vocabulary![value]);
|
|
183
|
+
}
|
|
184
|
+
return undefined;
|
|
185
|
+
}
|
|
186
|
+
|
|
187
|
+
function hasComponent(value: unknown, type: string): boolean {
|
|
188
|
+
if (Array.isArray(value)) return value.some((item) => hasComponent(item, type));
|
|
189
|
+
return isRecord(value) && (value.type === type || Object.values(value).some((item) => hasComponent(item, type)));
|
|
190
|
+
}
|
|
191
|
+
|
|
192
|
+
function field(value: unknown, key: string): unknown {
|
|
193
|
+
return isRecord(value) ? value[key] : undefined;
|
|
194
|
+
}
|
|
195
|
+
|
|
196
|
+
function asRecord(value: unknown, name: string): RecordLike {
|
|
197
|
+
if ((typeof value !== 'object' && typeof value !== 'function') || value === null) {
|
|
198
|
+
throw new TypeError(`${name} must be an object.`);
|
|
199
|
+
}
|
|
200
|
+
return value as RecordLike;
|
|
201
|
+
}
|
|
202
|
+
|
|
203
|
+
function isRecord(value: unknown): value is RecordLike {
|
|
204
|
+
return value !== null && typeof value === 'object' && !Array.isArray(value);
|
|
205
|
+
}
|
|
@@ -0,0 +1,12 @@
|
|
|
1
|
+
export type JSONValue = null | boolean | number | string | JSONValue[] | { [key: string]: JSONValue };
|
|
2
|
+
export type JSONSchema = boolean | { [key: string]: JSONValue };
|
|
3
|
+
export type TokenizerSource = object | ((...args: unknown[]) => unknown);
|
|
4
|
+
|
|
5
|
+
export interface ConstraintState<State> {
|
|
6
|
+
readonly initial: State;
|
|
7
|
+
transition(state: State, byte: number): State;
|
|
8
|
+
viable(state: State): boolean;
|
|
9
|
+
accepting(state: State): boolean;
|
|
10
|
+
stringCapacity?(state: State): number | undefined;
|
|
11
|
+
maskKey?(state: State): string | undefined;
|
|
12
|
+
}
|
package/src/index.ts
ADDED
|
@@ -0,0 +1,30 @@
|
|
|
1
|
+
import { type Tensor } from '@huggingface/transformers';
|
|
2
|
+
|
|
3
|
+
type LogitsData = Float32Array | Float64Array | number[];
|
|
4
|
+
|
|
5
|
+
export function applyMask(logits: Tensor, mask: Uint32Array, vocabSize: number): void {
|
|
6
|
+
const data = logits.data as LogitsData;
|
|
7
|
+
const stride = logits.dims.at(-1)!;
|
|
8
|
+
if (vocabSize > stride) {
|
|
9
|
+
throw new Error(`Constraint vocabulary size ${vocabSize} exceeds logits vocabulary size ${stride}.`);
|
|
10
|
+
}
|
|
11
|
+
for (let offset = 0; offset < data.length; offset += stride) {
|
|
12
|
+
const fullWords = vocabSize >>> 5;
|
|
13
|
+
for (let word = 0; word < fullWords; ++word) {
|
|
14
|
+
const bits = mask[word] | 0;
|
|
15
|
+
if (bits === -1) continue;
|
|
16
|
+
const start = offset + (word << 5);
|
|
17
|
+
if (bits === 0) {
|
|
18
|
+
data.fill(-Infinity, start, start + 32);
|
|
19
|
+
} else {
|
|
20
|
+
for (let bit = 0; bit < 32; ++bit) {
|
|
21
|
+
if (!(bits & (1 << bit))) data[start + bit] = -Infinity;
|
|
22
|
+
}
|
|
23
|
+
}
|
|
24
|
+
}
|
|
25
|
+
for (let tokenId = fullWords << 5; tokenId < vocabSize; ++tokenId) {
|
|
26
|
+
if (!(mask[tokenId >>> 5] & (1 << (tokenId & 31)))) data[offset + tokenId] = -Infinity;
|
|
27
|
+
}
|
|
28
|
+
data.fill(-Infinity, offset + vocabSize, offset + stride);
|
|
29
|
+
}
|
|
30
|
+
}
|