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
Original file line number Diff line number Diff line change
Expand Up @@ -260,7 +260,7 @@ extension ConstrainedDecodingStrategy.ConstrainedDecodedSequence {
private let inferenceEngine: any InferenceEngine
private let samplingConfiguration: SamplingConfiguration
private let stopSequences: StopSequences
private let constrainedOptions = InferenceOptions(maxTokens: 1, includeLogits: true)
private let constrainedOptions = InferenceOptions.guided(maxTokens: 1)

// Generation state, seeded eagerly from the prepared setup.
private var session: ConstrainedGenerationSession?
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -188,7 +188,7 @@ public struct ConstrainedGenerator: DecodingStrategy {
for _ in 0..<maxTokens {
if session.isTerminated { break }

let options = InferenceOptions(maxTokens: 1, includeLogits: true)
let options = InferenceOptions.guided(maxTokens: 1)

var rawLogits: [LogitsScalarType]? = nil
for try await output in try await inferenceEngine.generate(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -168,7 +168,7 @@ public protocol DecodingStrategy: Sendable {
/// - tokenizer: Tokenizer for encoding/decoding
/// - inferenceEngine: Engine for model inference
/// - samplingConfiguration: Sampling parameters (temperature, topK, etc.)
/// - options: Inference options (maxTokens, includeLogits)
/// - options: Inference options (maxTokens, tokens, logits)
/// - stopSequences: Token sequences that halt generation
/// - Returns: Stream of `GenerationResult` (text + optional logits)
func decode(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,7 @@ public struct VanillaDecodingStrategy: DecodingStrategy {
/// - tokenizer: Tokenizer for encoding/decoding
/// - inferenceEngine: Engine for model inference
/// - samplingConfiguration: Sampling parameters (temperature, topK, etc.)
/// - options: Inference options (maxTokens, includeLogits)
/// - options: Inference options (maxTokens, tokens, logits)
/// - stopSequences: Token sequences that halt generation
/// - Returns: Stream of `GenerationResult` (text + optional logits)
public func decode(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -117,7 +117,7 @@ final class CoreAIPipelinedEngine: InferenceEngine, ConstrainedGenerationCapable
samplingConfiguration: SamplingConfiguration,
inferenceOptions: InferenceOptions
) async throws -> GenerationSequence {
if inferenceOptions.includeLogits {
if inferenceOptions.returnsLogits {
throw InferenceRuntimeError.invalidArgument(
"CoreAI pipelined engine does not support logits (GPU-side sampling). "
+ "Use a sequential engine for constrained generation or evaluation."
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -603,7 +603,7 @@ extension CoreAISequentialEngine.GenerationSequence {
self.engine = engine
self.sessionState = sessionState
self.samplingConfiguration = samplingConfiguration.normalized()
self.returnsLogits = inferenceOptions.includeLogits
self.returnsLogits = inferenceOptions.returnsLogits
self.forcedContinuation = inferenceOptions.forcedContinuation
self.stopReasonStore = stopReasonStore
self.generationToken = generationToken
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -983,7 +983,7 @@ extension CoreAISequentialVLMEngine.GenerationSequence {
) {
self.engine = engine
self.samplingConfiguration = samplingConfiguration.normalized()
self.returnsLogits = inferenceOptions.includeLogits
self.returnsLogits = inferenceOptions.returnsLogits
self.forcedContinuation = inferenceOptions.forcedContinuation
self.stopReasonStore = stopReasonStore
self.generationToken = generationToken
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -793,7 +793,7 @@ extension StaticShapeEngine.GenerationSequence {
) {
self.engine = engine
self.samplingConfiguration = samplingConfiguration
self.returnsLogits = inferenceOptions.includeLogits
self.returnsLogits = inferenceOptions.returnsLogits
self.forcedContinuation = inferenceOptions.forcedContinuation
self.stopReasonStore = stopReasonStore
self.generationToken = generationToken
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,7 +21,7 @@ public typealias LogitsScalarType = Float
public struct InferenceOutput: Sendable {
public let tokenId: Int32

/// Populated when `InferenceOptions.includeLogits` is true. Shape: [vocabSize].
/// Populated when `InferenceOptions.logits` is not `.none`. Shape: [vocabSize].
public let logits: [LogitsScalarType]?

public init(tokenId: Int32, logits: [LogitsScalarType]? = nil) {
Expand All @@ -32,26 +32,111 @@ public struct InferenceOutput: Sendable {

// MARK: - Inference Options

/// Whether to sample a token this step.
public enum TokenRequest: Sendable, Equatable {
case none
case sample
}

/// Which positions' logits to return.
public enum LogitsRequest: Sendable, Equatable {
case none
case lastPosition
case allPositions
}

/// Controls what the engine produces and how much.
/// Struct-based for additive extensibility (future: embeddings, attention maps).
public struct InferenceOptions: Sendable {
/// Max tokens to generate. Nil = until EOS or context limit.
public var maxTokens: Int?
/// Include raw logits in each `InferenceOutput`. May incur GPU→CPU copy cost.
public var includeLogits: Bool
Comment thread
stikves marked this conversation as resolved.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Public API breaking when we remove this

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

let's mark this deprecated

/// When set, engines use these token IDs instead of sampling.
/// Used by MMLU-style evaluation to compute P(continuation|context).
public var tokens: TokenRequest
public var logits: LogitsRequest
/// Force these token IDs instead of sampling (MMLU-style scoring).
public var forcedContinuation: [Int32]?

/// Raw initializer over the full `(tokens, logits)` space.
///
/// The two axes are intentionally orthogonal and this initializer is
/// deliberately permissive: every `(tokens, logits)` pair is representable
/// and there are no invalid combinations to reject. The `prefill`, `extend`,
/// `eval`, `guided`, and `sampleWithLogits` presets below cover all the
/// meaningful pairings; prefer them to constructing options by hand. The
/// only representable pair without a dedicated preset is
/// `(tokens: .sample, logits: .allPositions)` — sampling a token while also

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

grepped for .tokens, .logits, .allPositions and .lastPosition in swift/ and internal/. Outside InferenceEngine.swift, the only readers are the tests. Every engine (sequential, VLM, static-shape, pipelined) reads returnsLogits, which is just logits != .none.

So tokens: .none and .allPositions change nothing at runtime, so these are declared but no engine reads them 🤔

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, this will be handled in a follow up.

(The engines do extra work today)

/// returning logits at every position — which is odd but harmless.
public init(
maxTokens: Int? = nil,
includeLogits: Bool = false,
tokens: TokenRequest = .sample,
logits: LogitsRequest = .none,
forcedContinuation: [Int32]? = nil
) {
self.maxTokens = maxTokens
self.includeLogits = includeLogits
self.tokens = tokens
self.logits = logits
Comment thread
stikves marked this conversation as resolved.
self.forcedContinuation = forcedContinuation
}

/// Legacy initializer bridging the old `includeLogits` boolean to the
/// `(tokens, logits)` axes: `true` maps to `logits: .lastPosition`, `false` to
/// `logits: .none`. Tokens are sampled, as they always were under this API.
@available(
*, deprecated,
message: "Use the tokens:/logits: initializer; includeLogits: true == logits: .lastPosition."
)
public init(
maxTokens: Int? = nil,
includeLogits: Bool,
forcedContinuation: [Int32]? = nil
) {
self.init(
maxTokens: maxTokens,
tokens: .sample,
logits: includeLogits ? .lastPosition : .none,
forcedContinuation: forcedContinuation)
}

/// Whether this request asks the engine to return logits at any position.
///
/// Equivalent to `logits != .none`; use this instead of open-coding the
/// comparison so the intent reads clearly at the call site.
public var returnsLogits: Bool { logits != .none }

/// Legacy accessor for the old `includeLogits` boolean. Reads as `returnsLogits`;
/// setting `true` requests last-position logits, `false` requests none.
@available(*, deprecated, message: "Use `logits` / `returnsLogits` instead.")
public var includeLogits: Bool {
get { returnsLogits }
set { logits = newValue ? .lastPosition : .none }
}
}

/// Presets for the common (tokens, logits) pairs.
extension InferenceOptions {
/// Warm the KV cache; does not need token or logits
public static func prefill(maxTokens: Int? = nil) -> InferenceOptions {
InferenceOptions(maxTokens: maxTokens, tokens: .none, logits: .none)
}
/// Decode one token.
public static func extend(maxTokens: Int? = nil) -> InferenceOptions {
InferenceOptions(maxTokens: maxTokens, tokens: .sample, logits: .none)
}
/// Same as `extend`.
public static func `default`(maxTokens: Int? = nil) -> InferenceOptions {
.extend(maxTokens: maxTokens)
}
/// Logits at every position, for eval / PPL.
public static func eval(maxTokens: Int? = nil, forcedContinuation: [Int32]? = nil) -> InferenceOptions {
InferenceOptions(
maxTokens: maxTokens, tokens: .none, logits: .allPositions, forcedContinuation: forcedContinuation)
}
/// Last-row logits for constrained decoding.
public static func guided(maxTokens: Int? = nil) -> InferenceOptions {
InferenceOptions(maxTokens: maxTokens, tokens: .none, logits: .lastPosition)
}
/// Decode, and also return the last-row logits.
public static func sampleWithLogits(maxTokens: Int? = nil) -> InferenceOptions {
InferenceOptions(maxTokens: maxTokens, tokens: .sample, logits: .lastPosition)
}
}

// MARK: - Configuration Data Structures
Expand Down Expand Up @@ -96,7 +181,7 @@ public protocol InferenceEngine: Sendable {
/// - Parameters:
/// - input: Token IDs (prompt, context, or continuation).
/// - sampling: Sampling configuration (temperature, topK, etc.).
/// - generation: Inference options (maxTokens, includeLogits).
/// - generation: Inference options (maxTokens, tokens, logits).
/// - Returns: An `InferenceOutputSequence` — iterate for tokens, read
/// `stopReason` after the loop to learn why generation ended.
func generate(
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -83,7 +83,7 @@ public class TextGenerator {
tokenizer: tokenizer,
inferenceEngine: inferenceEngine,
samplingConfiguration: samplingConfiguration,
options: InferenceOptions(maxTokens: maxTokens, includeLogits: true),
options: .sampleWithLogits(maxTokens: maxTokens),
stopSequences: effectiveStopSequences
)

Expand Down Expand Up @@ -134,9 +134,8 @@ public class TextGenerator {
let continuationTokens = Array(
encoding.tokens[encoding.continuationStartIndex..<encoding.tokens.count])

let options = InferenceOptions(
let options = InferenceOptions.eval(
maxTokens: continuationTokens.count,
includeLogits: true,
forcedContinuation: continuationTokens
)

Expand Down Expand Up @@ -193,9 +192,8 @@ public class TextGenerator {

try await inferenceEngine.reset()

let options = InferenceOptions(
let options = InferenceOptions.eval(
maxTokens: continuationTokens.count,
includeLogits: true,
forcedContinuation: continuationTokens
)

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -154,7 +154,7 @@ public struct CoreAIVLMExecutor: LanguageModelExecutor {
with: embeddedInput,
tokens: promptTokens,
samplingConfiguration: SamplingConfiguration(temperature: 1.0, topK: 1),
inferenceOptions: InferenceOptions(maxTokens: maxTokens, includeLogits: false)
inferenceOptions: .extend(maxTokens: maxTokens)
)

var generatedCount = 0
Expand Down
2 changes: 1 addition & 1 deletion swift/Sources/Tools/benchmark/BenchmarkMain.swift
Original file line number Diff line number Diff line change
Expand Up @@ -177,7 +177,7 @@ struct LLMBenchmark: AsyncParsableCommand {
try? await Task.sleep(for: .milliseconds(50))
try await engine.reset()

let options = InferenceOptions(maxTokens: generationTokens, includeLogits: false)
let options = InferenceOptions.extend(maxTokens: generationTokens)
let start = SuspendingClock.now
let stream = try await engine.generate(
with: prompt, samplingConfiguration: sampling, inferenceOptions: options
Expand Down
6 changes: 2 additions & 4 deletions swift/Sources/Tools/llm-runner/LLMRunnerMain.swift
Original file line number Diff line number Diff line change
Expand Up @@ -1126,10 +1126,8 @@ struct LLMRunner: AsyncParsableCommand, Sendable {
with: embeddedInput,
tokens: vlmTokens,
samplingConfiguration: samplingConfiguration,
inferenceOptions: InferenceOptions(
maxTokens: maxTokens,
includeLogits: printLogits || saveLogits != nil
)
inferenceOptions: (printLogits || saveLogits != nil)
? .sampleWithLogits(maxTokens: maxTokens) : .extend(maxTokens: maxTokens)
)

CLILogger.log("VLM generate started, maxTokens=\(maxTokens)", component: "VLM")
Expand Down
3 changes: 1 addition & 2 deletions swift/Sources/Tools/llm-server/CompletionHandler.swift
Original file line number Diff line number Diff line change
Expand Up @@ -143,9 +143,8 @@ private func processOnePrompt(
let paddingToken = allTokens[0]
let continuation = Array(allTokens.dropFirst()) + [paddingToken]

let options = InferenceOptions(
let options = InferenceOptions.eval(
maxTokens: continuation.count,
includeLogits: true,
forcedContinuation: continuation
)

Expand Down
2 changes: 1 addition & 1 deletion swift/Tests/LanguageModelsTests/TestUtilities.swift
Original file line number Diff line number Diff line change
Expand Up @@ -157,7 +157,7 @@ class MockEngine: InferenceEngine, @unchecked Sendable {
generationToken: GenerationToken
) {
self.engine = engine
self.returnsLogits = inferenceOptions.includeLogits
self.returnsLogits = inferenceOptions.returnsLogits
self.forcedContinuation = inferenceOptions.forcedContinuation
self.stopReasonStore = stopReasonStore
self.generationToken = generationToken
Expand Down
43 changes: 29 additions & 14 deletions swift/Tests/LanguageModelsTests/UnifiedGenerationAPITests.swift
Original file line number Diff line number Diff line change
Expand Up @@ -37,14 +37,32 @@ struct InferenceOptionsTests {
func defaultOptions() {
let opts = InferenceOptions()
#expect(opts.maxTokens == nil)
#expect(opts.includeLogits == false)
#expect(opts.logits == .none)
#expect(opts.tokens == .sample)
}

@Test("Custom options")
func customOptions() {
let opts = InferenceOptions(maxTokens: 50, includeLogits: true)
let opts = InferenceOptions(maxTokens: 50, logits: .lastPosition)
#expect(opts.maxTokens == 50)
#expect(opts.includeLogits == true)
#expect(opts.logits == .lastPosition)
}

@Test("Presets map to the right (tokens, logits) pairs")
func presets() {
#expect(InferenceOptions.prefill().tokens == .none)
#expect(InferenceOptions.prefill().logits == .none)
#expect(InferenceOptions.extend().tokens == .sample)
#expect(InferenceOptions.extend().logits == .none)
#expect(InferenceOptions.eval().tokens == .none)
#expect(InferenceOptions.eval().logits == .allPositions)
#expect(InferenceOptions.guided().tokens == .none)
#expect(InferenceOptions.guided().logits == .lastPosition)
#expect(InferenceOptions.sampleWithLogits().tokens == .sample)
#expect(InferenceOptions.sampleWithLogits().logits == .lastPosition)
// no-arg default == .extend (sample, no logits) — byte-identical to old includeLogits:false
#expect(InferenceOptions().tokens == .sample)
#expect(InferenceOptions().logits == .none)
}
}

Expand Down Expand Up @@ -75,11 +93,11 @@ struct GenerateDefaultExtensionTests {
#expect(outputs[4].tokenId == 20)
}

@Test("generate() returns nil logits when includeLogits is false")
@Test("generate() returns nil logits when logits is .none")
func noLogitsWhenNotRequested() async throws {
let engine = MockEngine(tokens: [42])

let generation = InferenceOptions(maxTokens: 1, includeLogits: false)
let generation = InferenceOptions(maxTokens: 1, logits: .none)
for try await output in try await engine.generate(
with: [1],
samplingConfiguration: SamplingConfiguration.greedy,
Expand All @@ -89,11 +107,11 @@ struct GenerateDefaultExtensionTests {
}
}

@Test("generate() returns logits when includeLogits is true")
@Test("generate() returns logits when logits is requested")
func logitsWhenRequested() async throws {
let engine = MockEngine(tokens: [42], vocabSize: 50)

let generation = InferenceOptions(maxTokens: 1, includeLogits: true)
let generation = InferenceOptions(maxTokens: 1, logits: .lastPosition)
for try await output in try await engine.generate(
with: [1],
samplingConfiguration: SamplingConfiguration.greedy,
Expand All @@ -108,7 +126,7 @@ struct GenerateDefaultExtensionTests {
func logitsHighProbOnTarget() async throws {
let engine = MockEngine(tokens: [5], vocabSize: 10)

let generation = InferenceOptions(maxTokens: 1, includeLogits: true)
let generation = InferenceOptions(maxTokens: 1, logits: .lastPosition)
for try await output in try await engine.generate(
with: [1],
samplingConfiguration: SamplingConfiguration.greedy,
Expand All @@ -128,7 +146,7 @@ struct GenerateDefaultExtensionTests {
func noLogitsWhenVocabSizeNil() async throws {
let engine = MockEngine(tokens: [42], vocabSize: nil)

let generation = InferenceOptions(maxTokens: 1, includeLogits: true)
let generation = InferenceOptions(maxTokens: 1, logits: .lastPosition)
for try await output in try await engine.generate(
with: [1],
samplingConfiguration: SamplingConfiguration.greedy,
Expand Down Expand Up @@ -208,7 +226,7 @@ struct GenerateMultiCallTests {
for try await output in try await engine.generate(
with: tokens,
samplingConfiguration: .greedy,
inferenceOptions: InferenceOptions(maxTokens: 1, includeLogits: true)
inferenceOptions: .guided(maxTokens: 1)
) {
got = output
break // Only consume 1 token (GG pattern)
Expand Down Expand Up @@ -260,10 +278,7 @@ struct GenerateMultiCallTests {
for try await output in try await engine.generate(
with: [1, 2, 3],
samplingConfiguration: .greedy,
inferenceOptions: InferenceOptions(
includeLogits: true,
forcedContinuation: forced
)
inferenceOptions: .eval(forcedContinuation: forced)
) {
outputs.append(output)
}
Expand Down