Repository navigation
Split inference includeLogits bool into orthogonal tokens + logits axes #297
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
9431478
6bdb59b
b0be4b0
95a6655
8cf894b
62a56b7
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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) { | ||
|
|
@@ -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 | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Public API breaking when we remove this
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 🤔
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 | ||
|
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 | ||
|
|
@@ -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( | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.