Skip to content

Commit f2ea965

Browse files
committed
updated applyChatTemplate with lazy memoization
1 parent ae0ead8 commit f2ea965

2 files changed

Lines changed: 42 additions & 1 deletion

File tree

Sources/Tokenizers/Tokenizer.swift

Lines changed: 13 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -284,6 +284,9 @@ public class PreTrainedTokenizer: Tokenizer {
284284
private let tokenizerConfig: Config
285285

286286
private let cleanUpTokenizationSpaces: Bool
287+
288+
// Cache for compiled Jinja templates keyed by their literal template string
289+
private var compiledChatTemplateCache: [String: Template] = [:]
287290

288291
public required init(tokenizerConfig: Config, tokenizerData: Config) throws {
289292
var addedTokens: [String: Int] = [:]
@@ -332,6 +335,15 @@ public class PreTrainedTokenizer: Tokenizer {
332335
model = try TokenizerModel.from(tokenizerConfig: tokenizerConfig, tokenizerData: tokenizerData, addedTokens: addedTokens)
333336
}
334337

338+
private func compiledTemplate(for templateString: String) throws -> Template {
339+
if let cached = compiledChatTemplateCache[templateString] {
340+
return cached
341+
}
342+
let compiled = try Template(templateString)
343+
compiledChatTemplateCache[templateString] = compiled
344+
return compiled
345+
}
346+
335347
func preTokenize(_ text: String, options: PreTokenizerOptions) -> [String] {
336348
guard let preTokenizer else { return [text] }
337349
return preTokenizer(text: text, options: options)
@@ -530,7 +542,7 @@ public class PreTrainedTokenizer: Tokenizer {
530542
throw TokenizerError.missingChatTemplate
531543
}
532544

533-
let template = try Template(selectedChatTemplate)
545+
let template = try compiledTemplate(for: selectedChatTemplate)
534546
var context: [String: Any] = [
535547
"messages": messages,
536548
"add_generation_prompt": addGenerationPrompt,

Tests/TokenizersTests/ChatTemplateTests.swift

Lines changed: 29 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -6,6 +6,7 @@
66
//
77

88
import Tokenizers
9+
import Foundation
910
import XCTest
1011

1112
class ChatTemplateTests: XCTestCase {
@@ -261,4 +262,32 @@ class ChatTemplateTests: XCTestCase {
261262
}
262263
}
263264
}
265+
266+
/// Performance: cached vs uncached template application
267+
func testApplyChatTemplatePerformanceCached() async throws {
268+
let tokenizer = try await AutoTokenizer.from(pretrained: "microsoft/Phi-3-mini-128k-instruct")
269+
270+
// Purposely reuse the same template literal to hit the memoized compiled template
271+
let mistral7BDefaultTemplate = "{{bos_token}}{% for message in messages %}{% if (message['role'] == 'user') != (loop.index0 % 2 == 0) %}{{ raise_exception('Conversation roles must alternate user/assistant/user/assistant/...') }}{% endif %}{% if message['role'] == 'user' %}{{ ' [INST] ' + message['content'] + ' [/INST]' }}{% elif message['role'] == 'assistant' %}{{ ' ' + message['content'] + ' ' + eos_token}}{% else %}{{ raise_exception('Only user and assistant roles are supported!') }}{% endif %}{% endfor %}"
272+
273+
// Prime cache once
274+
_ = try tokenizer.applyChatTemplate(messages: messages, chatTemplate: mistral7BDefaultTemplate)
275+
276+
measure(metrics: [XCTClockMetric()]) {
277+
_ = try! tokenizer.applyChatTemplate(messages: messages, chatTemplate: mistral7BDefaultTemplate)
278+
}
279+
}
280+
281+
/// Performance: simulate uncached runs by varying the template to bypass memoization
282+
func testApplyChatTemplatePerformanceUncached() async throws {
283+
let tokenizer = try await AutoTokenizer.from(pretrained: "microsoft/Phi-3-mini-128k-instruct")
284+
285+
let baseTemplate = "{{bos_token}}{% for message in messages %}{% if (message['role'] == 'user') != (loop.index0 % 2 == 0) %}{{ raise_exception('Conversation roles must alternate user/assistant/user/assistant/...') }}{% endif %}{% if message['role'] == 'user' %}{{ ' [INST] ' + message['content'] + ' [/INST]' }}{% elif message['role'] == 'assistant' %}{{ ' ' + message['content'] + ' ' + eos_token}}{% else %}{{ raise_exception('Only user and assistant roles are supported!') }}{% endif %}{% endfor %}"
286+
287+
measure(metrics: [XCTClockMetric()]) {
288+
// Make the template string unique each iteration to force a fresh compilation
289+
let uniqueTemplate = baseTemplate + "{# perf \(UUID().uuidString) #}"
290+
_ = try! tokenizer.applyChatTemplate(messages: messages, chatTemplate: uniqueTemplate)
291+
}
292+
}
264293
}

0 commit comments

Comments
 (0)