Skip to content
Merged
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
1 change: 1 addition & 0 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -641,6 +641,7 @@ curl http://localhost:5413/v1/chat/completions \
| `--num-draft-tokens` | `4` | Tokens per speculation round. Auto-capped to 1 when combined with `--stream-experts`. |
| `--dflash` | `false` | Enable DFlash block-diffusion speculative decoding. Requires a compatible DFlash draft model |
| `--dflash-block-size`| (auto) | Number of tokens per DFlash draft block. Defaults to draft model config |
| `--no-token-echo` | `false` | Stop echoing generated tokens to stdout as they stream. Request log lines (`srv ...`) are unaffected |

## 🔧 Per-Request API Parameters

Expand Down
79 changes: 65 additions & 14 deletions Sources/SwiftLM/Server.swift
Original file line number Diff line number Diff line change
Expand Up @@ -558,6 +558,9 @@ struct MLXServer: AsyncParsableCommand {
@Flag(name: .long, help: "Enable thinking/reasoning mode (Qwen3.5 etc). Default: disabled")
var thinking: Bool = false

@Flag(name: .long, help: "Do not echo generated tokens to stdout as they stream. Request log lines are unaffected.")
var noTokenEcho: Bool = false

@Flag(name: .long, help: "Enable VLM (vision-language model) mode for image inputs")
var vision: Bool = false

Expand Down Expand Up @@ -1355,6 +1358,7 @@ struct MLXServer: AsyncParsableCommand {
minP: self.minP,
repeatPenalty: self.repeatPenalty,
thinking: self.thinking,
tokenEcho: !self.noTokenEcho,
isVision: loadedAsVision,
prefillSize: self.prefillSize,
turboKV: self.turboKV,
Expand Down Expand Up @@ -1648,6 +1652,8 @@ struct ServerConfig: Sendable {
let minP: Float?
let repeatPenalty: Float?
let thinking: Bool
/// Echo each generated chunk to stdout (`--no-token-echo` turns it off).
let tokenEcho: Bool
let isVision: Bool
let prefillSize: Int
/// When true, each KVCacheSimple layer compresses history > 8192 tokens to 3-bit PolarQuant.
Expand Down Expand Up @@ -2347,7 +2353,7 @@ func handleChatCompletion(
startGeneration: startGeneration, modelId: modelId, stopSequences: stopSequences,
includeUsage: includeUsage, promptTokenCount: promptTokenCount,
enableThinking: enableThinking, thinkingPreOpened: thinkingPreOpened,
jsonMode: jsonMode, slot: slot,
jsonMode: jsonMode, tokenEcho: config.tokenEcho, slot: slot,
stats: stats, genStart: genStart, prefillStart: prefillStart,
emitPrefillProgress: emitPrefillProgress
)
Expand All @@ -2356,8 +2362,9 @@ func handleChatCompletion(
return try await handleChatNonStreaming(
stream: stream, modelId: modelId, stopSequences: stopSequences,
promptTokenCount: promptTokenCount, enableThinking: enableThinking,
thinkingPreOpened: thinkingPreOpened, jsonMode: jsonMode, slot: slot,
stats: stats, genStart: genStart, prefillStart: prefillStart, onPrefillDone: onPrefillDone
thinkingPreOpened: thinkingPreOpened, jsonMode: jsonMode, tokenEcho: config.tokenEcho,
slot: slot, stats: stats, genStart: genStart, prefillStart: prefillStart,
onPrefillDone: onPrefillDone
)
}
}
Expand Down Expand Up @@ -2537,6 +2544,7 @@ func handleChatStreaming(
enableThinking: Bool = false,
thinkingPreOpened: Bool = false,
jsonMode: Bool = false,
tokenEcho: Bool = true,
slot: GenerationSlot,
stats: ServerStats,
genStart: Date,
Expand Down Expand Up @@ -2602,6 +2610,10 @@ func handleChatStreaming(
// arriving in a later chunk merges with the previous one.
var emittedTextCount = 0
var heldStopTail = ""
let stopWindow = stopScanWindow(stopSequences)
// Characters appended since the last stop scan. Usually one chunk, but JSON-mode
// buffering skips the scan for its first chunks, so they must be covered later.
var unscannedStopChars = 0
// Unconditional cleanup: guarantees heartbeat is cancelled and the
// generation slot is returned on ALL exit paths (normal completion,
// startGeneration failure, client disconnect, or task cancellation).
Expand Down Expand Up @@ -2638,6 +2650,7 @@ func handleChatStreaming(
case .chunk(let text, _):
completionTokenCount += 1
fullText += text
unscannedStopChars += text.count
// GPU yield: prevent Metal from starving macOS WindowServer
if completionTokenCount % 8 == 0 {
try? await Task.sleep(for: .microseconds(50))
Expand All @@ -2657,8 +2670,10 @@ func handleChatStreaming(
if let onPrefillDone { await onPrefillDone() }
firstToken = false
}
print(text, terminator: "")
fflush(stdout)
if tokenEcho {
print(text, terminator: "")
fflush(stdout)
}

// ── JSON mode buffering: accumulate early tokens, strip prefix, then flush ──
if jsonBuffering {
Expand Down Expand Up @@ -2710,7 +2725,9 @@ func handleChatStreaming(
// shows no match and its tail is the opening of one. Emitting that tail
// hands the client the very text it asked to be cut (#126), so the
// ambiguous suffix is withheld until the next chunk resolves it.
let stopHit = checkStopSequences(fullText, stopSequences: stopSequences)
let stopHit = checkStopSequences(
fullText, stopSequences: stopSequences, lookback: stopWindow + unscannedStopChars)
unscannedStopChars = 0
var survivingText: String
if let (trimmedFull, _) = stopHit {
// The stop completed. Everything the client is owed is the part of
Expand Down Expand Up @@ -2745,6 +2762,8 @@ func handleChatStreaming(
}
}
cont.yield(sseChunk(modelId: modelId, reasoningContent: nil, content: nil, finishReason: "stop"))
// Stopping here means `.info` never arrives, so usage is the chunk
// count: a slight undercount if the decoder buffered any tokens.
let genDur = Date().timeIntervalSince(genStart)
let genTokPerSec = genDur > 0 ? Double(completionTokenCount) / genDur : 0
if includeUsage {
Expand Down Expand Up @@ -2783,6 +2802,9 @@ func handleChatStreaming(
print("[SwiftLM] Rejected tool call: reason=\(rejection.reason) tool=\(rejection.toolName ?? "?") detail=\(rejection.detail ?? "n/a")")

case .info(let info):
// The chunk count undercounts: tokens the decoder buffers (tool-call
// bodies, partial text) emit no chunk. `info` carries the real count.
completionTokenCount = info.generationTokenCount
heartbeatTask?.cancel()
heartbeatTask = nil
activePrefillProgressHook = nil
Expand Down Expand Up @@ -2889,6 +2911,7 @@ func handleChatNonStreaming(
enableThinking: Bool = false,
thinkingPreOpened: Bool = false,
jsonMode: Bool = false,
tokenEcho: Bool = true,
slot: GenerationSlot,
stats: ServerStats,
genStart: Date,
Expand Down Expand Up @@ -2920,8 +2943,10 @@ func handleChatNonStreaming(
if let onPrefillDone { await onPrefillDone() }
firstToken = false
}
print(text, terminator: "")
fflush(stdout)
if tokenEcho {
print(text, terminator: "")
fflush(stdout)
}
case .toolCall(let tc):
let argsJson = serializeToolCallArgs(tc.function.arguments)
collectedToolCalls.append(ToolCallResponse(
Expand All @@ -2935,6 +2960,7 @@ func handleChatNonStreaming(
print("[SwiftLM] Rejected tool call: reason=\(rejection.reason) tool=\(rejection.toolName ?? "?") detail=\(rejection.detail ?? "n/a")")
case .info(let info):
generationStopReason = info.stopReason
completionTokenCount = info.generationTokenCount // real count; see handleChatStreaming
}
}
print("") // end the real-time token stream line
Expand All @@ -2952,8 +2978,8 @@ func handleChatNonStreaming(
default:
finishReason = "stop"
}
if checkStopSequences(fullText, stopSequences: stopSequences) != nil {
fullText = checkStopSequences(fullText, stopSequences: stopSequences)!.0
if let (trimmedText, _) = checkStopSequences(fullText, stopSequences: stopSequences) {
fullText = trimmedText
finishReason = "stop"
}

Expand Down Expand Up @@ -3176,6 +3202,7 @@ func handleTextStreaming(
// a stop sequence (#133). Local to this loop; the chat path has its own pair.
var emittedTextCount = 0
var heldStopTail = ""
let stopWindow = stopScanWindow(stopSequences)
// Unconditional cleanup: cancels the heartbeat and returns the generation
// slot on ALL exit paths (completion, startGeneration failure, disconnect).
defer {
Expand Down Expand Up @@ -3217,7 +3244,9 @@ func handleTextStreaming(
// tail, and track characters actually released rather than deriving a
// position from chunk boundaries — that arithmetic was wrong in both
// directions on the chat path before it was replaced.
if let (trimmedText, _) = checkStopSequences(fullText, stopSequences: stopSequences) {
if let (trimmedText, _) = checkStopSequences(
fullText, stopSequences: stopSequences, lookback: stopWindow + text.count
) {
let remainder = String(
trimmedText.dropFirst(min(emittedTextCount, trimmedText.count)))
if !remainder.isEmpty {
Expand All @@ -3241,6 +3270,9 @@ func handleTextStreaming(
// Text-completion endpoint: tool calling has no wire representation here.
break
case .info(let info):
// The chunk count undercounts: tokens the decoder buffers (tool-call
// bodies, partial text) emit no chunk. `info` carries the real count.
completionTokenCount = info.generationTokenCount
heartbeatTask?.cancel()
heartbeatTask = nil
activePrefillProgressHook = nil
Expand Down Expand Up @@ -3303,7 +3335,9 @@ func handleTextNonStreaming(
if completionTokenCount % 8 == 0 {
try? await Task.sleep(for: .microseconds(50))
}
case .toolCall, .rejectedToolCall, .info:
case .info(let info):
completionTokenCount = info.generationTokenCount // real count; see handleChatStreaming
case .toolCall, .rejectedToolCall:
break
}
}
Expand Down Expand Up @@ -3533,10 +3567,20 @@ func mtpContext(main: ModelContext, assistant: (any DualModelMTP)?) -> ModelCont
/// happened to be listed first kept everything between the real stop and that one, so
/// `stop: ["\nUser:", "X"]` against `"abXc\nUser:"` streamed `"abXc"` when the client
/// had asked to stop at `X` (#126).
func checkStopSequences(_ text: String, stopSequences: [String]) -> (String, String)? {
///
/// `lookback` limits the search to the last `lookback` characters. Streaming callers
/// pass the length of text appended since their last scan plus `stopScanWindow(_:)`:
/// everything before it was already checked without a match, so a new match must end
/// inside the new text.
/// Without it, re-scanning the whole response on every chunk is O(n²). `nil` scans all.
func checkStopSequences(_ text: String, stopSequences: [String], lookback: Int? = nil) -> (String, String)? {
let searchStart = lookback.flatMap {
text.index(text.endIndex, offsetBy: -$0, limitedBy: text.startIndex)
} ?? text.startIndex
let searched = text[searchStart...]
var earliest: (index: String.Index, stop: String)?
for stop in stopSequences where !stop.isEmpty {
guard let range = text.range(of: stop) else { continue }
guard let range = searched.range(of: stop) else { continue }
if earliest == nil || range.lowerBound < earliest!.index {
earliest = (range.lowerBound, stop)
}
Expand All @@ -3545,6 +3589,13 @@ func checkStopSequences(_ text: String, stopSequences: [String]) -> (String, Str
return (String(text[text.startIndex..<earliest.index]), earliest.stop)
}

/// Characters of already-checked text a streaming stop scan must look back over: the
/// longest stop string, plus one because a chunk can merge with the previous character
/// into a single grapheme.
func stopScanWindow(_ stopSequences: [String]) -> Int {
(stopSequences.map(\.count).max() ?? 0) + 1
}

// ── Helpers ───────────────────────────────────────────────────────────────────

func jsonHeaders() -> HTTPFields {
Expand Down
85 changes: 85 additions & 0 deletions tests/SwiftLMTests/StopSequenceTests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -52,4 +52,89 @@ final class StopSequenceTests: XCTestCase {
let result = checkStopSequences("日本の首都は東京です<|end|>", stopSequences: ["<|end|>"])
XCTAssertEqual(result?.0, "日本の首都は東京です")
}

// MARK: Bounded streaming scan

/// Feeds `chunks` the way the streaming handlers do, scanning every `scanEvery` chunks.
/// `bounded: false` is the old full rescan, the oracle the bounded scan must match.
private func streamScan(
_ chunks: [String], stops: [String], bounded: Bool, scanEvery: Int = 1
) -> (String, String)? {
var full = ""
var unscanned = 0
for (i, chunk) in chunks.enumerated() {
full += chunk
unscanned += chunk.count
guard (i + 1) % scanEvery == 0 || i == chunks.count - 1 else { continue }
let lookback = bounded ? stopScanWindow(stops) + unscanned : nil
unscanned = 0
if let hit = checkStopSequences(full, stopSequences: stops, lookback: lookback) {
return hit
}
}
return nil
}

/// Splits `text` into chunks whose sizes cycle through `sizes`.
private func chunked(_ text: String, sizes: [Int]) -> [String] {
var chunks: [String] = []
var rest = Substring(text)
var i = 0
while !rest.isEmpty {
let n = sizes[i % sizes.count]
chunks.append(String(rest.prefix(n)))
rest = rest.dropFirst(n)
i += 1
}
return chunks
}

func testBoundedScanFindsStopStraddlingChunks() {
let result = streamScan(["answer\nUs", "er: next"], stops: ["\nUser:"], bounded: true)
XCTAssertEqual(result?.0, "answer")
XCTAssertEqual(result?.1, "\nUser:")
}

/// JSON mode skips the scan for its buffered chunks; one later scan must still see them.
func testBoundedScanCoversChunksSkippedSinceLastScan() {
let result = streamScan(
["{\"a\": 1}", "<|end|>", "tail", "more"], stops: ["<|end|>"], bounded: true, scanEvery: 4)
XCTAssertEqual(result?.0, "{\"a\": 1}")
}

func testBoundedScanMatchesFullRescanForEveryChunking() {
let cases: [(String, [String])] = [
("plain text then <|im_end|> and after", ["<|im_end|>", "<end_of_turn>"]),
("abXc\nUser: trailing", ["\nUser:", "X"]),
("ab<|eot_id|>cd<turn|>", ["<turn|>", "<|eot_id|>"]),
("日本の首都は東京です🎌<|end|>🎌", ["<|end|>"]),
("cafe\u{301} STOP", ["STOP"]),
("👨‍👩‍👧 family STOP", ["STOP", "family"]),
("no stop anywhere in this response", ["<|end|>", "\n\n"]),
("overlapping abcd", ["abcd", "bc"]),
]
let chunkings: [[Int]] = [[1], [2], [3], [5], [1, 4, 2], [7, 1], [100]]
for (text, stops) in cases {
for sizes in chunkings {
let chunks = chunked(text, sizes: sizes)
let bounded = streamScan(chunks, stops: stops, bounded: true)
let full = streamScan(chunks, stops: stops, bounded: false)
XCTAssertEqual(bounded?.0, full?.0, "\(text) chunked \(sizes)")
XCTAssertEqual(bounded?.1, full?.1, "\(text) chunked \(sizes)")
}
}
}

/// A combining mark opening a chunk merges into the previous chunk's last character.
func testBoundedScanSurvivesGraphemeMergeAcrossChunks() {
let chunks = ["cafe", "\u{301}ST", "OP"]
let bounded = streamScan(chunks, stops: ["STOP"], bounded: true)
XCTAssertEqual(bounded?.0, "caf\u{E9}")
XCTAssertEqual(bounded?.0, streamScan(chunks, stops: ["STOP"], bounded: false)?.0)
}

func testLookbackLongerThanTextScansAll() {
let result = checkStopSequences("x<|end|>", stopSequences: ["<|end|>"], lookback: 1000)
XCTAssertEqual(result?.0, "x")
}
}
Loading