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 @@ -136,6 +136,17 @@ public final class CoreAISequentialVLMEngine: MultimodalInferenceEngine, @unchec

// MARK: - Init

/// Validates the LLM decoder's state-name contract. Accepts 2–4 states (KV cache + optional
/// persistent conv/recurrent states), matching `CoreAISequentialEngine`; hybrid decoders carry
/// extra conv/recurrent states beyond the two KV-cache states.
static func validateLLMStateNames(_ stateNames: [String]) throws {
guard stateNames.count >= 2 && stateNames.count <= 4 else {
throw InferenceRuntimeError.invalidOutputType(
"VLM LLM function expected 2–4 states (KV cache + optional persistent states), "
+ "got \(stateNames.count): \(stateNames)")
}
}

/// Initialize the VLM engine with separate model assets for vision, embed, and LLM.
///
/// - Parameters:
Expand Down Expand Up @@ -232,11 +243,7 @@ public final class CoreAISequentialVLMEngine: MultimodalInferenceEngine, @unchec
"VLM LLM function expected 2 inputs (in_embeddings, position_ids), "
+ "got \(llmDesc.inputNames.count): \(llmDesc.inputNames)")
}
guard llmDesc.stateNames.count == 2 else {
throw InferenceRuntimeError.invalidOutputType(
"VLM LLM function expected 2 states (KV cache), "
+ "got \(llmDesc.stateNames.count): \(llmDesc.stateNames)")
}
try Self.validateLLMStateNames(llmDesc.stateNames)
guard llmDesc.outputNames.count >= 1 else {
throw InferenceRuntimeError.invalidOutputType(
"VLM LLM function expected at least 1 output (logits), "
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,40 @@
// Copyright 2026 Apple Inc.
//
// Use of this source code is governed by a BSD-3-clause license that can
// be found in the LICENSE file or at https://opensource.org/licenses/BSD-3-Clause

import Foundation
import Testing

@testable import CoreAILanguageModels

@Suite("CoreAISequentialVLMEngine state-name contract")
struct CoreAISequentialVLMEngineTests {
@Test("accepts a full-attention two-state KV decoder")
func acceptsTwoStates() throws {
try CoreAISequentialVLMEngine.validateLLMStateNames(["k_cache", "v_cache"])
}

@Test("accepts a hybrid four-state decoder (conv + recurrent states)")
func acceptsFourStates() throws {
// Hybrid decoders carry conv/recurrent states beyond the two KV-cache states; the
// engine previously rejected anything but exactly two states.
try CoreAISequentialVLMEngine.validateLLMStateNames(
["k_cache", "v_cache", "conv_states", "recurrent_states"])
}

@Test("rejects fewer than two states")
func rejectsOneState() {
#expect(throws: InferenceRuntimeError.self) {
try CoreAISequentialVLMEngine.validateLLMStateNames(["k_cache"])
}
}

@Test("rejects more than four states")
func rejectsFiveStates() {
#expect(throws: InferenceRuntimeError.self) {
try CoreAISequentialVLMEngine.validateLLMStateNames(
["s0", "s1", "s2", "s3", "s4"])
}
}
}