@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.
@@ -0,0 +1,103 @@
1
+ import { LogitsProcessor, LogitsProcessorList, type Tensor } from '@huggingface/transformers';
2
+
3
+ import {
4
+ createTokenConstraint,
5
+ prepareTokenizer,
6
+ type JSONSchema,
7
+ type TokenConstraint,
8
+ type TokenizerSource,
9
+ } from './engine';
10
+ import { applyMask } from './utils/mask';
11
+
12
+ export type ResponseFormat =
13
+ | { type: 'json_object' }
14
+ | { type: 'json_schema'; json_schema: JSONSchema }
15
+ | { type: 'regex'; regex: string };
16
+
17
+ type GenerationState = {
18
+ constraint: TokenConstraint;
19
+ processedInputLength?: number;
20
+ mask?: Uint32Array;
21
+ };
22
+
23
+ const WHITESPACE_REPETITION_PENALTY = 1.2;
24
+ const MAX_CONSECUTIVE_WHITESPACE_TOKENS = 4;
25
+
26
+ export class StructuredOutputProcessor extends LogitsProcessorList {
27
+ /**
28
+ * Precomputes the tokenizer-derived data structures used by every
29
+ * constraint. The first processor per tokenizer otherwise pays this cost
30
+ * (hundreds of milliseconds for large vocabularies) in its constructor;
31
+ * call this once after loading the model to pay it early instead.
32
+ */
33
+ static warmup(tokenizer: TokenizerSource): void {
34
+ prepareTokenizer(tokenizer);
35
+ }
36
+
37
+ constructor(tokenizer: TokenizerSource, responseFormat: ResponseFormat) {
38
+ super();
39
+ const state: GenerationState = {
40
+ constraint: createTokenConstraint(tokenizer, responseFormat),
41
+ };
42
+ this.push(new ConstraintLogitsProcessor(state));
43
+ }
44
+ }
45
+
46
+ class ConstraintLogitsProcessor extends LogitsProcessor {
47
+ constructor(private readonly state: GenerationState) {
48
+ super();
49
+ }
50
+
51
+ _call(inputIds: bigint[][], logits: Tensor) {
52
+ assertSingleSequence(inputIds.length);
53
+ const input = inputIds[0];
54
+ const start = this.state.processedInputLength ?? input.length;
55
+ for (let i = start; i < input.length; ++i) {
56
+ if (this.state.constraint.commit(Number(input[i]))) {
57
+ throw new Error(
58
+ 'StructuredOutputProcessor observed the tokenizer EOS token after generation continued. Ensure the model generation config uses the same eos_token_id as the tokenizer.',
59
+ );
60
+ }
61
+ }
62
+ this.state.processedInputLength = input.length;
63
+ const logitsVocabSize = logits.dims.at(-1);
64
+ if (logitsVocabSize === undefined || !Number.isInteger(logitsVocabSize) || logitsVocabSize <= 0) {
65
+ throw new Error('StructuredOutputProcessor requires logits with a vocabulary dimension.');
66
+ }
67
+ const words = Math.ceil(logitsVocabSize / 32);
68
+ if (this.state.mask?.length !== words) this.state.mask = new Uint32Array(words);
69
+ if (!this.state.constraint.fillMask(this.state.mask)) {
70
+ throw new Error('The constraint reached a dead end before producing a valid output.');
71
+ }
72
+ applyMask(logits, this.state.mask, this.state.constraint.vocabSize);
73
+ const repeatedWhitespace = this.state.constraint.repeatedWhitespace();
74
+ if (repeatedWhitespace !== undefined) {
75
+ discourageRepeatedWhitespace(logits, repeatedWhitespace.tokenIds, repeatedWhitespace.count);
76
+ }
77
+ return logits;
78
+ }
79
+ }
80
+
81
+ function assertSingleSequence(batchSize: number): void {
82
+ if (batchSize !== 1) {
83
+ throw new Error(`StructuredOutputProcessor currently supports batch size 1; received ${batchSize}.`);
84
+ }
85
+ }
86
+
87
+ function discourageRepeatedWhitespace(logits: Tensor, tokenIds: readonly number[], count: number): void {
88
+ const data = logits.data as Float32Array | Float64Array | number[];
89
+ const stride = logits.dims.at(-1)!;
90
+ const penalty = WHITESPACE_REPETITION_PENALTY ** count;
91
+ for (let offset = 0; offset < data.length; offset += stride) {
92
+ for (const tokenId of tokenIds) {
93
+ const index = offset + tokenId;
94
+ if (count >= MAX_CONSECUTIVE_WHITESPACE_TOKENS) {
95
+ data[index] = -Infinity;
96
+ } else if (data[index] < 0) {
97
+ data[index] *= penalty;
98
+ } else {
99
+ data[index] /= penalty;
100
+ }
101
+ }
102
+ }
103
+ }
@@ -0,0 +1,331 @@
1
+ import { compileJsonSchema } from './json';
2
+ import { compileRegex } from './regex';
3
+ import { extractTokenizer, type TokenizerData } from './tokenizer';
4
+ import type { ConstraintState, JSONSchema, TokenizerSource } from './types';
5
+
6
+ type TrieNode = { childBytes: number[]; childNodes: TrieNode[]; tokenIds: number[] };
7
+ type CachedTokenizer = {
8
+ data: TokenizerData;
9
+ trie: TrieNode;
10
+ whitespaceTokenIds: number[];
11
+ stringExceptionalTrie: TrieNode;
12
+ stringSafeMask: Uint32Array;
13
+ stringSafeCount: number;
14
+ stringSafeLengths: Uint32Array;
15
+ maxStringSafeLength: number;
16
+ maxTokenByteLength: number;
17
+ boundedStringMasks: Map<number, { mask: Uint32Array; count: number }>;
18
+ schemaMaskCaches: WeakMap<object, MaskCache>;
19
+ booleanSchemaMaskCaches: [MaskCache, MaskCache];
20
+ jsonObjectMaskCache: MaskCache;
21
+ regexMaskCaches: Map<string, MaskCache>;
22
+ };
23
+ type ResponseFormat =
24
+ | { type: 'json_object' }
25
+ | { type: 'json_schema'; json_schema: JSONSchema }
26
+ | { type: 'regex'; regex: string };
27
+
28
+ export type TokenConstraint = {
29
+ vocabSize: number;
30
+ fillMask(target: Uint32Array): boolean;
31
+ commit(tokenId: number): boolean;
32
+ repeatedWhitespace(): { tokenIds: readonly number[]; count: number } | undefined;
33
+ };
34
+
35
+ const tokenizerCache = new WeakMap<object, CachedTokenizer>();
36
+ const JSON_OBJECT_SCHEMA: JSONSchema = { type: 'object' };
37
+
38
+ /**
39
+ * Builds and caches the tokenizer-derived data structures (token tries, string
40
+ * masks) ahead of time. This is the expensive part of creating the first
41
+ * constraint for a tokenizer (hundreds of milliseconds for a 256k vocabulary),
42
+ * so calling this right after loading a model moves that cost off the first
43
+ * generation. Subsequent calls with the same tokenizer are free.
44
+ */
45
+ export function prepareTokenizer(tokenizerSource: TokenizerSource): void {
46
+ cachedTokenizer(tokenizerSource);
47
+ }
48
+
49
+ export function createTokenConstraint(
50
+ tokenizerSource: TokenizerSource,
51
+ responseFormat: ResponseFormat,
52
+ ): TokenConstraint {
53
+ const tokenizer = cachedTokenizer(tokenizerSource);
54
+ const machine = createMachine(responseFormat, tokenizer);
55
+ const maskCache = cacheFor(tokenizer, responseFormat);
56
+ let state = machine.initial;
57
+ // Post-transition states discovered during the trie walk, so commit() can
58
+ // reuse them. Entries are only valid when their stamp matches the current
59
+ // fillMask() generation; bumping the stamp invalidates all of them at once
60
+ // without refilling the vocabulary-sized array on every step.
61
+ const tokenStates: Array<unknown> = new Array(tokenizer.data.tokens.length);
62
+ const tokenStamps = new Int32Array(tokenizer.data.tokens.length);
63
+ let stamp = 0;
64
+ let consecutiveWhitespace = 0;
65
+ const tracksJsonWhitespace = responseFormat.type !== 'regex';
66
+
67
+ return {
68
+ vocabSize: tokenizer.data.tokens.length,
69
+ fillMask(target) {
70
+ const words = Math.ceil(tokenizer.data.tokens.length / 32);
71
+ if (target.length < words) throw new RangeError(`Mask target requires at least ${words} words.`);
72
+ target.fill(0);
73
+ stamp++;
74
+ const cacheKey = machine.maskKey?.(state);
75
+ const cachedMask = cacheKey === undefined ? undefined : maskCache?.get(cacheKey);
76
+ if (cachedMask !== undefined) {
77
+ target.set(cachedMask);
78
+ return true;
79
+ }
80
+ let allowed = 0;
81
+ if (machine.accepting(state)) {
82
+ setBit(target, tokenizer.data.eosTokenId);
83
+ allowed++;
84
+ }
85
+ const stringCapacity = machine.stringCapacity?.(state);
86
+ if (stringCapacity !== undefined) {
87
+ const safe = boundedStringMask(tokenizer, stringCapacity);
88
+ target.set(safe.mask);
89
+ allowed += safe.count;
90
+ }
91
+ const nodes: TrieNode[] = [stringCapacity === undefined ? tokenizer.trie : tokenizer.stringExceptionalTrie];
92
+ const states: unknown[] = [state];
93
+ while (nodes.length > 0) {
94
+ const node = nodes.pop()!;
95
+ const current = states.pop()!;
96
+ for (const tokenId of node.tokenIds) {
97
+ if (tokenizer.data.specialTokenIds.has(tokenId)) continue;
98
+ setBit(target, tokenId);
99
+ tokenStates[tokenId] = current;
100
+ tokenStamps[tokenId] = stamp;
101
+ allowed++;
102
+ }
103
+ for (let index = 0; index < node.childNodes.length; ++index) {
104
+ const next = machine.transition(current, node.childBytes[index]);
105
+ if (!machine.viable(next)) continue;
106
+ nodes.push(node.childNodes[index]);
107
+ states.push(next);
108
+ }
109
+ }
110
+ if (allowed > 0 && cacheKey !== undefined) {
111
+ maskCache?.set(cacheKey, target.subarray(0, words));
112
+ }
113
+ return allowed > 0;
114
+ },
115
+ commit(tokenId) {
116
+ if (!Number.isInteger(tokenId) || tokenId < 0 || tokenId >= tokenizer.data.tokens.length) {
117
+ throw new RangeError(`Token ${tokenId} is outside the tokenizer vocabulary.`);
118
+ }
119
+ if (tokenId === tokenizer.data.eosTokenId) {
120
+ if (!machine.accepting(state)) throw new Error(`Token ${tokenId} does not satisfy the constraint.`);
121
+ return true;
122
+ }
123
+ if (tokenizer.data.specialTokenIds.has(tokenId)) {
124
+ throw new Error(`Token ${tokenId} does not satisfy the constraint.`);
125
+ }
126
+ let next: unknown;
127
+ if (tokenStamps[tokenId] === stamp && stamp > 0) {
128
+ next = tokenStates[tokenId];
129
+ } else {
130
+ next = state;
131
+ for (const byte of tokenizer.data.tokens[tokenId]) next = machine.transition(next, byte);
132
+ }
133
+ stamp++;
134
+ if (!machine.viable(next)) throw new Error(`Token ${tokenId} does not satisfy the constraint.`);
135
+ consecutiveWhitespace =
136
+ tracksJsonWhitespace && next === state && isJsonWhitespace(tokenizer.data.tokens[tokenId])
137
+ ? consecutiveWhitespace + 1
138
+ : 0;
139
+ state = next;
140
+ return false;
141
+ },
142
+ repeatedWhitespace() {
143
+ if (consecutiveWhitespace === 0) return undefined;
144
+ return { tokenIds: tokenizer.whitespaceTokenIds, count: consecutiveWhitespace };
145
+ },
146
+ };
147
+ }
148
+
149
+ function createMachine(responseFormat: ResponseFormat, tokenizer: CachedTokenizer): ConstraintState<unknown> {
150
+ if (responseFormat?.type === 'regex') {
151
+ if (typeof responseFormat.regex !== 'string') throw new TypeError('response_format.regex must be a string.');
152
+ return compileRegex(responseFormat.regex) as ConstraintState<unknown>;
153
+ }
154
+ if (responseFormat?.type === 'json_schema') {
155
+ return compileJsonSchema(responseFormat.json_schema, tokenizer.maxTokenByteLength) as ConstraintState<unknown>;
156
+ }
157
+ if (responseFormat?.type === 'json_object') {
158
+ return compileJsonSchema(JSON_OBJECT_SCHEMA, tokenizer.maxTokenByteLength) as ConstraintState<unknown>;
159
+ }
160
+ throw new TypeError(`Unsupported response format: ${String((responseFormat as { type?: unknown })?.type)}.`);
161
+ }
162
+
163
+ function cachedTokenizer(source: TokenizerSource): CachedTokenizer {
164
+ let cached = tokenizerCache.get(source as object);
165
+ if (cached === undefined) {
166
+ const data = extractTokenizer(source);
167
+ const stringExceptionalTokenIds: number[] = [];
168
+ const whitespaceTokenIds: number[] = [];
169
+ const stringSafeMask = new Uint32Array(Math.ceil(data.tokens.length / 32));
170
+ const stringSafeLengths = new Uint32Array(data.tokens.length);
171
+ let stringSafeCount = 0;
172
+ let maxStringSafeLength = 0;
173
+ let maxTokenByteLength = 0;
174
+ for (let tokenId = 0; tokenId < data.tokens.length; ++tokenId) {
175
+ const special = data.specialTokenIds.has(tokenId);
176
+ if (!special && isJsonWhitespace(data.tokens[tokenId])) whitespaceTokenIds.push(tokenId);
177
+ if (!special && data.tokens[tokenId].length > maxTokenByteLength) {
178
+ maxTokenByteLength = data.tokens[tokenId].length;
179
+ }
180
+ const length = special ? undefined : safeStringTokenLength(data.tokens[tokenId]);
181
+ if (length !== undefined) {
182
+ stringSafeCount++;
183
+ stringSafeLengths[tokenId] = length;
184
+ if (length > maxStringSafeLength) maxStringSafeLength = length;
185
+ setBit(stringSafeMask, tokenId);
186
+ } else {
187
+ stringExceptionalTokenIds.push(tokenId);
188
+ }
189
+ }
190
+ cached = {
191
+ data,
192
+ trie: createTrie(data.tokens),
193
+ whitespaceTokenIds,
194
+ stringExceptionalTrie: createTrie(data.tokens, stringExceptionalTokenIds),
195
+ stringSafeMask,
196
+ stringSafeCount,
197
+ stringSafeLengths,
198
+ maxStringSafeLength,
199
+ maxTokenByteLength,
200
+ boundedStringMasks: new Map(),
201
+ schemaMaskCaches: new WeakMap(),
202
+ booleanSchemaMaskCaches: [new MaskCache(), new MaskCache()],
203
+ jsonObjectMaskCache: new MaskCache(),
204
+ regexMaskCaches: new Map(),
205
+ };
206
+ tokenizerCache.set(source as object, cached);
207
+ }
208
+ return cached;
209
+ }
210
+
211
+ function cacheFor(tokenizer: CachedTokenizer, responseFormat: ResponseFormat): MaskCache | undefined {
212
+ if (responseFormat.type === 'regex') {
213
+ let cache = tokenizer.regexMaskCaches.get(responseFormat.regex);
214
+ if (cache === undefined) {
215
+ cache = new MaskCache();
216
+ if (tokenizer.regexMaskCaches.size >= 16) {
217
+ tokenizer.regexMaskCaches.delete(tokenizer.regexMaskCaches.keys().next().value!);
218
+ }
219
+ tokenizer.regexMaskCaches.set(responseFormat.regex, cache);
220
+ }
221
+ return cache;
222
+ }
223
+ if (responseFormat.type === 'json_object') return tokenizer.jsonObjectMaskCache;
224
+ const schema = responseFormat.json_schema;
225
+ if (typeof schema === 'boolean') return tokenizer.booleanSchemaMaskCaches[schema ? 1 : 0];
226
+ let cache = tokenizer.schemaMaskCaches.get(schema);
227
+ if (cache === undefined) {
228
+ cache = new MaskCache();
229
+ tokenizer.schemaMaskCaches.set(schema, cache);
230
+ }
231
+ return cache;
232
+ }
233
+
234
+ class MaskCache {
235
+ private readonly masks = new Map<string, Uint32Array>();
236
+ private words = 0;
237
+
238
+ get(key: string): Uint32Array | undefined {
239
+ const mask = this.masks.get(key);
240
+ if (mask === undefined) return undefined;
241
+ this.masks.delete(key);
242
+ this.masks.set(key, mask);
243
+ return mask;
244
+ }
245
+
246
+ set(key: string, source: Uint32Array): void {
247
+ const mask = source.slice();
248
+ const previous = this.masks.get(key);
249
+ if (previous !== undefined) {
250
+ this.words -= previous.length;
251
+ this.masks.delete(key);
252
+ }
253
+ this.masks.set(key, mask);
254
+ this.words += mask.length;
255
+ while (this.masks.size > 256 || this.words > 1_048_576) {
256
+ const oldestKey = this.masks.keys().next().value!;
257
+ const oldest = this.masks.get(oldestKey)!;
258
+ this.masks.delete(oldestKey);
259
+ this.words -= oldest.length;
260
+ }
261
+ }
262
+ }
263
+
264
+ function createTrie(tokens: Uint8Array[], tokenIds?: number[]): TrieNode {
265
+ const root: TrieNode = { childBytes: [], childNodes: [], tokenIds: [] };
266
+ const size = tokenIds === undefined ? tokens.length : tokenIds.length;
267
+ for (let index = 0; index < size; ++index) {
268
+ const tokenId = tokenIds === undefined ? index : tokenIds[index];
269
+ const bytes = tokens[tokenId];
270
+ let node = root;
271
+ for (let position = 0; position < bytes.length; ++position) {
272
+ const byte = bytes[position];
273
+ const childIndex = node.childBytes.indexOf(byte);
274
+ if (childIndex === -1) {
275
+ const child: TrieNode = { childBytes: [], childNodes: [], tokenIds: [] };
276
+ node.childBytes.push(byte);
277
+ node.childNodes.push(child);
278
+ node = child;
279
+ } else {
280
+ node = node.childNodes[childIndex];
281
+ }
282
+ }
283
+ node.tokenIds.push(tokenId);
284
+ }
285
+ return root;
286
+ }
287
+
288
+ const safeStringDecoder = new TextDecoder('utf-8', { fatal: true });
289
+
290
+ function safeStringTokenLength(bytes: Uint8Array): number | undefined {
291
+ if (bytes.length === 0) return undefined;
292
+ let length = 0;
293
+ for (const byte of bytes) {
294
+ if (byte < 0x20 || byte === 0x22 || byte === 0x5c) return undefined;
295
+ // Code points equal non-continuation bytes in valid UTF-8.
296
+ if ((byte & 0xc0) !== 0x80) length++;
297
+ }
298
+ try {
299
+ safeStringDecoder.decode(bytes);
300
+ } catch {
301
+ return undefined;
302
+ }
303
+ return length;
304
+ }
305
+
306
+ function boundedStringMask(tokenizer: CachedTokenizer, capacity: number): { mask: Uint32Array; count: number } {
307
+ if (capacity >= tokenizer.maxStringSafeLength) {
308
+ return { mask: tokenizer.stringSafeMask, count: tokenizer.stringSafeCount };
309
+ }
310
+ let cached = tokenizer.boundedStringMasks.get(capacity);
311
+ if (cached !== undefined) return cached;
312
+ const mask = new Uint32Array(Math.ceil(tokenizer.data.tokens.length / 32));
313
+ let count = 0;
314
+ for (let tokenId = 0; tokenId < tokenizer.stringSafeLengths.length; ++tokenId) {
315
+ const length = tokenizer.stringSafeLengths[tokenId];
316
+ if (length === 0 || length > capacity) continue;
317
+ setBit(mask, tokenId);
318
+ count++;
319
+ }
320
+ cached = { mask, count };
321
+ tokenizer.boundedStringMasks.set(capacity, cached);
322
+ return cached;
323
+ }
324
+
325
+ function isJsonWhitespace(bytes: Uint8Array): boolean {
326
+ return bytes.length > 0 && bytes.every((byte) => byte === 0x09 || byte === 0x0a || byte === 0x0d || byte === 0x20);
327
+ }
328
+
329
+ function setBit(mask: Uint32Array, tokenId: number): void {
330
+ mask[tokenId >>> 5] |= 1 << (tokenId & 31);
331
+ }
@@ -0,0 +1,2 @@
1
+ export { createTokenConstraint, prepareTokenizer, type TokenConstraint } from './constraint';
2
+ export type { JSONSchema, TokenizerSource } from './types';