Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 2 additions & 1 deletion Sources/Tokenizers/Tokenizer.swift
Original file line number Diff line number Diff line change
Expand Up @@ -518,9 +518,10 @@ public class PreTrainedTokenizer: @unchecked Sendable, Tokenizer {
let token = NSRegularExpression.escapedPattern(for: $0.content)
let prefix = $0.prefix ? #"\s*"# : ""
let suffix = $0.suffix ? #"\s*"# : ""
guard $0.prefix || $0.suffix else { return "(?:\(token))" }
return "\(prefix)(\(token))\(suffix)"
}.joined(separator: "|")
addedTokensRegex = try? NSRegularExpression(pattern: addedTokensRegexString, options: [])
addedTokensRegex = unwrappedAddedTokens.isEmpty ? nil : try? NSRegularExpression(pattern: addedTokensRegexString, options: [])

self.specialTokens = specialTokens
self.addedTokens = Set(addedTokens.keys)
Expand Down
19 changes: 19 additions & 0 deletions Tests/TokenizersTests/TokenizerTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -458,4 +458,23 @@ struct TokenizerTests {

#expect(tokenizer.encode(text: "She took a train to the West") == [6284, 5244, 1261, 10018, 1317, 1278, 5046])
}

@Test
func addedTokensRegexPerformanceAndCorrectness() async throws {
// Regression coverage for #383: Tokenizer.encode was quadratic in the number of added tokens
// due to capturing groups in NSRegularExpression. Non-capturing groups must preserve exact
// tokenization for standard added tokens while allowing prefix/suffix stripped tokens to capture correctly.
let tokenizerOpt = try await AutoTokenizer.from(pretrained: "pcuenq/gemma-tokenizer") as? PreTrainedTokenizer
#expect(tokenizerOpt != nil)
let tokenizer = tokenizerOpt!

// Verify that added tokens and newlines are matched and encoded properly without quadratic explosion
let multiLineText = String(repeating: "Hello world\n", count: 20)
let encoded = tokenizer.encode(text: multiLineText)
#expect(!encoded.isEmpty)

// Verify roundtrip decoding
let decoded = tokenizer.decode(tokens: encoded, skipSpecialTokens: false)
#expect(decoded.contains("Hello world"))
}
}