From b80dbd8a63b20a532553885be9b64c5b38dcfa54 Mon Sep 17 00:00:00 2001 From: lightseek-bot <243258330+lightseek-bot@users.noreply.github.com> Date: Wed, 15 Jul 2026 06:38:42 +0000 Subject: [PATCH 1/8] feat(inkling): add end-to-end SMG support Signed-off-by: lightseek-bot <243258330+lightseek-bot@users.noreply.github.com> --- crates/multimodal/src/audio/mod.rs | 4 +- .../src/audio/processors/inkling.rs | 432 ++++++++++++++ crates/multimodal/src/audio/processors/mod.rs | 2 + crates/multimodal/src/registry/inkling.rs | 380 +++++++++++++ crates/multimodal/src/registry/mod.rs | 3 + crates/multimodal/src/types.rs | 35 +- crates/multimodal/src/vision/processor.rs | 10 + .../src/vision/processors/inkling.rs | 456 +++++++++++++++ .../multimodal/src/vision/processors/mod.rs | 3 + crates/multimodal/src/vision/transforms.rs | 150 ++++- .../inkling_preprocess_fingerprints.json | 73 +++ .../tests/inkling_preprocess_golden.rs | 126 +++++ crates/protocols/src/chat.rs | 101 +++- crates/reasoning_parser/src/factory.rs | 17 +- crates/reasoning_parser/src/lib.rs | 4 +- .../reasoning_parser/src/parsers/inkling.rs | 334 +++++++++++ crates/reasoning_parser/src/parsers/mod.rs | 2 + crates/reasoning_parser/src/traits.rs | 9 + crates/tool_parser/src/factory.rs | 13 +- crates/tool_parser/src/lib.rs | 6 +- crates/tool_parser/src/parsers/inkling.rs | 532 ++++++++++++++++++ crates/tool_parser/src/parsers/mod.rs | 2 + .../tool_parser/tests/tool_parser_inkling.rs | 179 ++++++ .../smg_grpc_servicer/tokenspeed/servicer.py | 11 + model_gateway/src/routers/grpc/context.rs | 3 +- .../src/routers/grpc/multimodal/process.rs | 112 +++- .../src/routers/grpc/regular/processor.rs | 24 +- .../grpc/regular/responses/conversions.rs | 1 + .../grpc/regular/stages/chat/preparation.rs | 35 +- .../regular/stages/messages/preparation.rs | 33 +- .../src/routers/grpc/regular/streaming.rs | 22 +- model_gateway/src/routers/grpc/router.rs | 1 + .../src/routers/grpc/utils/chat_utils.rs | 516 ++++++++++++++--- model_gateway/src/routers/grpc/utils/mod.rs | 2 +- .../src/routers/grpc/utils/parsers.rs | 31 + model_gateway/src/server.rs | 6 +- .../tests/api/request_formats_test.rs | 75 +++ 37 files changed, 3601 insertions(+), 144 deletions(-) create mode 100644 crates/multimodal/src/audio/processors/inkling.rs create mode 100644 crates/multimodal/src/registry/inkling.rs create mode 100644 crates/multimodal/src/vision/processors/inkling.rs create mode 100644 crates/multimodal/tests/fixtures/golden/inkling_preprocess_fingerprints.json create mode 100644 crates/multimodal/tests/inkling_preprocess_golden.rs create mode 100644 crates/reasoning_parser/src/parsers/inkling.rs create mode 100644 crates/tool_parser/src/parsers/inkling.rs create mode 100644 crates/tool_parser/tests/tool_parser_inkling.rs diff --git a/crates/multimodal/src/audio/mod.rs b/crates/multimodal/src/audio/mod.rs index dff8c6d76b..1c9c0ef07d 100644 --- a/crates/multimodal/src/audio/mod.rs +++ b/crates/multimodal/src/audio/mod.rs @@ -7,4 +7,6 @@ pub(crate) mod transforms; pub use decode::{decode_audio_mono_f32, DecodedAudio}; pub use processor::AudioPreProcessor; -pub use processors::{Qwen3AudioParams, Qwen3AudioProcessor}; +pub use processors::{ + InklingAudioEncoderParams, InklingAudioProcessor, Qwen3AudioParams, Qwen3AudioProcessor, +}; diff --git a/crates/multimodal/src/audio/processors/inkling.rs b/crates/multimodal/src/audio/processors/inkling.rs new file mode 100644 index 0000000000..1df6c48554 --- /dev/null +++ b/crates/multimodal/src/audio/processors/inkling.rs @@ -0,0 +1,432 @@ +//! Inkling audio preprocessing. +//! +//! Implements the Inkling feature-extraction pipeline for the model-facing part: +//! audio bytes are decoded to mono f32, resampled to the configured sample +//! rate, non-zero quiet signals are raised to the configured RMS floor, +//! transformed to Slaney-normalized log-mel magnitudes, then quantized to dMel +//! bin ids. The returned encoder input stores those integer bin ids as f32, +//! matching the checkpoint feature-extractor contract and allowing the transport +//! to apply the same configurable floating-point wire dtype as other modalities. + +use ndarray::Array2; +use rustfft::{num_complex::Complex32, FftPlanner}; + +use crate::{ + audio::{ + transforms::{bandlimited_resample, hann_window, mel_basis}, + AudioPreProcessor, DecodedAudio, + }, + encoder_inputs::{ModelSpecificValue, PreprocessedEncoderInputs}, + error::TransformError, + types::AudioClip, +}; + +#[derive(Debug, Clone)] +pub struct InklingAudioEncoderParams { + pub sample_rate: usize, + pub window_size_multiplier: f64, + pub n_fft: Option, + pub n_mels: usize, + pub num_dmel_bins: usize, + pub dmel_min_value: f64, + pub dmel_max_value: f64, + pub audio_token_duration_s: f64, + pub audio_rms_norm_floor: f64, +} + +impl Default for InklingAudioEncoderParams { + fn default() -> Self { + Self { + sample_rate: 16_000, + window_size_multiplier: 2.0, + n_fft: None, + n_mels: 80, + num_dmel_bins: 16, + dmel_min_value: -7.0, + dmel_max_value: 2.0, + audio_token_duration_s: 0.05, + audio_rms_norm_floor: 0.01, + } + } +} + +impl InklingAudioEncoderParams { + pub fn from_model_config(config: &serde_json::Value) -> Self { + let mut params = Self::default(); + let Some(audio_config) = config.get("audio_config") else { + return params; + }; + if let Some(v) = audio_config.get("n_mel_bins").and_then(|v| v.as_u64()) { + params.n_mels = v as usize; + } + if let Some(v) = audio_config.get("mel_vocab_size").and_then(|v| v.as_u64()) { + params.num_dmel_bins = v as usize; + } + if let Some(v) = audio_config.get("dmel_min_value").and_then(|v| v.as_f64()) { + params.dmel_min_value = v; + } + if let Some(v) = audio_config.get("dmel_max_value").and_then(|v| v.as_f64()) { + params.dmel_max_value = v; + } + if let Some(v) = audio_config + .get("audio_rms_norm_floor") + .and_then(|v| v.as_f64()) + { + params.audio_rms_norm_floor = v; + } + params + } + + fn hop_length(&self) -> Result { + exact_sample_count( + self.audio_token_duration_s * self.sample_rate as f64, + "audio_token_duration_s * sample_rate", + ) + } + + fn window_size(&self) -> Result { + exact_sample_count( + self.audio_token_duration_s * self.window_size_multiplier * self.sample_rate as f64, + "audio_token_duration_s * window_size_multiplier * sample_rate", + ) + } +} + +#[derive(Debug, Clone)] +pub struct InklingAudioProcessor { + params: InklingAudioEncoderParams, +} + +impl Default for InklingAudioProcessor { + fn default() -> Self { + Self::new() + } +} + +impl InklingAudioProcessor { + pub fn new() -> Self { + Self { + params: InklingAudioEncoderParams::default(), + } + } + + pub fn from_model_config(config: &serde_json::Value) -> Self { + Self { + params: InklingAudioEncoderParams::from_model_config(config), + } + } + + pub fn preprocess_decoded_clips( + &self, + clips: Vec, + ) -> Result { + if clips.is_empty() { + return Err(TransformError::EmptyBatch); + } + + let mut all_bins = Vec::new(); + let mut token_counts = Vec::with_capacity(clips.len()); + let mut item_sizes = Vec::with_capacity(clips.len()); + let mut tokens_per_item = Vec::with_capacity(clips.len()); + + for clip in clips { + let bins = self.preprocess_decoded(clip)?; + let shape = bins.shape(); + let num_tokens = shape[0]; + let n_mels = shape[1]; + token_counts.push(num_tokens); + tokens_per_item.push(num_tokens as i64); + item_sizes.push((n_mels as u32, num_tokens as u32)); + all_bins.extend(bins.into_raw_vec_and_offset().0); + } + + let total_tokens: usize = token_counts.iter().sum(); + let encoder_input = Array2::from_shape_vec((total_tokens, self.params.n_mels), all_bins) + .map_err(|e| { + TransformError::ShapeError(format!( + "failed to create Inkling audio encoder input [{total_tokens}, {}]: {e}", + self.params.n_mels + )) + })?; + + Ok( + PreprocessedEncoderInputs::new(encoder_input, token_counts, item_sizes).with_extra( + "tokens_per_item", + ModelSpecificValue::int_1d(tokens_per_item), + ), + ) + } + + fn preprocess_decoded(&self, decoded: DecodedAudio) -> Result, TransformError> { + if decoded.sample_rate == 0 { + return Err(TransformError::ShapeError( + "decoded audio sample rate must be positive".to_string(), + )); + } + if decoded.samples.is_empty() { + return Err(TransformError::ShapeError( + "decoded audio contains no samples".to_string(), + )); + } + let mut samples = if decoded.sample_rate == self.params.sample_rate { + decoded.samples + } else { + bandlimited_resample( + &decoded.samples, + decoded.sample_rate, + self.params.sample_rate, + )? + }; + normalize_audio_rms(&mut samples, self.params.audio_rms_norm_floor)?; + dmel_bins(&samples, &self.params) + } +} + +impl AudioPreProcessor for InklingAudioProcessor { + fn preprocess( + &self, + clips: &[std::sync::Arc], + ) -> Result { + self.preprocess_decoded_clips(clips.iter().map(|clip| clip.decoded().clone()).collect()) + } +} + +fn exact_sample_count(value: f64, name: &str) -> Result { + let rounded = value.round(); + if (value - rounded).abs() > 1e-6 { + return Err(TransformError::ShapeError(format!( + "{name} must resolve to an integer sample count, got {value}" + ))); + } + if rounded <= 0.0 { + return Err(TransformError::ShapeError(format!( + "{name} must be positive, got {rounded}" + ))); + } + Ok(rounded as usize) +} + +fn normalize_audio_rms(samples: &mut [f32], floor: f64) -> Result<(), TransformError> { + if !floor.is_finite() || floor < 0.0 { + return Err(TransformError::ShapeError(format!( + "audio_rms_norm_floor must be finite and non-negative, got {floor}" + ))); + } + if floor == 0.0 { + return Ok(()); + } + + let mean_square = samples + .iter() + .map(|&sample| { + let sample = f64::from(sample); + sample * sample + }) + .sum::() + / samples.len() as f64; + let rms = mean_square.sqrt(); + if rms > 0.0 && rms < floor { + let scale = floor / rms; + for sample in samples { + *sample = (f64::from(*sample) * scale) as f32; + } + } + Ok(()) +} + +fn dmel_bins( + samples: &[f32], + params: &InklingAudioEncoderParams, +) -> Result, TransformError> { + if params.n_mels == 0 || params.num_dmel_bins == 0 { + return Err(TransformError::ShapeError( + "n_mels and num_dmel_bins must be positive".to_string(), + )); + } + + let hop_length = params.hop_length()?; + let window_size = params.window_size()?; + let n_fft = params.n_fft.unwrap_or(window_size); + if n_fft < window_size { + return Err(TransformError::ShapeError(format!( + "n_fft ({n_fft}) must be greater than or equal to window_size ({window_size})" + ))); + } + if samples.is_empty() { + return Err(TransformError::ShapeError( + "audio preprocessing requires at least one sample".to_string(), + )); + } + + let right_pad = samples.len().div_ceil(hop_length) * hop_length - samples.len(); + let left_pad = n_fft.saturating_sub(hop_length); + let padded_len = left_pad + samples.len() + right_pad; + let mut padded = vec![0.0_f32; padded_len]; + padded[left_pad..left_pad + samples.len()].copy_from_slice(samples); + + let frame_count = (padded_len - n_fft) / hop_length + 1; + let fft_bins = n_fft / 2 + 1; + let window = hann_window(window_size); + let mel_basis = mel_basis(params.sample_rate, n_fft, params.n_mels); + let mut planner = FftPlanner::::new(); + let fft = planner.plan_fft_forward(n_fft); + let mut buffer = vec![Complex32::new(0.0, 0.0); n_fft]; + let mut magnitudes = vec![0.0_f32; fft_bins * frame_count]; + + for frame in 0..frame_count { + buffer.fill(Complex32::new(0.0, 0.0)); + let start = frame * hop_length; + for i in 0..window_size { + buffer[i].re = padded[start + i] * window[i]; + } + fft.process(&mut buffer); + for bin in 0..fft_bins { + let value = buffer[bin]; + magnitudes[bin * frame_count + frame] = + (value.re.mul_add(value.re, value.im * value.im)) + .max(1e-10) + .sqrt(); + } + } + + let mut output = Vec::with_capacity(frame_count * params.n_mels); + for frame in 0..frame_count { + for mel in 0..params.n_mels { + let basis_row = &mel_basis[mel * fft_bins..(mel + 1) * fft_bins]; + let mut value = 0.0_f32; + for bin in 0..fft_bins { + value += basis_row[bin] * magnitudes[bin * frame_count + frame]; + } + let log_mel = f64::from(value.max(1e-10).log10()) + .clamp(params.dmel_min_value, params.dmel_max_value); + output.push(quantize_dmel(log_mel, params) as f32); + } + } + + Array2::from_shape_vec((frame_count, params.n_mels), output).map_err(|e| { + TransformError::ShapeError(format!( + "failed to create Inkling dMel bins [{frame_count}, {}]: {e}", + params.n_mels + )) + }) +} + +fn quantize_dmel(value: f64, params: &InklingAudioEncoderParams) -> usize { + if params.num_dmel_bins <= 1 { + return 0; + } + let span = params.dmel_max_value - params.dmel_min_value; + if span <= 0.0 { + return 0; + } + let scaled = (value - params.dmel_min_value) / span * (params.num_dmel_bins - 1) as f64; + // Nearest-bin selection keeps the lower bin on an exact tie. + (scaled - 0.5) + .ceil() + .clamp(0.0, (params.num_dmel_bins - 1) as f64) as usize +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn rejects_empty_or_invalid_rate_audio() { + let processor = InklingAudioProcessor::new(); + assert!(processor + .preprocess_decoded(DecodedAudio { + samples: Vec::new(), + sample_rate: 16_000, + }) + .is_err()); + assert!(processor + .preprocess_decoded(DecodedAudio { + samples: vec![0.0], + sample_rate: 0, + }) + .is_err()); + } + + #[test] + fn model_config_overrides_audio_shape_params() { + let config = serde_json::json!({ + "audio_config": { + "n_mel_bins": 8, + "mel_vocab_size": 4, + "dmel_min_value": -5.0, + "dmel_max_value": 1.0, + "audio_rms_norm_floor": 0.02 + } + }); + let processor = InklingAudioProcessor::from_model_config(&config); + assert_eq!(processor.params.n_mels, 8); + assert_eq!(processor.params.num_dmel_bins, 4); + assert_eq!(processor.params.dmel_min_value, -5.0); + assert_eq!(processor.params.dmel_max_value, 1.0); + assert_eq!(processor.params.audio_rms_norm_floor, 0.02); + } + + fn seeded_signal(scale: f32) -> Vec { + (0..1600) + .map(|i| { + let raw = ((i * 73 + 17 * 977) % 65_536) as i32 - 32_768; + raw as f32 / 32_768.0 * scale + }) + .collect() + } + + #[test] + fn quiet_nonzero_audio_is_normalized_to_rms_floor() { + let processor = InklingAudioProcessor::new(); + let very_quiet = processor + .preprocess_decoded(DecodedAudio { + samples: seeded_signal(0.001), + sample_rate: 16_000, + }) + .unwrap(); + let less_quiet = processor + .preprocess_decoded(DecodedAudio { + samples: seeded_signal(0.004), + sample_rate: 16_000, + }) + .unwrap(); + + assert_eq!(very_quiet, less_quiet); + } + + #[test] + fn rms_floor_leaves_silence_and_loud_audio_unchanged() { + let processor = InklingAudioProcessor::new(); + let mut without_floor = processor.clone(); + without_floor.params.audio_rms_norm_floor = 0.0; + + for samples in [vec![0.0; 1600], seeded_signal(1.0)] { + let with_floor = processor + .preprocess_decoded(DecodedAudio { + samples: samples.clone(), + sample_rate: 16_000, + }) + .unwrap(); + let without_floor = without_floor + .preprocess_decoded(DecodedAudio { + samples, + sample_rate: 16_000, + }) + .unwrap(); + assert_eq!(with_floor, without_floor); + } + } + + #[test] + fn rejects_invalid_rms_floor() { + let mut processor = InklingAudioProcessor::new(); + processor.params.audio_rms_norm_floor = -0.01; + let error = processor + .preprocess_decoded(DecodedAudio { + samples: vec![0.1; 800], + sample_rate: 16_000, + }) + .unwrap_err(); + assert!(error.to_string().contains("audio_rms_norm_floor")); + } +} diff --git a/crates/multimodal/src/audio/processors/mod.rs b/crates/multimodal/src/audio/processors/mod.rs index 77d27a3f64..6a78b81700 100644 --- a/crates/multimodal/src/audio/processors/mod.rs +++ b/crates/multimodal/src/audio/processors/mod.rs @@ -1,5 +1,7 @@ //! Model-specific audio preprocessing implementations. +mod inkling; mod qwen3_audio; +pub use inkling::{InklingAudioEncoderParams, InklingAudioProcessor}; pub use qwen3_audio::{Qwen3AudioParams, Qwen3AudioProcessor}; diff --git a/crates/multimodal/src/registry/inkling.rs b/crates/multimodal/src/registry/inkling.rs new file mode 100644 index 0000000000..717567118a --- /dev/null +++ b/crates/multimodal/src/registry/inkling.rs @@ -0,0 +1,380 @@ +use std::collections::HashMap; + +use serde_json::{json, Value}; + +use crate::{ + audio::{AudioPreProcessor, InklingAudioProcessor}, + encoder_inputs::PreprocessedEncoderInputs, + registry::{ModelMetadata, ModelProcessorSpec, ModelRegistryError, RegistryResult}, + types::{EncoderFieldLayouts, FieldLayout, Modality, PromptReplacement, TokenId}, + vision::PreProcessorConfig, +}; + +pub(super) struct InklingSpec; + +impl InklingSpec { + const CONTENT_IMAGE: &'static str = "<|content_image|>"; + const CONTENT_AUDIO_INPUT: &'static str = "<|content_audio_input|>"; + const AUDIO: &'static str = "<|audio|>"; + + fn image_transport_fill_id(metadata: &ModelMetadata) -> RegistryResult { + metadata.token_id(Self::CONTENT_IMAGE) + } + + fn audio_placeholder_id(metadata: &ModelMetadata) -> RegistryResult { + metadata.token_id(Self::AUDIO) + } + + fn tower_enabled(metadata: &ModelMetadata, config_key: &str) -> bool { + metadata + .config + .get(config_key) + .and_then(|config| config.get("decoder_dmodel")) + .is_some_and(|value| !value.is_null()) + } +} + +impl ModelProcessorSpec for InklingSpec { + fn name(&self) -> &'static str { + "inkling" + } + + fn matches(&self, metadata: &ModelMetadata) -> bool { + let id = metadata.model_id.to_ascii_lowercase(); + id.contains("inkling") + || metadata + .config_model_type() + .is_some_and(|mt| mt == "inkling_mm_model") + } + + fn placeholder_token(&self, _metadata: &ModelMetadata) -> RegistryResult { + Ok(Self::CONTENT_IMAGE.to_string()) + } + + fn placeholder_token_id(&self, metadata: &ModelMetadata) -> RegistryResult { + Self::image_transport_fill_id(metadata) + } + + fn placeholder_token_for( + &self, + _metadata: &ModelMetadata, + modality: Modality, + ) -> RegistryResult { + match modality { + Modality::Image => Ok(Self::CONTENT_IMAGE.to_string()), + Modality::Audio => Ok(Self::CONTENT_AUDIO_INPUT.to_string()), + _ => Err(ModelRegistryError::UnsupportedModality { + spec: self.name(), + modality, + }), + } + } + + fn placeholder_token_id_for( + &self, + metadata: &ModelMetadata, + modality: Modality, + ) -> RegistryResult { + match modality { + Modality::Image => Self::image_transport_fill_id(metadata), + Modality::Audio => Self::audio_placeholder_id(metadata), + _ => Err(ModelRegistryError::UnsupportedModality { + spec: self.name(), + modality, + }), + } + } + + fn modality_limits( + &self, + metadata: &ModelMetadata, + ) -> RegistryResult> { + let mut limits = HashMap::new(); + if Self::tower_enabled(metadata, "vision_config") { + limits.insert(Modality::Image, 10); + } + if Self::tower_enabled(metadata, "audio_config") { + limits.insert(Modality::Audio, 10); + } + Ok(limits) + } + + fn processor_kwargs(&self, _metadata: &ModelMetadata) -> RegistryResult { + Ok(json!({})) + } + + fn audio_processor( + &self, + model_config: &Value, + _preprocessor_config: &PreProcessorConfig, + ) -> Option> { + Some(Box::new(InklingAudioProcessor::from_model_config( + model_config, + ))) + } + + fn prompt_replacements( + &self, + metadata: &ModelMetadata, + preprocessed: &PreprocessedEncoderInputs, + ) -> RegistryResult> { + self.prompt_replacements_for(metadata, preprocessed, Modality::Image) + } + + fn prompt_replacements_for( + &self, + metadata: &ModelMetadata, + preprocessed: &PreprocessedEncoderInputs, + modality: Modality, + ) -> RegistryResult> { + match modality { + Modality::Image => { + let content_image_id = metadata.token_id(Self::CONTENT_IMAGE)?; + Ok(preprocessed + .feature_token_counts + .iter() + .map(|&num_tokens| { + let mut tokens = Vec::with_capacity(num_tokens + 1); + tokens.push(content_image_id); + // TML transports images as a typed span and does not + // define a positive image target token. TokenSpeed only + // needs real ids until it rewrites the explicit feature + // offsets to content-derived MM pad values, so reuse the + // public content token as an internal transport fill. + tokens.extend(std::iter::repeat_n(content_image_id, num_tokens)); + PromptReplacement::sequence(Modality::Image, Self::CONTENT_IMAGE, tokens) + .with_feature_span(1, num_tokens) + }) + .collect()) + } + Modality::Audio => { + let content_audio_id = metadata.token_id(Self::CONTENT_AUDIO_INPUT)?; + let audio_id = Self::audio_placeholder_id(metadata)?; + Ok(preprocessed + .feature_token_counts + .iter() + .map(|&num_tokens| { + let mut tokens = Vec::with_capacity(num_tokens + 1); + tokens.push(content_audio_id); + tokens.extend(std::iter::repeat_n(audio_id, num_tokens)); + PromptReplacement::sequence( + Modality::Audio, + Self::CONTENT_AUDIO_INPUT, + tokens, + ) + .with_feature_span(1, num_tokens) + }) + .collect()) + } + _ => Err(ModelRegistryError::UnsupportedModality { + spec: self.name(), + modality, + }), + } + } + + fn encoder_field_layouts_for(&self, modality: Modality) -> EncoderFieldLayouts { + match modality { + Modality::Image | Modality::Audio => EncoderFieldLayouts::new( + FieldLayout::flat("tokens_per_item"), + HashMap::from([("tokens_per_item".to_string(), FieldLayout::Batched)]), + ), + Modality::Video | Modality::ImageEmbeds => EncoderFieldLayouts::default(), + } + } +} + +#[cfg(test)] +mod tests { + use std::collections::HashMap; + + use serde_json::json; + + use super::InklingSpec; + use crate::{ + registry::{test_helpers::*, ModelMetadata, ModelProcessorSpec, ModelRegistry}, + types::ImageSize, + }; + + fn tokenizer() -> TestTokenizer { + TestTokenizer::new(&[ + ("<|content_image|>", 200005), + ("<|content_audio_input|>", 200020), + ("<|audio|>", 200023), + ]) + } + + fn config() -> serde_json::Value { + json!({ + "model_type": "inkling_mm_model", + "architectures": ["InklingForConditionalGeneration"], + "vision_config": {"decoder_dmodel": 6144}, + "audio_config": {"decoder_dmodel": 6144} + }) + } + + #[test] + fn inkling_matches_model_type() { + let tokenizer = tokenizer(); + let config = config(); + let metadata = ModelMetadata { + model_id: "local-checkpoint", + tokenizer: &tokenizer, + config: &config, + }; + let registry = ModelRegistry::new(); + let spec = registry.lookup(&metadata).expect("inkling spec"); + assert_eq!(spec.name(), "inkling"); + } + + #[test] + fn inkling_matches_family_name_without_model_type() { + let tokenizer = tokenizer(); + let config = json!({}); + let metadata = ModelMetadata { + model_id: "org/inkling-chat", + tokenizer: &tokenizer, + config: &config, + }; + let registry = ModelRegistry::new(); + let spec = registry.lookup(&metadata).expect("inkling spec"); + assert_eq!(spec.name(), "inkling"); + } + + #[test] + fn inkling_spec_builds_audio_processor() { + use std::sync::Arc; + + use bytes::Bytes; + + use crate::{ + audio::DecodedAudio, + types::{AudioClip, AudioSource}, + vision::PreProcessorConfig, + }; + + let config = json!({"audio_config": {"n_mel_bins": 8}}); + let processor = InklingSpec + .audio_processor(&config, &PreProcessorConfig::default()) + .expect("inkling spec must provide an audio processor"); + + let clip = Arc::new(AudioClip::new( + Bytes::from_static(b"audio"), + DecodedAudio { + samples: vec![0.0; 800], + sample_rate: 16_000, + }, + AudioSource::InlineBytes, + "audio-hash".to_string(), + )); + let result = processor.preprocess(&[clip]).unwrap(); + assert_eq!(result.encoder_input.shape(), &[1, 8]); + } + + #[test] + fn modality_limits_follow_configured_towers() { + let tokenizer = tokenizer(); + for (config, expected) in [ + ( + json!({ + "vision_config": {"decoder_dmodel": 6144}, + "audio_config": {"decoder_dmodel": 6144} + }), + HashMap::from([ + (crate::types::Modality::Image, 10), + (crate::types::Modality::Audio, 10), + ]), + ), + ( + json!({ + "vision_config": {"decoder_dmodel": null}, + "audio_config": {"decoder_dmodel": 6144} + }), + HashMap::from([(crate::types::Modality::Audio, 10)]), + ), + ( + json!({ + "vision_config": {"decoder_dmodel": 6144}, + "audio_config": {} + }), + HashMap::from([(crate::types::Modality::Image, 10)]), + ), + (json!({}), HashMap::new()), + ] { + let metadata = ModelMetadata { + model_id: "inkling-test", + tokenizer: &tokenizer, + config: &config, + }; + assert_eq!(InklingSpec.modality_limits(&metadata).unwrap(), expected); + } + } + + #[test] + fn image_replacement_preserves_content_token_and_adds_patch_span() { + let tokenizer = tokenizer(); + let config = config(); + let metadata = ModelMetadata { + model_id: "local-checkpoint", + tokenizer: &tokenizer, + config: &config, + }; + let registry = ModelRegistry::new(); + let spec = registry.lookup(&metadata).expect("inkling spec"); + + let replacements = spec + .prompt_replacements_for( + &metadata, + &test_preprocessed_with_tokens(&[ImageSize::new(80, 40)], &[3]), + crate::types::Modality::Image, + ) + .unwrap(); + assert_eq!(replacements[0].tokens, vec![200005, 200005, 200005, 200005]); + assert_eq!( + replacements[0].feature_ranges, + Some(vec![crate::types::PlaceholderRange { + offset: 1, + length: 3 + }]) + ); + assert_eq!( + spec.placeholder_token_id_for(&metadata, crate::types::Modality::Image) + .unwrap(), + 200005 + ); + } + + #[test] + fn audio_replacement_preserves_content_token_and_adds_placeholder_span() { + let tokenizer = tokenizer(); + let config = config(); + let metadata = ModelMetadata { + model_id: "local-checkpoint", + tokenizer: &tokenizer, + config: &config, + }; + let registry = ModelRegistry::new(); + let spec = registry.lookup(&metadata).expect("inkling spec"); + + let replacements = spec + .prompt_replacements_for( + &metadata, + &test_preprocessed_with_tokens(&[ImageSize::new(80, 2)], &[2]), + crate::types::Modality::Audio, + ) + .unwrap(); + assert_eq!(replacements[0].tokens, vec![200020, 200023, 200023]); + assert_eq!( + replacements[0].feature_ranges, + Some(vec![crate::types::PlaceholderRange { + offset: 1, + length: 2 + }]) + ); + assert_eq!( + spec.placeholder_token_id_for(&metadata, crate::types::Modality::Audio) + .unwrap(), + 200023 + ); + } +} diff --git a/crates/multimodal/src/registry/mod.rs b/crates/multimodal/src/registry/mod.rs index 6fb5a19028..85cf7afffe 100644 --- a/crates/multimodal/src/registry/mod.rs +++ b/crates/multimodal/src/registry/mod.rs @@ -1,3 +1,4 @@ +mod inkling; mod kimi_k25; mod llama4; mod llava; @@ -8,6 +9,7 @@ mod qwen3_vl; mod qwen_vl; mod traits; +use inkling::InklingSpec; use kimi_k25::KimiK25VisionSpec; use llama4::Llama4Spec; use llava::{LlavaNextSpec, LlavaSpec}; @@ -28,6 +30,7 @@ impl ModelRegistry { pub fn new() -> Self { Self { specs: vec![ + LazySpec::new(|| Box::new(InklingSpec)), LazySpec::new(|| Box::new(KimiK25VisionSpec)), LazySpec::new(|| Box::new(Llama4Spec)), // LlavaNext must be registered before Llava so "llava_next" model_type matches first. diff --git a/crates/multimodal/src/types.rs b/crates/multimodal/src/types.rs index fe49170e81..83db848c81 100644 --- a/crates/multimodal/src/types.rs +++ b/crates/multimodal/src/types.rs @@ -461,7 +461,7 @@ impl ImageSize { } } -#[derive(Debug, Clone, Serialize, Deserialize, Default)] +#[derive(Debug, Clone, Serialize, Deserialize, Default, PartialEq, Eq)] pub struct PlaceholderRange { pub offset: usize, pub length: usize, @@ -472,6 +472,13 @@ pub struct PromptReplacement { pub modality: Modality, pub placeholder_token: String, pub tokens: Vec, + /// Feature-token ranges relative to the start of `tokens`. + /// + /// Most model specs leave this unset and let the gateway discover feature + /// ranges by scanning for the configured placeholder token id. Specs whose + /// transport fill token is also a structural token can set explicit ranges + /// so the structural header is not mistaken for an encoder position. + pub feature_ranges: Option>, /// Number of structural tokens the chat template emits *immediately before* /// this placeholder (e.g. Qwen's leading `<|vision_start|>`) that belong to /// the placeholder's range. `expand_tokens` folds them into the reported @@ -494,6 +501,7 @@ impl PromptReplacement { modality, placeholder_token: placeholder_token.to_string(), tokens: vec![token_id; count], + feature_ranges: None, structural_prefix: 0, } } @@ -503,10 +511,25 @@ impl PromptReplacement { modality, placeholder_token: placeholder_token.to_string(), tokens: sequence, + feature_ranges: None, structural_prefix: 0, } } + /// Declare the encoder-feature ranges inside the replacement sequence. + /// Offsets are relative to the first replacement token. + #[must_use] + pub fn with_feature_ranges(mut self, ranges: Vec) -> Self { + self.feature_ranges = Some(ranges); + self + } + + /// Declare one contiguous encoder-feature span inside the replacement. + #[must_use] + pub fn with_feature_span(self, offset: usize, length: usize) -> Self { + self.with_feature_ranges(vec![PlaceholderRange { offset, length }]) + } + /// Declare that `n` template-emitted structural tokens precede this /// placeholder and should be included in its reported range. See /// [`Self::structural_prefix`]. @@ -535,6 +558,16 @@ mod tests { fn prompt_replacement_builders() { let rep = PromptReplacement::repeated(Modality::Image, "", 100, 3); assert_eq!(rep.tokens, vec![100, 100, 100]); + assert!(rep.feature_ranges.is_none()); + + let rep = rep.with_feature_span(1, 2); + assert_eq!( + rep.feature_ranges, + Some(vec![PlaceholderRange { + offset: 1, + length: 2 + }]) + ); } #[test] diff --git a/crates/multimodal/src/vision/processor.rs b/crates/multimodal/src/vision/processor.rs index de025dddb2..27654ae877 100644 --- a/crates/multimodal/src/vision/processor.rs +++ b/crates/multimodal/src/vision/processor.rs @@ -272,6 +272,11 @@ impl VisionProcessorRegistry { Box::new(super::processors::Qwen3VLProcessor::new()), ); + // Inkling/TML multimodal model. + registry.register( + "inkling", + Box::new(super::processors::InklingImageProcessor::new()), + ); // Register Qwen2-VL (matches Qwen/Qwen2-VL-*, etc.) registry.register( "qwen2-vl", @@ -374,6 +379,11 @@ mod tests { // Get the processor and check model name let processor = registry.find("llava-hf/llava-1.5-7b-hf", None).unwrap(); assert_eq!(processor.model_name(), "llava"); + + let processor = registry + .find("org/inkling-chat", None) + .expect("Inkling model family"); + assert_eq!(processor.model_name(), "inkling"); } #[test] diff --git a/crates/multimodal/src/vision/processors/inkling.rs b/crates/multimodal/src/vision/processors/inkling.rs new file mode 100644 index 0000000000..31d3393b27 --- /dev/null +++ b/crates/multimodal/src/vision/processors/inkling.rs @@ -0,0 +1,456 @@ +//! Inkling image processor. +//! +//! Implements `InklingImageProcessor`: optional aspect-preserving +//! resize, RGB HWC input, CLIP normalization, and one extra patch column on +//! every row. The output tensor is +//! `[sum_patches, temporal_patch_size, patch_size, patch_size, 3]`. + +use image::{DynamicImage, GenericImageView, RgbImage}; +use ndarray::{ArrayD, IxDyn}; + +use crate::vision::{ + preprocessor_config::PreProcessorConfig, + processor::{ModelSpecificValue, PreprocessedEncoderInputs, VisionPreProcessor}, + transforms::{self, TransformError}, +}; + +pub const INKLING_IMAGE_MEAN: [f64; 3] = [0.48145466, 0.4578275, 0.40821073]; +pub const INKLING_IMAGE_STD: [f64; 3] = [0.26862954, 0.2613026, 0.2757771]; +pub const DEFAULT_PATCH_SIZE: usize = 40; +pub const DEFAULT_TEMPORAL_PATCH_SIZE: usize = 2; +pub const DEFAULT_RESCALE_IMAGE_MAX_UPSCALED_LONG_EDGE: u32 = 2048; +const PAD_RAW_VALUE: f32 = -1.0 / 255.0; + +#[derive(Debug, Clone)] +pub struct InklingImageProcessor { + patch_size: usize, + temporal_patch_size: usize, + rescale_image_frac: Option, + rescale_image_max_upscaled_long_edge: Option, +} + +impl Default for InklingImageProcessor { + fn default() -> Self { + Self::new() + } +} + +impl InklingImageProcessor { + pub fn new() -> Self { + Self { + patch_size: DEFAULT_PATCH_SIZE, + temporal_patch_size: DEFAULT_TEMPORAL_PATCH_SIZE, + // Count patches at the authored image dimensions. Resizing remains + // an explicit preprocessor override. + rescale_image_frac: None, + rescale_image_max_upscaled_long_edge: Some( + DEFAULT_RESCALE_IMAGE_MAX_UPSCALED_LONG_EDGE, + ), + } + } + + fn with_preprocessor_config(&self, config: &PreProcessorConfig) -> Self { + let mut processor = self.clone(); + if config.patch_size.is_some() { + processor.patch_size = config.get_patch_size(self.patch_size); + } + if let Some(temporal_patch_size) = config.temporal_patch_size { + processor.temporal_patch_size = temporal_patch_size; + } + if let Some(value) = config.extra.get("rescale_image_frac") { + processor.rescale_image_frac = value.as_f64(); + } + if let Some(value) = config.extra.get("rescale_image_max_upscaled_long_edge") { + processor.rescale_image_max_upscaled_long_edge = + value.as_u64().and_then(|v| u32::try_from(v).ok()); + } + processor + } + + fn patch_grid(&self, width: usize, height: usize) -> (usize, usize, usize) { + let nph = height.div_ceil(self.patch_size); + let npw = width / self.patch_size + 1; + let num_patches = nph * npw; + (nph, npw, num_patches) + } + + fn num_patch_features(&self) -> usize { + self.temporal_patch_size * self.patch_size * self.patch_size * 3 + } + + fn validate(&self) -> Result<(), TransformError> { + if self.patch_size == 0 || self.temporal_patch_size == 0 { + return Err(TransformError::ShapeError( + "Inkling patch_size and temporal_patch_size must be positive".to_string(), + )); + } + if let Some(frac) = self.rescale_image_frac { + if !frac.is_finite() || frac <= 0.0 { + return Err(TransformError::ShapeError(format!( + "Inkling rescale_image_frac must be positive and finite, got {frac}" + ))); + } + } + if matches!(self.rescale_image_max_upscaled_long_edge, Some(0)) { + return Err(TransformError::ShapeError( + "Inkling rescale_image_max_upscaled_long_edge must be positive".to_string(), + )); + } + Ok(()) + } + + fn scaled_image_dimensions(&self, width: u32, height: u32) -> (u32, u32) { + let Some(frac) = self.rescale_image_frac else { + return (width, height); + }; + let long_edge = width.max(height); + if long_edge == 0 { + return (width, height); + } + + let mut target_long_edge = long_edge as f64 * frac; + if let Some(max_upscaled_long_edge) = self.rescale_image_max_upscaled_long_edge { + let effective_cap = max_upscaled_long_edge.max(long_edge); + target_long_edge = target_long_edge.min(effective_cap as f64); + } + let ratio = target_long_edge / long_edge as f64; + if (ratio - 1.0).abs() < f64::EPSILON { + return (width, height); + } + + let scale_dim = |dim: u32| -> u32 { ((dim as f64 * ratio + 0.5).floor() as u32).max(1) }; + (scale_dim(width), scale_dim(height)) + } + + fn prepare_rgb_image(&self, image: &DynamicImage) -> RgbImage { + let (width, height) = image.dimensions(); + let (scaled_width, scaled_height) = self.scaled_image_dimensions(width, height); + if (scaled_width, scaled_height) == (width, height) { + image.to_rgb8() + } else { + transforms::resize_lanczos_pil(image, scaled_width, scaled_height).to_rgb8() + } + } + + fn append_image_patches( + &self, + image: &DynamicImage, + mean: &[f64; 3], + std: &[f64; 3], + output: &mut Vec, + ) { + let rgb = self.prepare_rgb_image(image); + let width = rgb.width() as usize; + let height = rgb.height() as usize; + let raw = rgb.as_raw(); + let (_, npw, _) = self.patch_grid(width, height); + let pad_norm: [f32; 3] = + std::array::from_fn(|c| ((PAD_RAW_VALUE as f64 - mean[c]) / std[c]) as f32); + let inv255 = 1.0_f32 / 255.0; + + for patch_idx in 0..height.div_ceil(self.patch_size) * npw { + let patch_y = patch_idx / npw; + let patch_x = patch_idx - patch_y * npw; + let y_base = patch_y * self.patch_size; + let x_base = patch_x * self.patch_size; + let patch_start = output.len(); + + for y in 0..self.patch_size { + let iy = y_base + y; + for x in 0..self.patch_size { + let ix = x_base + x; + if iy < height && ix < width { + let offset = (iy * width + ix) * 3; + for c in 0..3 { + let raw_value = raw[offset + c] as f32 * inv255; + output.push(((raw_value as f64 - mean[c]) / std[c]) as f32); + } + } else { + output.extend_from_slice(&pad_norm); + } + } + } + + let single_temporal_patch_len = self.patch_size * self.patch_size * 3; + for _ in 1..self.temporal_patch_size { + output.extend_from_within(patch_start..patch_start + single_temporal_patch_len); + } + } + } +} + +impl VisionPreProcessor for InklingImageProcessor { + fn default_mean(&self) -> [f64; 3] { + INKLING_IMAGE_MEAN + } + + fn default_std(&self) -> [f64; 3] { + INKLING_IMAGE_STD + } + + fn preprocess( + &self, + images: &[DynamicImage], + config: &PreProcessorConfig, + ) -> Result { + if images.is_empty() { + return Err(TransformError::EmptyBatch); + } + + let processor = self.with_preprocessor_config(config); + processor.validate()?; + + let mean = config + .image_mean + .as_ref() + .map_or(INKLING_IMAGE_MEAN, |values| { + if values.len() >= 3 { + [values[0], values[1], values[2]] + } else { + INKLING_IMAGE_MEAN + } + }); + let std = config + .image_std + .as_ref() + .map_or(INKLING_IMAGE_STD, |values| { + if values.len() >= 3 { + [values[0], values[1], values[2]] + } else { + INKLING_IMAGE_STD + } + }); + + let mut feature_token_counts = Vec::with_capacity(images.len()); + let mut tokens_per_item = Vec::with_capacity(images.len()); + let mut item_sizes = Vec::with_capacity(images.len()); + let mut total_patches = 0usize; + for image in images { + let (width, height) = image.dimensions(); + let (scaled_width, scaled_height) = processor.scaled_image_dimensions(width, height); + let (_, _, num_patches) = + processor.patch_grid(scaled_width as usize, scaled_height as usize); + feature_token_counts.push(num_patches); + tokens_per_item.push(num_patches as i64); + item_sizes.push((scaled_width, scaled_height)); + total_patches += num_patches; + } + + let mut patches = Vec::with_capacity(total_patches * processor.num_patch_features()); + for image in images { + processor.append_image_patches(image, &mean, &std, &mut patches); + } + + let encoder_input = ArrayD::from_shape_vec( + IxDyn(&[ + total_patches, + processor.temporal_patch_size, + processor.patch_size, + processor.patch_size, + 3, + ]), + patches, + ) + .map_err(|e| { + TransformError::ShapeError(format!( + "failed to create Inkling image encoder input [{total_patches}, {}, {}, {}, 3]: {e}", + processor.temporal_patch_size, processor.patch_size, processor.patch_size + )) + })?; + + Ok( + PreprocessedEncoderInputs::new(encoder_input, feature_token_counts, item_sizes) + .with_extra( + "tokens_per_item", + ModelSpecificValue::int_1d(tokens_per_item), + ), + ) + } + + fn calculate_num_tokens(&self, width: u32, height: u32, config: &PreProcessorConfig) -> usize { + let processor = self.with_preprocessor_config(config); + let (scaled_width, scaled_height) = processor.scaled_image_dimensions(width, height); + processor + .patch_grid(scaled_width as usize, scaled_height as usize) + .2 + } + + fn model_name(&self) -> &'static str { + "inkling" + } + + fn get_processed_size(&self, _config: &PreProcessorConfig) -> Option<(u32, u32)> { + None + } +} + +#[cfg(test)] +mod tests { + use image::{Rgb, RgbImage}; + + use super::*; + use crate::vision::{preprocessor_config::PatchSize, processor::ModelSpecificValue}; + + fn test_image() -> DynamicImage { + let mut img = RgbImage::new(2, 2); + img.put_pixel(0, 0, Rgb([0, 127, 255])); + img.put_pixel(1, 0, Rgb([255, 127, 0])); + img.put_pixel(0, 1, Rgb([10, 20, 30])); + img.put_pixel(1, 1, Rgb([40, 50, 60])); + DynamicImage::ImageRgb8(img) + } + + fn processor_without_rescale() -> InklingImageProcessor { + InklingImageProcessor { + patch_size: DEFAULT_PATCH_SIZE, + temporal_patch_size: DEFAULT_TEMPORAL_PATCH_SIZE, + rescale_image_frac: None, + rescale_image_max_upscaled_long_edge: None, + } + } + + fn norm(raw: f32, channel: usize) -> f32 { + ((raw / 255.0) as f64 - INKLING_IMAGE_MEAN[channel]) as f32 + / INKLING_IMAGE_STD[channel] as f32 + } + + fn pad_norm(channel: usize) -> f32 { + (PAD_RAW_VALUE as f64 - INKLING_IMAGE_MEAN[channel]) as f32 + / INKLING_IMAGE_STD[channel] as f32 + } + + #[test] + fn inkling_image_shape_includes_extra_patch_column_and_temporal_duplication() { + let processor = processor_without_rescale(); + let config = PreProcessorConfig { + patch_size: Some(PatchSize { + height: Some(2), + width: Some(2), + }), + temporal_patch_size: Some(2), + ..PreProcessorConfig::default() + }; + let result = processor.preprocess(&[test_image()], &config).unwrap(); + assert_eq!(result.encoder_input.shape(), &[2, 2, 2, 2, 3]); + assert_eq!(result.feature_token_counts, vec![2]); + + let flat = result.encoder_input.as_slice().unwrap(); + assert!((flat[0] - norm(0.0, 0)).abs() < 1e-6); + assert!((flat[1] - norm(127.0, 1)).abs() < 1e-6); + assert!((flat[2] - norm(255.0, 2)).abs() < 1e-6); + + let single_temporal_len = 2 * 2 * 3; + assert_eq!( + &flat[0..single_temporal_len], + &flat[single_temporal_len..2 * single_temporal_len] + ); + + let second_patch_start = 2 * single_temporal_len; + assert!((flat[second_patch_start] - pad_norm(0)).abs() < 1e-6); + assert!((flat[second_patch_start + 1] - pad_norm(1)).abs() < 1e-6); + assert!((flat[second_patch_start + 2] - pad_norm(2)).abs() < 1e-6); + } + + #[test] + fn calculate_num_tokens_matches_expected_grid_rule() { + let processor = processor_without_rescale(); + let config = PreProcessorConfig { + patch_size: Some(PatchSize { + height: Some(40), + width: Some(40), + }), + ..PreProcessorConfig::default() + }; + assert_eq!(processor.calculate_num_tokens(80, 40, &config), 3); + assert_eq!(processor.calculate_num_tokens(81, 41, &config), 6); + } + + #[test] + fn calculate_num_tokens_uses_authored_dimensions_by_default() { + let processor = InklingImageProcessor::new(); + let config = PreProcessorConfig::default(); + + assert_eq!(processor.calculate_num_tokens(640, 480, &config), 204); + assert_eq!(processor.calculate_num_tokens(1200, 600, &config), 465); + assert_eq!(processor.calculate_num_tokens(3000, 2000, &config), 3800); + } + + #[test] + fn config_can_explicitly_enable_image_rescale() { + let processor = InklingImageProcessor::new(); + let config = PreProcessorConfig::from_json( + r#"{"patch_size":40,"rescale_image_frac":2.0,"rescale_image_max_upscaled_long_edge":2048}"#, + ) + .unwrap(); + + assert_eq!(processor.calculate_num_tokens(640, 480, &config), 792); + } + + #[test] + fn preprocess_preserves_authored_size_by_default() { + let processor = InklingImageProcessor::new(); + let config = PreProcessorConfig { + patch_size: Some(PatchSize { + height: Some(2), + width: Some(2), + }), + temporal_patch_size: Some(2), + ..PreProcessorConfig::default() + }; + + let result = processor.preprocess(&[test_image()], &config).unwrap(); + + assert_eq!(result.encoder_input.shape(), &[2, 2, 2, 2, 3]); + assert_eq!(result.feature_token_counts, vec![2]); + assert_eq!(result.item_sizes, vec![(2, 2)]); + assert!(matches!( + result.model_specific.get("tokens_per_item"), + Some(ModelSpecificValue::IntTensor { data, shape }) + if data.as_slice() == [2] && shape.as_slice() == [1] + )); + } + + #[test] + fn preprocess_applies_explicit_resize_before_patchifying() { + let processor = InklingImageProcessor::new(); + let config = PreProcessorConfig::from_json( + r#"{"patch_size":2,"temporal_patch_size":2,"rescale_image_frac":2.0}"#, + ) + .unwrap(); + + let result = processor.preprocess(&[test_image()], &config).unwrap(); + + assert_eq!(result.encoder_input.shape(), &[6, 2, 2, 2, 3]); + assert_eq!(result.feature_token_counts, vec![6]); + assert_eq!(result.item_sizes, vec![(4, 4)]); + } + + #[test] + fn scaled_dimensions_match_expected_rounding_and_cap_rule() { + let processor = InklingImageProcessor { + patch_size: DEFAULT_PATCH_SIZE, + temporal_patch_size: DEFAULT_TEMPORAL_PATCH_SIZE, + rescale_image_frac: Some(1.5), + rescale_image_max_upscaled_long_edge: None, + }; + + assert_eq!(processor.scaled_image_dimensions(3, 2), (5, 3)); + + let capped_processor = InklingImageProcessor { + patch_size: DEFAULT_PATCH_SIZE, + temporal_patch_size: DEFAULT_TEMPORAL_PATCH_SIZE, + rescale_image_frac: Some(2.0), + rescale_image_max_upscaled_long_edge: Some( + DEFAULT_RESCALE_IMAGE_MAX_UPSCALED_LONG_EDGE, + ), + }; + assert_eq!( + capped_processor.scaled_image_dimensions(1200, 600), + (2048, 1024) + ); + assert_eq!( + capped_processor.scaled_image_dimensions(3000, 2000), + (3000, 2000) + ); + } +} diff --git a/crates/multimodal/src/vision/processors/mod.rs b/crates/multimodal/src/vision/processors/mod.rs index 4c8c4135e8..5ff53bbb37 100644 --- a/crates/multimodal/src/vision/processors/mod.rs +++ b/crates/multimodal/src/vision/processors/mod.rs @@ -11,12 +11,14 @@ //! - **Qwen2.5-VL** (`qwen2_vl`): Same processor as Qwen2-VL (identical preprocessing) //! - **Qwen3-VL** (`qwen3_vl`): Similar to Qwen2-VL but with patch_size=16 and [0.5,0.5,0.5] normalization //! - **Qwen3-Omni** (`qwen3_omni_vision`): Qwen3 vision preprocessing with Omni video limits and timing metadata +//! - **Inkling** (`inkling`): Optional aspect-preserving resize with CLIP normalization and padded patch columns //! - **Kimi-K2.5** (`kimi_k25`): MoonViT resize and zero-padding to patch alignment //! - **Phi3-Vision** (`phi3_vision`): Dynamic HD transform with 336x336 tiles //! - **Phi4-Vision** (`phi4_vision`): Dynamic HD transform with 448x448 tiles and SiGLIP encoder //! - **LLaMA 4 Vision** (`llama4_vision`): Tile-based processing with 336x336 tiles and global tile //! - **Pixtral/Mistral3** (`pixtral`): CLIP-based preprocessing with dynamic resolution +pub mod inkling; pub mod kimi_k25; pub mod llama4_vision; pub mod llava; @@ -28,6 +30,7 @@ pub mod qwen3_omni_vision; pub mod qwen3_vl; pub mod qwen_vl_base; +pub use inkling::InklingImageProcessor; pub use kimi_k25::KimiK25Processor; pub use llama4_vision::Llama4VisionProcessor; pub use llava::{ImageAspectRatio, LlavaNextProcessor, LlavaProcessor}; diff --git a/crates/multimodal/src/vision/transforms.rs b/crates/multimodal/src/vision/transforms.rs index ee02174972..10c89ac5ed 100644 --- a/crates/multimodal/src/vision/transforms.rs +++ b/crates/multimodal/src/vision/transforms.rs @@ -3,7 +3,7 @@ //! This module provides composable transforms that match HuggingFace image processor //! behavior, enabling pure Rust preprocessing without Python dependencies. -use std::cell::RefCell; +use std::{cell::RefCell, f64::consts::PI}; use fast_image_resize::{ images::{Image as FirImage, ImageRef as FirImageRef}, @@ -318,6 +318,29 @@ fn fir_image_to_dynamic( // algorithm exactly, validated against Pillow. const PIL_PRECISION_BITS: i64 = 32 - 8 - 2; const PIL_BICUBIC_SUPPORT: f64 = 2.0; +const PIL_LANCZOS_SUPPORT: f64 = 3.0; + +#[derive(Clone, Copy)] +enum PilResizeFilter { + Bicubic, + Lanczos, +} + +impl PilResizeFilter { + fn support(self) -> f64 { + match self { + Self::Bicubic => PIL_BICUBIC_SUPPORT, + Self::Lanczos => PIL_LANCZOS_SUPPORT, + } + } + + fn weight(self, x: f64) -> f64 { + match self { + Self::Bicubic => pil_cubic(x), + Self::Lanczos => pil_lanczos(x), + } + } +} #[inline] fn pil_cubic(x: f64) -> f64 { @@ -333,12 +356,36 @@ fn pil_cubic(x: f64) -> f64 { } } +#[inline] +fn pil_sinc(x: f64) -> f64 { + if x == 0.0 { + 1.0 + } else { + let x = x * PI; + x.sin() / x + } +} + +#[inline] +fn pil_lanczos(x: f64) -> f64 { + let x = x.abs(); + if x < PIL_LANCZOS_SUPPORT { + pil_sinc(x) * pil_sinc(x / PIL_LANCZOS_SUPPORT) + } else { + 0.0 + } +} + /// Pillow `precompute_coeffs` for one axis: integer (fixed-point) kernels plus /// per-output bounds `(start, count)`. -fn pil_precompute_coeffs(in_size: usize, out_size: usize) -> (Vec<(usize, usize)>, Vec>) { +fn pil_precompute_coeffs( + in_size: usize, + out_size: usize, + filter: PilResizeFilter, +) -> (Vec<(usize, usize)>, Vec>) { let scale = in_size as f64 / out_size as f64; let filterscale = if scale >= 1.0 { scale } else { 1.0 }; - let support = PIL_BICUBIC_SUPPORT * filterscale; + let support = filter.support() * filterscale; let inv = 1.0 / filterscale; let coeff_scale = (1_i64 << PIL_PRECISION_BITS) as f64; @@ -360,7 +407,7 @@ fn pil_precompute_coeffs(in_size: usize, out_size: usize) -> (Vec<(usize, usize) let mut w = vec![0.0_f64; xmax]; let mut tot = 0.0; for (x, wx) in w.iter_mut().enumerate() { - let v = pil_cubic(((x + xmin) as f64 - center + 0.5) * inv); + let v = filter.weight(((x + xmin) as f64 - center + 0.5) * inv); *wx = v; tot += v; } @@ -489,8 +536,9 @@ fn pil_resample_horizontal( in_w: usize, out_w: usize, channels: usize, + filter: PilResizeFilter, ) -> Vec { - let (bounds, kernels) = pil_precompute_coeffs(in_w, out_w); + let (bounds, kernels) = pil_precompute_coeffs(in_w, out_w, filter); let half = 1_i64 << (PIL_PRECISION_BITS - 1); let row_out = out_w * channels; let mut out = vec![0_u8; rows * row_out]; @@ -520,8 +568,14 @@ fn pil_resample_horizontal( out } -fn pil_resample_horizontal_rgb(src: &[u8], rows: usize, in_w: usize, out_w: usize) -> Vec { - let (bounds, kernels) = pil_precompute_coeffs(in_w, out_w); +fn pil_resample_horizontal_rgb( + src: &[u8], + rows: usize, + in_w: usize, + out_w: usize, + filter: PilResizeFilter, +) -> Vec { + let (bounds, kernels) = pil_precompute_coeffs(in_w, out_w, filter); let half = 1_i64 << (PIL_PRECISION_BITS - 1); let row_out = out_w * 3; let mut out = vec![0_u8; rows * row_out]; @@ -641,8 +695,9 @@ fn pil_resample_vertical( width: usize, out_h: usize, channels: usize, + filter: PilResizeFilter, ) -> Vec { - let (bounds, kernels) = pil_precompute_coeffs(in_h, out_h); + let (bounds, kernels) = pil_precompute_coeffs(in_h, out_h, filter); let half = 1_i64 << (PIL_PRECISION_BITS - 1); let row_out = width * channels; let mut out = vec![0_u8; out_h * row_out]; @@ -668,8 +723,14 @@ fn pil_resample_vertical( out } -fn pil_resample_vertical_rgb(src: &[u8], in_h: usize, width: usize, out_h: usize) -> Vec { - let (bounds, kernels) = pil_precompute_coeffs(in_h, out_h); +fn pil_resample_vertical_rgb( + src: &[u8], + in_h: usize, + width: usize, + out_h: usize, + filter: PilResizeFilter, +) -> Vec { + let (bounds, kernels) = pil_precompute_coeffs(in_h, out_h, filter); let half = 1_i64 << (PIL_PRECISION_BITS - 1); let row_out = width * 3; let mut out = vec![0_u8; out_h * row_out]; @@ -702,7 +763,15 @@ fn pil_resample_vertical_rgb(src: &[u8], in_h: usize, width: usize, out_h: usize pub fn resize_bicubic_pil(image: &DynamicImage, out_w: u32, out_h: u32) -> DynamicImage { let rgb = image.to_rgb8(); let (in_w, in_h) = rgb.dimensions(); - let output = resize_bicubic_pil_bytes(rgb.as_raw(), in_w, in_h, out_w, out_h, false); + let output = resize_pil_bytes( + rgb.as_raw(), + in_w, + in_h, + out_w, + out_h, + false, + PilResizeFilter::Bicubic, + ); #[expect( clippy::expect_used, reason = "output is exactly out_w*out_h*3 bytes by construction" @@ -732,7 +801,15 @@ pub fn resize_bicubic_pil_rgb( data.len() ))); } - let output = resize_bicubic_pil_bytes(data, width, height, out_w, out_h, true); + let output = resize_pil_bytes( + data, + width, + height, + out_w, + out_h, + true, + PilResizeFilter::Bicubic, + ); RgbImage::from_raw(out_w, out_h, output).ok_or_else(|| { TransformError::ShapeError(format!( "failed to build PIL bicubic RGB image for {out_w}x{out_h}" @@ -740,35 +817,59 @@ pub fn resize_bicubic_pil_rgb( }) } -fn resize_bicubic_pil_bytes( +/// Pillow-exact LANCZOS resize (RGB8), matching +/// `PIL.Image.resize(.., LANCZOS)`. +pub fn resize_lanczos_pil(image: &DynamicImage, out_w: u32, out_h: u32) -> DynamicImage { + let rgb = image.to_rgb8(); + let (in_w, in_h) = rgb.dimensions(); + let output = resize_pil_bytes( + rgb.as_raw(), + in_w, + in_h, + out_w, + out_h, + false, + PilResizeFilter::Lanczos, + ); + #[expect( + clippy::expect_used, + reason = "output is exactly out_w*out_h*3 bytes by construction" + )] + DynamicImage::ImageRgb8( + RgbImage::from_raw(out_w, out_h, output).expect("pil resize buffer size"), + ) +} + +fn resize_pil_bytes( data: &[u8], in_w: u32, in_h: u32, out_w: u32, out_h: u32, joint_rgb: bool, + filter: PilResizeFilter, ) -> Vec { let (in_w, in_h, out_w, out_h) = (in_w as usize, in_h as usize, out_w as usize, out_h as usize); if in_w == out_w && in_h == out_h { data.to_vec() } else if in_w == out_w { if joint_rgb { - pil_resample_vertical_rgb(data, in_h, in_w, out_h) + pil_resample_vertical_rgb(data, in_h, in_w, out_h, filter) } else { - pil_resample_vertical(data, in_h, in_w, out_h, 3) + pil_resample_vertical(data, in_h, in_w, out_h, 3, filter) } } else { let horiz = if joint_rgb { - pil_resample_horizontal_rgb(data, in_h, in_w, out_w) + pil_resample_horizontal_rgb(data, in_h, in_w, out_w, filter) } else { - pil_resample_horizontal(data, in_h, in_w, out_w, 3) + pil_resample_horizontal(data, in_h, in_w, out_w, 3, filter) }; if in_h == out_h { horiz } else if joint_rgb { - pil_resample_vertical_rgb(&horiz, in_h, out_w, out_h) + pil_resample_vertical_rgb(&horiz, in_h, out_w, out_h, filter) } else { - pil_resample_vertical(&horiz, in_h, out_w, out_h, 3) + pil_resample_vertical(&horiz, in_h, out_w, out_h, 3, filter) } } } @@ -1059,14 +1160,21 @@ mod tests { } for (out_w, out_h) in [(src_w, 17), (19, src_h), (src_w, src_h)] { - let horizontal = - pil_resample_horizontal(&data, src_h as usize, src_w as usize, out_w as usize, 3); + let horizontal = pil_resample_horizontal( + &data, + src_h as usize, + src_w as usize, + out_w as usize, + 3, + PilResizeFilter::Bicubic, + ); let expected = pil_resample_vertical( &horizontal, src_h as usize, out_w as usize, out_h as usize, 3, + PilResizeFilter::Bicubic, ); let actual = resize_bicubic_pil_rgb(&data, src_w, src_h, out_w, out_h) .unwrap() diff --git a/crates/multimodal/tests/fixtures/golden/inkling_preprocess_fingerprints.json b/crates/multimodal/tests/fixtures/golden/inkling_preprocess_fingerprints.json new file mode 100644 index 0000000000..9ef4753537 --- /dev/null +++ b/crates/multimodal/tests/fixtures/golden/inkling_preprocess_fingerprints.json @@ -0,0 +1,73 @@ +{ + "audio_cases": [ + { + "feature_token_counts": [ + 1 + ], + "fnv1a_i32": "23bf6f06afce8925", + "name": "silence_50ms", + "sample_rate": 16000, + "samples": 800, + "seed": 0, + "shape": [ + 1, + 80 + ], + "tokens_per_item": [ + 1 + ] + }, + { + "feature_token_counts": [ + 2 + ], + "fnv1a_i32": "28c53d4365adb6d6", + "name": "seeded_100ms", + "sample_rate": 16000, + "samples": 1600, + "seed": 17, + "shape": [ + 2, + 80 + ], + "tokens_per_item": [ + 2 + ] + }, + { + "feature_token_counts": [ + 2 + ], + "fnv1a_i32": "9fe4630677ca10a0", + "name": "seeded_100ms_44100hz", + "sample_rate": 44100, + "samples": 4410, + "seed": 23, + "shape": [ + 2, + 80 + ], + "tokens_per_item": [ + 2 + ] + }, + { + "amplitude": 0.001, + "feature_token_counts": [ + 2 + ], + "fnv1a_i32": "4d98d5274acbd938", + "name": "quiet_seeded_100ms", + "sample_rate": 16000, + "samples": 1600, + "seed": 17, + "shape": [ + 2, + 80 + ], + "tokens_per_item": [ + 2 + ] + } + ] +} diff --git a/crates/multimodal/tests/inkling_preprocess_golden.rs b/crates/multimodal/tests/inkling_preprocess_golden.rs new file mode 100644 index 0000000000..796daeb1b9 --- /dev/null +++ b/crates/multimodal/tests/inkling_preprocess_golden.rs @@ -0,0 +1,126 @@ +//! Inkling audio preprocessing regression checks. + +#![allow(clippy::expect_used, clippy::panic)] + +use llm_multimodal::{ + audio::{DecodedAudio, InklingAudioProcessor}, + vision::processor::{ModelSpecificValue, PreprocessedEncoderInputs}, +}; +use serde::Deserialize; + +#[derive(Deserialize)] +struct GoldenDocument { + audio_cases: Vec, +} + +#[derive(Deserialize)] +struct AudioCase { + name: String, + sample_rate: usize, + samples: usize, + seed: u32, + #[serde(default = "default_audio_amplitude")] + amplitude: f32, + shape: Vec, + tokens_per_item: Vec, + feature_token_counts: Vec, + fnv1a_i32: String, +} + +fn default_audio_amplitude() -> f32 { + 1.0 +} + +fn make_audio_samples(count: usize, seed: u32, amplitude: f32) -> Vec { + if seed == 0 { + return vec![0.0; count]; + } + (0..count) + .map(|i| { + let raw = ((i as u32 * 73 + seed * 977) % 65_536) as i32 - 32_768; + raw as f32 / 32_768.0 * amplitude + }) + .collect() +} + +fn tokens_per_item(result: &PreprocessedEncoderInputs) -> Vec { + match result.model_specific.get("tokens_per_item") { + Some(ModelSpecificValue::IntTensor { data, shape }) => { + assert_eq!(shape, &[data.len()]); + data.clone() + } + value => panic!("expected tokens_per_item IntTensor, got {value:?}"), + } +} + +fn fnv1a_update(mut hash: u64, bytes: &[u8]) -> u64 { + for byte in bytes { + hash ^= u64::from(*byte); + hash = hash.wrapping_mul(0x0000_0100_0000_01b3); + } + hash +} + +fn fnv1a_i32(values: &[f32]) -> String { + let mut hash = 0xcbf2_9ce4_8422_2325_u64; + for value in values { + let bin = *value as i32; + assert!( + (*value - bin as f32).abs() < f32::EPSILON, + "audio dMel bin is not an integer: {value}" + ); + hash = fnv1a_update(hash, &bin.to_le_bytes()); + } + format!("{hash:016x}") +} + +fn load_golden() -> GoldenDocument { + serde_json::from_str(include_str!( + "fixtures/golden/inkling_preprocess_fingerprints.json" + )) + .expect("invalid checked-in Inkling golden fixture") +} + +#[test] +fn inkling_audio_preprocess_matches_checked_in_golden() { + let golden = load_golden(); + let processor = InklingAudioProcessor::new(); + for case in &golden.audio_cases { + let decoded = DecodedAudio { + samples: make_audio_samples(case.samples, case.seed, case.amplitude), + sample_rate: case.sample_rate, + }; + let result = processor + .preprocess_decoded_clips(vec![decoded]) + .expect("Inkling audio preprocessing failed"); + + assert_eq!( + result.encoder_input.shape(), + case.shape, + "audio shape changed for {}", + case.name + ); + assert_eq!( + result.feature_token_counts, case.feature_token_counts, + "audio token counts changed for {}", + case.name + ); + assert_eq!( + tokens_per_item(&result), + case.tokens_per_item, + "audio tokens_per_item changed for {}", + case.name + ); + + let values = result + .encoder_input + .as_slice_memory_order() + .expect("Inkling audio encoder input must be contiguous"); + assert_eq!( + fnv1a_i32(values), + case.fnv1a_i32, + "Inkling audio int32 fingerprint changed for {}", + case.name + ); + } +} diff --git a/crates/protocols/src/chat.rs b/crates/protocols/src/chat.rs index 7c153c1cfe..f49e9cea5a 100644 --- a/crates/protocols/src/chat.rs +++ b/crates/protocols/src/chat.rs @@ -199,9 +199,28 @@ pub struct ChatCompletionRequest { /// Cache key for prompts (beta feature) pub prompt_cache_key: Option, - /// Effort level for reasoning models (low, medium, high) + /// Effort level for reasoning models. + /// + /// OpenAI-compatible callers normally send a named string, while some + /// model integrations accept a numeric value. Keep the public Rust shape + /// as a string for compatibility, but accept either JSON representation at + /// the HTTP boundary; model-specific normalization happens in the gateway. + #[serde(default, deserialize_with = "deserialize_reasoning_effort")] pub reasoning_effort: Option, + /// Internal request-origin marker. It is set only by the + /// `/v1/chat/completions` handler and is never accepted from or emitted to + /// clients. Shared Chat pipelines (for example `/v1/responses`) therefore + /// do not accidentally inherit Chat-only defaults. + #[doc(hidden)] + #[serde( + default, + skip_serializing, + deserialize_with = "ignore_chat_completions_api_request" + )] + #[schemars(skip)] + pub chat_completions_api_request: bool, + /// An object specifying the format that the model must output pub response_format: Option, @@ -342,6 +361,44 @@ pub fn thinking_from_reasoning_effort(reasoning_effort: Option<&str>) -> Option< } } +fn deserialize_reasoning_effort<'de, D>(deserializer: D) -> Result, D::Error> +where + D: serde::Deserializer<'de>, +{ + let value = Option::::deserialize(deserializer)?; + match value { + None | Some(Value::Null) => Ok(None), + Some(Value::String(value)) => Ok(Some(value)), + Some(Value::Number(value)) => Ok(Some(value.to_string())), + Some(_) => Err(serde::de::Error::custom( + "reasoning_effort must be a string, number, or null", + )), + } +} + +fn ignore_chat_completions_api_request<'de, D>(deserializer: D) -> Result +where + D: serde::Deserializer<'de>, +{ + let _ = serde::de::IgnoredAny::deserialize(deserializer)?; + Ok(false) +} + +impl ChatCompletionRequest { + /// Mark this value as originating at the public Chat Completions endpoint. + /// + /// This marker is deliberately separate from serde normalization so + /// internally constructed Chat requests used by other endpoints remain + /// distinguishable. + pub fn mark_chat_completions_api_request(&mut self) { + self.chat_completions_api_request = true; + } + + pub fn is_chat_completions_api_request(&self) -> bool { + self.chat_completions_api_request + } +} + // ============================================================================ // Validation Functions // ============================================================================ @@ -805,6 +862,48 @@ mod tests { assert_eq!(thinking_from_reasoning_effort(Some("bogus")), None); } + #[test] + fn reasoning_effort_accepts_scalar_json_and_rejects_other_types() { + for (value, expected) in [ + (json!("high"), Some("high")), + (json!(0.2), Some("0.2")), + (json!(0.99), Some("0.99")), + (Value::Null, None), + ] { + let request = request_with_output_fields(&[("reasoning_effort", value)]); + assert_eq!(request.reasoning_effort.as_deref(), expected); + } + + for value in [json!(true), json!([]), json!({"level": "high"})] { + let mut request = json!({ + "model": "test-model", + "messages": [{"role": "user", "content": "hello"}], + }); + request["reasoning_effort"] = value; + let error = serde_json::from_value::(request).unwrap_err(); + assert!(error + .to_string() + .contains("reasoning_effort must be a string, number, or null")); + } + } + + #[test] + fn chat_origin_marker_cannot_be_spoofed_or_serialized() { + let mut request: ChatCompletionRequest = serde_json::from_value(json!({ + "model": "test-model", + "messages": [{"role": "user", "content": "hello"}], + "chat_completions_api_request": true, + })) + .expect("request must deserialize"); + assert!(!request.is_chat_completions_api_request()); + assert!(!request.other.contains_key("chat_completions_api_request")); + + request.mark_chat_completions_api_request(); + assert!(request.is_chat_completions_api_request()); + let serialized = serde_json::to_value(request).expect("request must serialize"); + assert!(serialized.get("chat_completions_api_request").is_none()); + } + #[test] fn return_audio_preserves_explicit_values() { for fields in [vec![], vec![("return_audio", Value::Null)]] { diff --git a/crates/reasoning_parser/src/factory.rs b/crates/reasoning_parser/src/factory.rs index 7dc4ba89af..0e9c2b57f0 100644 --- a/crates/reasoning_parser/src/factory.rs +++ b/crates/reasoning_parser/src/factory.rs @@ -6,9 +6,9 @@ use parking_lot::RwLock; use crate::{ parsers::{ - BaseReasoningParser, CohereCmdParser, DeepSeekR1Parser, Glm45Parser, KimiParser, - MiniMaxParser, NanoV3Parser, PassthroughParser, Qwen3Parser, QwenThinkingParser, - Step3Parser, + BaseReasoningParser, CohereCmdParser, DeepSeekR1Parser, Glm45Parser, InklingParser, + KimiParser, MiniMaxParser, NanoV3Parser, PassthroughParser, Qwen3Parser, + QwenThinkingParser, Step3Parser, }, traits::{ParserConfig, ReasoningParser, DEFAULT_MAX_BUFFER_SIZE}, }; @@ -148,6 +148,8 @@ impl ParserFactory { registry.register_parser("nano_v3", || Box::new(NanoV3Parser::new())); + registry.register_parser("inkling", || Box::new(InklingParser::new())); + // standard think tokens, always_in_reasoning=false registry.register_parser("deepseek_v31", || { let config = ParserConfig { @@ -211,6 +213,9 @@ impl ParserFactory { registry.register_pattern("nemotron-super", "nano_v3"); registry.register_pattern("nano-v3", "nano_v3"); + // Inkling checkpoints use the model-family name in their ID or config. + registry.register_pattern("inkling", "inkling"); + Self { registry } } @@ -275,6 +280,12 @@ mod tests { assert_eq!(parser.model_type(), "kimi"); } + #[test] + fn test_factory_creates_inkling() { + let factory = ParserFactory::new(); + assert_eq!(factory.create("inkling-chat").model_type(), "inkling"); + } + #[test] fn test_factory_fallback_to_passthrough() { let factory = ParserFactory::new(); diff --git a/crates/reasoning_parser/src/lib.rs b/crates/reasoning_parser/src/lib.rs index 8b291627a8..58a090518a 100644 --- a/crates/reasoning_parser/src/lib.rs +++ b/crates/reasoning_parser/src/lib.rs @@ -4,8 +4,8 @@ pub mod traits; pub use factory::{ParserFactory, ParserRegistry}; pub use parsers::{ - BaseReasoningParser, CohereCmdParser, DeepSeekR1Parser, Glm45Parser, KimiParser, MiniMaxParser, - NanoV3Parser, PassthroughParser, Qwen3Parser, QwenThinkingParser, Step3Parser, + BaseReasoningParser, CohereCmdParser, DeepSeekR1Parser, Glm45Parser, InklingParser, KimiParser, + MiniMaxParser, NanoV3Parser, PassthroughParser, Qwen3Parser, QwenThinkingParser, Step3Parser, }; pub use traits::{ ParseError, ParserConfig, ParserResult, ReasoningParser, DEFAULT_MAX_BUFFER_SIZE, diff --git a/crates/reasoning_parser/src/parsers/inkling.rs b/crates/reasoning_parser/src/parsers/inkling.rs new file mode 100644 index 0000000000..6b6c3d7823 --- /dev/null +++ b/crates/reasoning_parser/src/parsers/inkling.rs @@ -0,0 +1,334 @@ +//! Inkling/TML typed-content reasoning parser. + +use crate::traits::{ParseError, ParserResult, ReasoningParser, DEFAULT_MAX_BUFFER_SIZE}; + +const CONTENT_THINKING: &str = "<|content_thinking|>"; +const CONTENT_TEXT: &str = "<|content_text|>"; +const CONTENT_INVOKE_TOOL_JSON: &str = "<|content_invoke_tool_json|>"; +const CONTENT_INVOKE_TOOL_TEXT: &str = "<|content_invoke_tool_text|>"; +const CONTENT_MODEL_END_SAMPLING: &str = "<|content_model_end_sampling|>"; +const END_MESSAGE: &str = "<|end_message|>"; +const MESSAGE_MODEL: &str = "<|message_model|>"; + +// Keep this in sync with the Inkling TMLv0 tokenizer control tokens. +const CONTROL_TOKENS: &[&str] = &[ + "<|endoftext|>", + "<|message_user|>", + MESSAGE_MODEL, + "<|message_system|>", + "<|message_tool|>", + CONTENT_TEXT, + "<|content_image|>", + CONTENT_MODEL_END_SAMPLING, + CONTENT_THINKING, + "<|content_audio_input|>", + "<|content_tool_error|>", + "<|content_xml|>", + CONTENT_INVOKE_TOOL_JSON, + CONTENT_INVOKE_TOOL_TEXT, + END_MESSAGE, + "<|audio_end|>", +]; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum BlockKind { + /// TML message headers may carry an author/tool name between the role and + /// content-type tokens. It is protocol metadata, not assistant content. + Header, + Reasoning, + Content, + Tool, + UnsupportedTool, +} + +#[derive(Debug, Clone, Copy)] +enum ControlCandidate { + Complete { start: usize, token: &'static str }, + Partial { start: usize }, +} + +/// Parser for Inkling's TML typed output blocks. +/// +/// Unlike `` formats, Inkling emits a sequence of typed blocks. +/// Structured JSON tool blocks remain in normal output verbatim so the +/// tool-call parser can consume them after reasoning separation. Unsupported +/// headerless text-mode invocations are suppressed as protocol data. +#[derive(Debug, Clone)] +pub struct InklingParser { + block_kind: Option, + buffer: String, + max_buffer_size: usize, +} + +impl InklingParser { + pub fn new() -> Self { + Self { + block_kind: None, + buffer: String::new(), + max_buffer_size: DEFAULT_MAX_BUFFER_SIZE, + } + } + + fn find_control_candidate(text: &str) -> Option { + for (start, _) in text.match_indices('<') { + let suffix = &text[start..]; + + if let Some(token) = CONTROL_TOKENS + .iter() + .copied() + .find(|token| suffix.starts_with(token)) + { + return Some(ControlCandidate::Complete { start, token }); + } + + if CONTROL_TOKENS.iter().any(|token| token.starts_with(suffix)) { + return Some(ControlCandidate::Partial { start }); + } + } + + None + } + + fn emit_text(&self, text: &str, result: &mut ParserResult) { + match self.block_kind { + Some(BlockKind::Reasoning) => result.reasoning_text.push_str(text), + Some(BlockKind::Header | BlockKind::UnsupportedTool) => {} + Some(BlockKind::Content | BlockKind::Tool) | None => { + result.normal_text.push_str(text); + } + } + } + + fn handle_control(&mut self, token: &str, result: &mut ParserResult) { + if token == MESSAGE_MODEL { + self.block_kind = Some(BlockKind::Header); + return; + } + + if token == CONTENT_INVOKE_TOOL_JSON { + self.block_kind = Some(BlockKind::Tool); + result.normal_text.push_str(token); + return; + } + + if token == CONTENT_INVOKE_TOOL_TEXT { + // The OpenAI function-call shape cannot losslessly represent + // headerless text-mode invocations. Suppress the entire frame so + // its payload cannot be mistaken for assistant answer text. + self.block_kind = Some(BlockKind::UnsupportedTool); + return; + } + + if self.block_kind == Some(BlockKind::Tool) { + result.normal_text.push_str(token); + if Self::is_end_token(token) { + self.block_kind = None; + } + return; + } + + if self.block_kind == Some(BlockKind::UnsupportedTool) { + if Self::is_end_token(token) { + self.block_kind = None; + } + return; + } + + match token { + CONTENT_THINKING => self.block_kind = Some(BlockKind::Reasoning), + CONTENT_TEXT => self.block_kind = Some(BlockKind::Content), + END_MESSAGE | CONTENT_MODEL_END_SAMPLING => self.block_kind = None, + _ => {} + } + } + + fn is_end_token(token: &str) -> bool { + matches!(token, END_MESSAGE | CONTENT_MODEL_END_SAMPLING) + } + + fn parse_buffer(&mut self, finalize: bool) -> ParserResult { + let text = std::mem::take(&mut self.buffer); + let mut result = ParserResult::default(); + let mut pos = 0; + + while pos < text.len() { + let remaining = &text[pos..]; + match Self::find_control_candidate(remaining) { + Some(ControlCandidate::Complete { start, token }) => { + self.emit_text(&remaining[..start], &mut result); + self.handle_control(token, &mut result); + pos += start + token.len(); + } + Some(ControlCandidate::Partial { start }) => { + self.emit_text(&remaining[..start], &mut result); + if finalize { + self.emit_text(&remaining[start..], &mut result); + } else { + self.buffer.push_str(&remaining[start..]); + } + break; + } + None => { + self.emit_text(remaining, &mut result); + break; + } + } + } + + result + } +} + +impl Default for InklingParser { + fn default() -> Self { + Self::new() + } +} + +impl ReasoningParser for InklingParser { + fn detect_and_parse_reasoning(&mut self, text: &str) -> Result { + if text.len() > self.max_buffer_size { + return Err(ParseError::BufferOverflow(text.len())); + } + + // Complete parsing is independent of any prior streaming state. + let mut parser = Self::new(); + parser.max_buffer_size = self.max_buffer_size; + parser.buffer.push_str(text); + Ok(parser.parse_buffer(true)) + } + + fn parse_reasoning_streaming_incremental( + &mut self, + text: &str, + ) -> Result { + let buffered_size = self.buffer.len() + text.len(); + if buffered_size > self.max_buffer_size { + return Err(ParseError::BufferOverflow(buffered_size)); + } + + self.buffer.push_str(text); + Ok(self.parse_buffer(false)) + } + + fn reset(&mut self) { + self.block_kind = None; + self.buffer.clear(); + } + + fn model_type(&self) -> &str { + "inkling" + } + + fn requires_special_tokens(&self) -> bool { + true + } + + fn is_in_reasoning(&self) -> bool { + self.block_kind == Some(BlockKind::Reasoning) + } + + fn mark_reasoning_started(&mut self) { + self.block_kind = Some(BlockKind::Reasoning); + } + + fn mark_think_start_stripped(&mut self) { + // Inkling uses typed blocks rather than a separately injected start tag. + } +} + +#[cfg(test)] +mod tests { + use super::*; + + const TOOL_BLOCK: &str = concat!( + "<|content_invoke_tool_json|>", + r#"{"name":"search","args":{"q":"rust"}}"#, + "<|end_message|>" + ); + // Canonical typed-content fixture. + // Structured tool calls carry their tool name in the model-message header + // as well as in the JSON payload; only the payload reaches the tool parser. + const TYPED_OUTPUT: &str = concat!( + "<|message_model|>", + "<|content_thinking|>check sources<|end_message|>", + "<|message_model|>", + "<|content_text|>Here is the answer.<|end_message|>", + "<|message_model|>search", + "<|content_invoke_tool_json|>", + r#"{"name":"search","args":{"q":"rust"}}"#, + "<|end_message|>", + "<|content_model_end_sampling|>" + ); + + #[test] + fn parses_typed_blocks_and_preserves_tool_framing() { + let mut parser = InklingParser::new(); + let result = parser.detect_and_parse_reasoning(TYPED_OUTPUT).unwrap(); + + assert_eq!(result.reasoning_text, "check sources"); + assert_eq!( + result.normal_text, + format!("Here is the answer.{TOOL_BLOCK}") + ); + assert!(parser.requires_special_tokens()); + } + + #[test] + fn safely_suppresses_headerless_text_tool_frame() { + let output = concat!( + "<|content_invoke_tool_text|>search for rust<|end_message|>", + "<|content_model_end_sampling|>" + ); + let mut parser = InklingParser::new(); + let result = parser.detect_and_parse_reasoning(output).unwrap(); + + assert_eq!(result.reasoning_text, ""); + assert_eq!(result.normal_text, ""); + } + + #[test] + fn streaming_buffers_control_tokens_at_every_chunk_boundary() { + let expected_normal = format!("Here is the answer.{TOOL_BLOCK}"); + + for split in TYPED_OUTPUT + .char_indices() + .map(|(index, _)| index) + .chain(std::iter::once(TYPED_OUTPUT.len())) + { + let mut parser = InklingParser::new(); + let first = parser + .parse_reasoning_streaming_incremental(&TYPED_OUTPUT[..split]) + .unwrap(); + let second = parser + .parse_reasoning_streaming_incremental(&TYPED_OUTPUT[split..]) + .unwrap(); + + assert_eq!( + format!("{}{}", first.reasoning_text, second.reasoning_text), + "check sources", + "reasoning mismatch at split {split}" + ); + assert_eq!( + format!("{}{}", first.normal_text, second.normal_text), + expected_normal, + "content mismatch at split {split}" + ); + } + } + + #[test] + fn reset_clears_partial_token_and_reasoning_state() { + let mut parser = InklingParser::new(); + parser + .parse_reasoning_streaming_incremental("<|content_thinking|>work<|end_mes") + .unwrap(); + assert!(parser.is_in_reasoning()); + + parser.reset(); + let result = parser + .parse_reasoning_streaming_incremental("plain answer") + .unwrap(); + assert_eq!(result, ParserResult::normal("plain answer".to_string())); + } +} diff --git a/crates/reasoning_parser/src/parsers/mod.rs b/crates/reasoning_parser/src/parsers/mod.rs index c46177cf8b..4479425777 100644 --- a/crates/reasoning_parser/src/parsers/mod.rs +++ b/crates/reasoning_parser/src/parsers/mod.rs @@ -2,6 +2,7 @@ pub mod base; pub mod cohere_cmd; pub mod deepseek_r1; pub mod glm45; +pub mod inkling; pub mod kimi; pub mod minimax; pub mod nano_v3; @@ -13,6 +14,7 @@ pub use base::BaseReasoningParser; pub use cohere_cmd::CohereCmdParser; pub use deepseek_r1::DeepSeekR1Parser; pub use glm45::Glm45Parser; +pub use inkling::InklingParser; pub use kimi::KimiParser; pub use minimax::MiniMaxParser; pub use nano_v3::NanoV3Parser; diff --git a/crates/reasoning_parser/src/traits.rs b/crates/reasoning_parser/src/traits.rs index 47689936c6..5c876c6331 100644 --- a/crates/reasoning_parser/src/traits.rs +++ b/crates/reasoning_parser/src/traits.rs @@ -70,6 +70,15 @@ pub trait ReasoningParser: Send + Sync { /// Get the model type this parser is designed for. fn model_type(&self) -> &str; + /// Whether detokenization must preserve model control tokens for this parser. + /// + /// Most reasoning formats use ordinary text delimiters such as ``. + /// Typed-output formats can instead use tokenizer special tokens, which must + /// reach the parser verbatim rather than being removed by detokenization. + fn requires_special_tokens(&self) -> bool { + false + } + /// Check if the parser is currently in reasoning mode. /// /// Returns true if the parser is currently parsing reasoning content. diff --git a/crates/tool_parser/src/factory.rs b/crates/tool_parser/src/factory.rs index d1c4124f16..e1f18f1391 100644 --- a/crates/tool_parser/src/factory.rs +++ b/crates/tool_parser/src/factory.rs @@ -10,8 +10,8 @@ use tokio::sync::Mutex; use crate::{ parsers::{ CohereParser, DeepSeek31Parser, DeepSeekDsmlParser, DeepSeekParser, Glm4MoeParser, - JsonParser, KimiK2Parser, LlamaParser, MinimaxM2Parser, MistralParser, PassthroughParser, - PythonicParser, QwenParser, QwenXmlParser, Step3Parser, + InklingParser, JsonParser, KimiK2Parser, LlamaParser, MinimaxM2Parser, MistralParser, + PassthroughParser, PythonicParser, QwenParser, QwenXmlParser, Step3Parser, }, traits::ToolParser, }; @@ -326,6 +326,11 @@ impl ParserFactory { || Box::new(KimiK2Parser::new()), KimiK2Parser::build_structural_tag, ); + registry.register_parser_with_structural_tag( + "inkling", + || Box::new(InklingParser::new()), + InklingParser::build_structural_tag, + ); registry.register_parser("minimax_m2", || Box::new(MinimaxM2Parser::new())); registry.register_parser("cohere", || Box::new(CohereParser::new())); @@ -402,6 +407,10 @@ impl ParserFactory { registry.map_model("Kimi-K2*", "kimik2"); registry.map_model("moonshot*/Kimi-K2*", "kimik2"); + // Inkling models use TML JSON tool calls. + registry.map_model("inkling*", "inkling"); + registry.map_model("Inkling*", "inkling"); + // MiniMax models registry.map_model("minimax*", "minimax_m2"); registry.map_model("MiniMax*", "minimax_m2"); diff --git a/crates/tool_parser/src/lib.rs b/crates/tool_parser/src/lib.rs index 2a08890a0d..c33677d457 100644 --- a/crates/tool_parser/src/lib.rs +++ b/crates/tool_parser/src/lib.rs @@ -17,9 +17,9 @@ mod tests; // Re-export types used outside this module pub use factory::{ParserFactory, PooledParser, ToolConstraint}; pub use parsers::{ - CohereParser, DeepSeek31Parser, DeepSeekDsmlParser, DeepSeekParser, Glm4MoeParser, JsonParser, - KimiK2Parser, LlamaParser, MinimaxM2Parser, MistralParser, PythonicParser, QwenParser, - Step3Parser, + CohereParser, DeepSeek31Parser, DeepSeekDsmlParser, DeepSeekParser, Glm4MoeParser, + InklingParser, JsonParser, KimiK2Parser, LlamaParser, MinimaxM2Parser, MistralParser, + PythonicParser, QwenParser, Step3Parser, }; pub use traits::ToolParser; pub use types::{FunctionCall, StreamingParseResult, ToolCall}; diff --git a/crates/tool_parser/src/parsers/inkling.rs b/crates/tool_parser/src/parsers/inkling.rs new file mode 100644 index 0000000000..0de34ad3d2 --- /dev/null +++ b/crates/tool_parser/src/parsers/inkling.rs @@ -0,0 +1,532 @@ +use std::collections::HashSet; + +use async_trait::async_trait; +use openai_protocol::common::Tool; +use serde_json::Value; + +use crate::{ + errors::ParserResult, + partial_json::PartialJson, + traits::ToolParser, + types::{FunctionCall, StreamingParseResult, ToolCall, ToolCallItem}, +}; + +const TOOL_CALL_JSON_START: &str = "<|content_invoke_tool_json|>"; +const TOOL_CALL_TEXT_START: &str = "<|content_invoke_tool_text|>"; +const END_MESSAGE: &str = "<|end_message|>"; +const MESSAGE_MODEL: &str = "<|message_model|>"; +const CONTENT_TEXT: &str = "<|content_text|>"; +const CONTENT_THINKING: &str = "<|content_thinking|>"; +const MODEL_END_SAMPLING: &str = "<|content_model_end_sampling|>"; + +const STREAM_CONTROL_TOKENS: [&str; 7] = [ + TOOL_CALL_JSON_START, + TOOL_CALL_TEXT_START, + MESSAGE_MODEL, + CONTENT_TEXT, + CONTENT_THINKING, + END_MESSAGE, + MODEL_END_SAMPLING, +]; + +const HEADER_CONTROL_TOKENS: [&str; 7] = [ + TOOL_CALL_JSON_START, + TOOL_CALL_TEXT_START, + MESSAGE_MODEL, + CONTENT_TEXT, + CONTENT_THINKING, + END_MESSAGE, + MODEL_END_SAMPLING, +]; + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum StreamingState { + Text, + MessageHeader, + ToolCallJson, + ToolCallEnd, + DiscardToolCall, +} + +/// Parser for Inkling's TML tool-call format. +/// +/// A call is emitted as: +/// `<|message_model|>name<|content_invoke_tool_json|>{"name":"...",` +/// `"args":{...}}<|end_message|>`. +/// TML response framing tokens around ordinary assistant text are removed before +/// that text is returned to the OpenAI-compatible API. Headerless text-mode +/// invocations cannot be represented losslessly by the OpenAI function-call +/// shape, so their frame is safely suppressed instead of becoming answer text. +pub struct InklingParser { + partial_json: PartialJson, + buffer: String, + state: StreamingState, + current_tool_index: usize, + current_tool_name: Option, +} + +impl InklingParser { + pub fn new() -> Self { + Self { + partial_json: PartialJson::default(), + buffer: String::new(), + state: StreamingState::Text, + current_tool_index: 0, + current_tool_name: None, + } + } + + /// Build an xgrammar structural tag that constrains the JSON arguments for + /// each declared tool while retaining Inkling's native TML framing. + pub fn build_structural_tag(tools: &[Tool], at_least_one: bool) -> Value { + let tags = tools + .iter() + .filter(|tool| !tool.function.name.is_empty()) + .map(|tool| { + let name = serde_json::to_string(&tool.function.name).unwrap_or_default(); + serde_json::json!({ + "begin": format!( + "{TOOL_CALL_JSON_START}{{\"name\":{name},\"args\":" + ), + "content": { + "type": "json_schema", + "json_schema": &tool.function.parameters, + }, + "end": format!("}}{END_MESSAGE}"), + }) + }) + .collect::>(); + + serde_json::json!({ + "format": { + "type": "triggered_tags", + "triggers": [TOOL_CALL_JSON_START], + "tags": tags, + "at_least_one": at_least_one, + } + }) + } + + fn clean_normal_text(text: &str) -> String { + let controls = [ + MESSAGE_MODEL, + CONTENT_TEXT, + CONTENT_THINKING, + END_MESSAGE, + MODEL_END_SAMPLING, + ]; + let mut output = String::new(); + let mut cursor = 0; + let mut in_header = false; + + while cursor < text.len() { + let remaining = &text[cursor..]; + let Some((start, token)) = find_earliest_token(remaining, &controls) else { + if !in_header { + output.push_str(remaining); + } + break; + }; + + if !in_header { + output.push_str(&remaining[..start]); + } + cursor += start + token.len(); + in_header = token == MESSAGE_MODEL; + } + + output + } + + fn parse_complete_impl( + text: &str, + allowed_tools: Option<&HashSet<&str>>, + ) -> (String, Vec) { + let tool_markers = [TOOL_CALL_JSON_START, TOOL_CALL_TEXT_START]; + let mut normal_text = String::new(); + let mut calls = Vec::new(); + let mut cursor = 0; + + while let Some((relative_start, marker)) = + find_earliest_token(&text[cursor..], &tool_markers) + { + let marker_start = cursor + relative_start; + normal_text.push_str(&Self::clean_normal_text(&text[cursor..marker_start])); + let payload_start = marker_start + marker.len(); + + if marker == TOOL_CALL_TEXT_START { + tracing::warn!( + "Ignoring TML text-mode tool invocation; OpenAI tool calls require structured JSON" + ); + cursor = after_discarded_tool_frame(text, payload_start); + continue; + } + + let payload = &text[payload_start..]; + let whitespace = payload.len() - payload.trim_start().len(); + let json_start = payload_start + whitespace; + let Some(json_len) = complete_json_object_len(&text[json_start..]) else { + tracing::warn!("Ignoring malformed TML JSON tool invocation"); + cursor = after_discarded_tool_frame(text, payload_start); + continue; + }; + let json_end = json_start + json_len; + let after_json = &text[json_end..]; + let whitespace = after_json.len() - after_json.trim_start().len(); + let end_start = json_end + whitespace; + + let Some(call) = parse_tool_call_json(&text[json_start..json_end], allowed_tools) + else { + // Malformed and undefined calls are protocol data. Suppress the + // whole frame, then resume at the next typed message. + cursor = after_discarded_tool_frame(text, json_end); + continue; + }; + calls.push(call); + + if text[end_start..].starts_with(END_MESSAGE) { + cursor = end_start + END_MESSAGE.len(); + } else if text[end_start..].starts_with(MODEL_END_SAMPLING) { + cursor = end_start + MODEL_END_SAMPLING.len(); + } else { + // Match the streaming parser's permissive ToolCallEnd state: + // after a complete valid JSON object, a missing end marker does + // not turn subsequent bytes into tool payload. + cursor = end_start; + } + } + + normal_text.push_str(&Self::clean_normal_text(&text[cursor..])); + (normal_text, calls) + } + + fn process_text(&mut self, result: &mut StreamingParseResult) -> bool { + let state_markers = [MESSAGE_MODEL, TOOL_CALL_JSON_START, TOOL_CALL_TEXT_START]; + if let Some((start, marker)) = find_earliest_token(&self.buffer, &state_markers) { + let normal_text = Self::clean_normal_text(&self.buffer[..start]); + result.normal_text.push_str(&normal_text); + self.buffer.drain(..start + marker.len()); + self.state = if marker == MESSAGE_MODEL { + StreamingState::MessageHeader + } else if marker == TOOL_CALL_JSON_START { + StreamingState::ToolCallJson + } else { + tracing::warn!( + "Ignoring TML text-mode tool invocation; OpenAI tool calls require structured JSON" + ); + StreamingState::DiscardToolCall + }; + return true; + } + + let keep = longest_partial_token_suffix(&self.buffer, &STREAM_CONTROL_TOKENS); + let safe_len = self.buffer.len() - keep; + if safe_len > 0 { + let safe_text = self.buffer[..safe_len].to_string(); + self.buffer.drain(..safe_len); + result + .normal_text + .push_str(&Self::clean_normal_text(&safe_text)); + } + false + } + + fn process_message_header(&mut self) -> bool { + if let Some((start, token)) = find_earliest_token(&self.buffer, &HEADER_CONTROL_TOKENS) { + // Everything before the content-type token is an author, tool + // name, or channel header. It is metadata and must not be emitted + // as assistant content. + self.buffer.drain(..start + token.len()); + self.state = if token == MESSAGE_MODEL { + StreamingState::MessageHeader + } else if token == TOOL_CALL_JSON_START { + StreamingState::ToolCallJson + } else if token == TOOL_CALL_TEXT_START { + tracing::warn!( + "Ignoring TML text-mode tool invocation; OpenAI tool calls require structured JSON" + ); + StreamingState::DiscardToolCall + } else { + StreamingState::Text + }; + return true; + } + + // Header bytes are never user-visible. Drop author/channel bytes while + // retaining only a possible control-token prefix across chunk splits. + let keep = longest_partial_token_suffix(&self.buffer, &HEADER_CONTROL_TOKENS); + let safe_len = self.buffer.len() - keep; + if safe_len > 0 { + self.buffer.drain(..safe_len); + } + false + } + + fn process_tool_call( + &mut self, + allowed_tools: &HashSet<&str>, + result: &mut StreamingParseResult, + ) -> bool { + let whitespace = self.buffer.len() - self.buffer.trim_start().len(); + if whitespace > 0 { + self.buffer.drain(..whitespace); + } + if self.buffer.is_empty() { + return false; + } + + if !self.buffer.starts_with('{') { + self.current_tool_name = None; + self.state = StreamingState::DiscardToolCall; + return true; + } + + if self.current_tool_name.is_none() { + if let Ok((Value::Object(object), _)) = + self.partial_json.parse_value(&self.buffer, false) + { + if let Some(name) = object.get("name").and_then(Value::as_str) { + if allowed_tools.contains(name) { + self.current_tool_name = Some(name.to_string()); + result.calls.push(ToolCallItem { + tool_index: self.current_tool_index, + name: Some(name.to_string()), + parameters: String::new(), + }); + } else { + tracing::debug!("Inkling attempted to call undefined tool: {}", name); + self.state = StreamingState::DiscardToolCall; + return true; + } + } + } + } + + let Some(json_len) = complete_json_object_len(&self.buffer) else { + return false; + }; + let json = &self.buffer[..json_len]; + let Some(call) = parse_tool_call_json(json, Some(allowed_tools)) else { + self.current_tool_name = None; + self.state = StreamingState::DiscardToolCall; + return true; + }; + + // A complete object may expose the name and arguments in the same + // chunk. Emit the name first in that case, as required by stream APIs. + if self.current_tool_name.is_none() { + result.calls.push(ToolCallItem { + tool_index: self.current_tool_index, + name: Some(call.function.name.clone()), + parameters: String::new(), + }); + } + result.calls.push(ToolCallItem { + tool_index: self.current_tool_index, + name: None, + parameters: call.function.arguments, + }); + + self.buffer.drain(..json_len); + self.current_tool_index += 1; + self.current_tool_name = None; + self.state = StreamingState::ToolCallEnd; + true + } + + fn process_tool_call_end(&mut self) -> bool { + let whitespace = self.buffer.len() - self.buffer.trim_start().len(); + if whitespace > 0 { + self.buffer.drain(..whitespace); + } + if self.buffer.is_empty() { + return false; + } + if self.buffer.starts_with(END_MESSAGE) { + self.buffer.drain(..END_MESSAGE.len()); + self.state = StreamingState::Text; + return true; + } + if END_MESSAGE.starts_with(&self.buffer) { + return false; + } + + // Be permissive when a backend omits the end token after an otherwise + // complete call. The remaining bytes still need normal text handling. + self.state = StreamingState::Text; + true + } + + fn discard_tool_call(&mut self) -> bool { + let end_tokens = [END_MESSAGE, MODEL_END_SAMPLING]; + if let Some((end, token)) = find_earliest_token(&self.buffer, &end_tokens) { + self.buffer.drain(..end + token.len()); + self.current_tool_name = None; + self.state = StreamingState::Text; + return true; + } + + let keep = longest_partial_token_suffix(&self.buffer, &end_tokens); + let safe_len = self.buffer.len() - keep; + if safe_len > 0 { + self.buffer.drain(..safe_len); + } + false + } +} + +impl Default for InklingParser { + fn default() -> Self { + Self::new() + } +} + +#[async_trait] +impl ToolParser for InklingParser { + async fn parse_complete(&self, text: &str) -> ParserResult<(String, Vec)> { + Ok(Self::parse_complete_impl(text, None)) + } + + async fn parse_complete_with_tools( + &self, + text: &str, + tools: &[Tool], + ) -> ParserResult<(String, Vec)> { + let allowed_tools = tools + .iter() + .map(|tool| tool.function.name.as_str()) + .collect::>(); + Ok(Self::parse_complete_impl(text, Some(&allowed_tools))) + } + + async fn parse_incremental( + &mut self, + chunk: &str, + tools: &[Tool], + ) -> ParserResult { + self.buffer.push_str(chunk); + let allowed_tools = tools + .iter() + .map(|tool| tool.function.name.as_str()) + .collect::>(); + let mut result = StreamingParseResult::default(); + + loop { + let progressed = match self.state { + StreamingState::Text => self.process_text(&mut result), + StreamingState::MessageHeader => self.process_message_header(), + StreamingState::ToolCallJson => self.process_tool_call(&allowed_tools, &mut result), + StreamingState::ToolCallEnd => self.process_tool_call_end(), + StreamingState::DiscardToolCall => self.discard_tool_call(), + }; + if !progressed { + break; + } + } + + Ok(result) + } + + fn has_tool_markers(&self, text: &str) -> bool { + text.contains(TOOL_CALL_JSON_START) || text.contains(TOOL_CALL_TEXT_START) + } + + fn reset(&mut self) { + self.buffer.clear(); + self.state = StreamingState::Text; + self.current_tool_index = 0; + self.current_tool_name = None; + } +} + +fn parse_tool_call_json(json: &str, allowed_tools: Option<&HashSet<&str>>) -> Option { + let value = serde_json::from_str::(json).ok()?; + let object = value.as_object()?; + let name = object.get("name")?.as_str()?; + let args = object.get("args")?.as_object()?; + + if allowed_tools.is_some_and(|tools| !tools.contains(name)) { + tracing::debug!("Inkling attempted to call undefined tool: {}", name); + return None; + } + + Some(ToolCall { + function: FunctionCall { + name: name.to_string(), + arguments: serde_json::to_string(args).ok()?, + }, + }) +} + +/// Find the earliest complete protocol token, preferring the token list order +/// when two entries start at the same byte offset. +fn find_earliest_token<'a>(text: &str, tokens: &'a [&'a str]) -> Option<(usize, &'a str)> { + tokens + .iter() + .filter_map(|token| text.find(token).map(|start| (start, *token))) + .min_by_key(|(start, _)| *start) +} + +/// Return the first byte after a discarded tool frame. Both end markers match +/// the streaming parser's `DiscardToolCall` recovery behavior. If a malformed +/// frame is unterminated, all remaining bytes belong to that protocol frame. +fn after_discarded_tool_frame(text: &str, payload_start: usize) -> usize { + let end_tokens = [END_MESSAGE, MODEL_END_SAMPLING]; + find_earliest_token(&text[payload_start..], &end_tokens) + .map_or(text.len(), |(relative_end, token)| { + payload_start + relative_end + token.len() + }) +} + +/// Return the byte length of one complete top-level JSON object. Braces inside +/// quoted strings are ignored, and trailing data is left for the caller. +fn complete_json_object_len(text: &str) -> Option { + if !text.starts_with('{') { + return None; + } + + let mut depth = 0usize; + let mut in_string = false; + let mut escaped = false; + + for (index, byte) in text.bytes().enumerate() { + if in_string { + if escaped { + escaped = false; + } else if byte == b'\\' { + escaped = true; + } else if byte == b'"' { + in_string = false; + } + continue; + } + + match byte { + b'"' => in_string = true, + b'{' => depth += 1, + b'}' => { + depth = depth.checked_sub(1)?; + if depth == 0 { + return Some(index + 1); + } + } + _ => {} + } + } + None +} + +fn longest_partial_token_suffix(buffer: &str, tokens: &[&str]) -> usize { + tokens + .iter() + .flat_map(|token| { + token + .char_indices() + .skip(1) + .map(move |(index, _)| &token[..index]) + }) + .filter(|prefix| buffer.ends_with(prefix)) + .map(str::len) + .max() + .unwrap_or(0) +} diff --git a/crates/tool_parser/src/parsers/mod.rs b/crates/tool_parser/src/parsers/mod.rs index 216473741e..50b6c57998 100644 --- a/crates/tool_parser/src/parsers/mod.rs +++ b/crates/tool_parser/src/parsers/mod.rs @@ -8,6 +8,7 @@ pub mod deepseek; pub mod deepseek31; pub mod deepseek_dsml; pub mod glm4_moe; +pub mod inkling; pub mod json; pub mod kimik2; pub mod llama; @@ -28,6 +29,7 @@ pub use deepseek::DeepSeekParser; pub use deepseek31::DeepSeek31Parser; pub use deepseek_dsml::DeepSeekDsmlParser; pub use glm4_moe::Glm4MoeParser; +pub use inkling::InklingParser; pub use json::JsonParser; pub use kimik2::KimiK2Parser; pub use llama::LlamaParser; diff --git a/crates/tool_parser/tests/tool_parser_inkling.rs b/crates/tool_parser/tests/tool_parser_inkling.rs new file mode 100644 index 0000000000..dfefc645c9 --- /dev/null +++ b/crates/tool_parser/tests/tool_parser_inkling.rs @@ -0,0 +1,179 @@ +mod common; + +use common::create_test_tools; +use serde_json::json; +use tool_parser::{InklingParser, ParserFactory, ToolParser}; + +#[tokio::test] +async fn test_inkling_complete_parsing() { + let parser = InklingParser::new(); + let tools = create_test_tools(); + let input = concat!( + "<|message_model|><|content_text|>I'll check.<|end_message|>", + "<|message_model|>search<|content_invoke_tool_json|>{\"name\":\"search\",\"args\":", + "{\"query\":\"Rust {language}\"}}<|end_message|>", + "<|message_model|>get_weather<|content_invoke_tool_json|>{\"name\":\"get_weather\",\"args\":", + "{\"city\":\"Tokyo\"}}<|end_message|>", + "<|content_model_end_sampling|>" + ); + + let (normal_text, calls) = parser + .parse_complete_with_tools(input, &tools) + .await + .unwrap(); + + assert_eq!(normal_text, "I'll check."); + assert_eq!(calls.len(), 2); + assert_eq!(calls[0].function.name, "search"); + assert_eq!(calls[1].function.name, "get_weather"); + assert_eq!( + serde_json::from_str::(&calls[0].function.arguments).unwrap(), + json!({"query": "Rust {language}"}) + ); +} + +#[tokio::test] +async fn test_inkling_rejects_undefined_tool() { + let parser = InklingParser::new(); + let tools = create_test_tools(); + let input = concat!( + "Before", + "<|content_invoke_tool_json|>", + "{\"name\":\"not_declared\",\"args\":{}}", + "<|end_message|>" + ); + + let (normal_text, calls) = parser + .parse_complete_with_tools(input, &tools) + .await + .unwrap(); + assert!(calls.is_empty()); + assert_eq!(normal_text, "Before"); +} + +#[tokio::test] +async fn test_inkling_streaming_buffers_split_control_tokens() { + let tools = create_test_tools(); + let mut parser = InklingParser::new(); + let chunks = [ + "<|message_", + "model|><|content_", + "text|>Before<|end_mes", + "sage|><|message_model|>sea", + "rch<|content_invoke_tool_json|>{\"name\":\"sea", + "rch\",\"args\":{\"query\":\"Rust {lang}", + "\"}}<|end_mes", + "sage|><|content_", + "text|>After<|content_model_end_", + "sampling|>", + ]; + + let mut normal_text = String::new(); + let mut names = Vec::new(); + let mut arguments = String::new(); + for chunk in chunks { + let result = parser.parse_incremental(chunk, &tools).await.unwrap(); + normal_text.push_str(&result.normal_text); + for call in result.calls { + if let Some(name) = call.name { + names.push((call.tool_index, name)); + } + arguments.push_str(&call.parameters); + } + } + + assert_eq!(normal_text, "BeforeAfter"); + assert_eq!(names, vec![(0, "search".to_string())]); + assert_eq!( + serde_json::from_str::(&arguments).unwrap(), + json!({"query": "Rust {lang}"}) + ); +} + +#[tokio::test] +async fn test_inkling_streaming_skips_undefined_call() { + let tools = create_test_tools(); + let mut parser = InklingParser::new(); + let chunks = [ + "<|content_invoke_tool_json|>{\"name\":\"unknown\",", + "\"args\":{}}<|end_mes", + "sage|><|content_invoke_tool_json|>", + "{\"name\":\"search\",\"args\":{}}<|end_message|>", + ]; + + let mut calls = Vec::new(); + for chunk in chunks { + calls.extend(parser.parse_incremental(chunk, &tools).await.unwrap().calls); + } + + assert_eq!(calls.len(), 2); + assert_eq!(calls[0].tool_index, 0); + assert_eq!(calls[0].name.as_deref(), Some("search")); + assert_eq!(calls[1].parameters, "{}"); +} + +#[test] +fn test_inkling_factory_and_structural_tag() { + let factory = ParserFactory::new(); + assert!(factory.has_parser("inkling")); + assert_eq!( + factory + .registry() + .resolve_model_to_parser("Inkling-7B") + .as_deref(), + Some("inkling") + ); + assert!(factory.registry().has_structural_tag("inkling")); + + let tools = create_test_tools(); + let tag = InklingParser::build_structural_tag(&tools[..1], true); + assert_eq!(tag["format"]["type"], "triggered_tags"); + assert_eq!( + tag["format"]["tags"][0]["begin"], + "<|content_invoke_tool_json|>{\"name\":\"search\",\"args\":" + ); + assert_eq!(tag["format"]["tags"][0]["end"], "}<|end_message|>"); + assert_eq!(tag["format"]["at_least_one"], true); +} + +#[tokio::test] +async fn test_inkling_complete_recovers_after_malformed_tool_frame() { + let input = concat!( + "<|message_model|><|content_text|>Before<|end_message|>", + "<|content_invoke_tool_json|>not-json<|end_message|>", + "<|message_model|><|content_text|>After malformed.<|end_message|>", + "<|content_invoke_tool_json|>{\"name\":\"search\",\"args\":{}}<|end_message|>", + "<|message_model|><|content_text|>Tail<|end_message|>", + "<|content_model_end_sampling|>" + ); + let tools = create_test_tools(); + let parser = InklingParser::new(); + + let (normal_text, calls) = parser + .parse_complete_with_tools(input, &tools) + .await + .unwrap(); + + assert_eq!(normal_text, "BeforeAfter malformed.Tail"); + assert_eq!(calls.len(), 1); + assert_eq!(calls[0].function.name, "search"); + assert_eq!(calls[0].function.arguments, "{}"); +} + +#[tokio::test] +async fn test_inkling_text_mode_is_safely_suppressed() { + let input = concat!( + "<|message_model|><|content_text|>Before<|end_message|>", + "<|content_invoke_tool_text|>search for weather in SF<|end_message|>", + "<|content_model_end_sampling|>" + ); + let tools = create_test_tools(); + + let parser = InklingParser::new(); + let (normal_text, calls) = parser + .parse_complete_with_tools(input, &tools) + .await + .unwrap(); + assert_eq!(normal_text, "Before"); + assert!(calls.is_empty()); +} diff --git a/grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py b/grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py index 6f1763f79c..83c18cacf4 100644 --- a/grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py +++ b/grpc_servicer/smg_grpc_servicer/tokenspeed/servicer.py @@ -481,6 +481,17 @@ def config_value(config, name: str, default=None): audio_modality = common_pb2.AUDIO video_modality = common_pb2.VIDEO model_type = config_value(hf_config, "model_type", "") if hf_config is not None else "" + is_inkling = model_type == "inkling_mm_model" + + if is_inkling: + supported = [] + vision_config = config_value(hf_config, "vision_config") + if config_value(vision_config, "decoder_dmodel") is not None: + supported.append(image_modality) + audio_config = config_value(hf_config, "audio_config") + if config_value(audio_config, "decoder_dmodel") is not None: + supported.append(audio_modality) + return supported thinker_config = config_value(hf_config, "thinker_config") if model_type == "qwen3_asr": diff --git a/model_gateway/src/routers/grpc/context.rs b/model_gateway/src/routers/grpc/context.rs index c62bf501c8..77d457561b 100644 --- a/model_gateway/src/routers/grpc/context.rs +++ b/model_gateway/src/routers/grpc/context.rs @@ -97,10 +97,11 @@ impl std::fmt::Display for FinalResponse { pub(crate) struct SharedComponents { pub tokenizer_registry: Arc, pub tool_parser_factory: ToolParserFactory, - #[expect(dead_code)] pub reasoning_parser_factory: ReasoningParserFactory, /// Configured tool parser name (from CLI `--tool-call-parser`) pub configured_tool_parser: Option, + /// Configured reasoning parser name (from CLI `--reasoning-parser`) + pub configured_reasoning_parser: Option, /// Multimodal processing components (initialized at router creation) pub multimodal: Option>, } diff --git a/model_gateway/src/routers/grpc/multimodal/process.rs b/model_gateway/src/routers/grpc/multimodal/process.rs index f41b7c4904..718196ef31 100644 --- a/model_gateway/src/routers/grpc/multimodal/process.rs +++ b/model_gateway/src/routers/grpc/multimodal/process.rs @@ -6,7 +6,7 @@ use std::{collections::HashMap, sync::Arc, time::Instant}; -use anyhow::Result; +use anyhow::{Context, Result}; use futures::future::try_join_all; use llm_multimodal::{ AsyncMultiModalTracker, AudioClip, EncoderFieldLayouts, ImageFrame, Modality, ModelMetadata, @@ -607,7 +607,16 @@ fn expand_tokens_for_modalities( .collect::>>()?; let offset = expanded.len(); let length = replacement_tokens.len(); - let patches = patch_ranges(offset, &replacement_tokens, expansion.placeholder_token_id); + let patches = if let Some(feature_ranges) = &replacement.feature_ranges { + explicit_feature_ranges(offset, length, feature_ranges).with_context(|| { + format!( + "Invalid explicit feature ranges in {} replacement {item_index}", + expansion.modality + ) + })? + } else { + patch_ranges(offset, &replacement_tokens, expansion.placeholder_token_id) + }; expanded.extend(replacement_tokens); let prefix = replacement.structural_prefix.min(offset); bindings[idx].push(PromptBinding { @@ -678,6 +687,46 @@ fn patch_ranges( ranges } +fn explicit_feature_ranges( + replacement_offset: usize, + replacement_length: usize, + relative_ranges: &[PlaceholderRange], +) -> Result> { + anyhow::ensure!( + !relative_ranges.is_empty(), + "explicit feature ranges must not be empty" + ); + let mut absolute_ranges = Vec::with_capacity(relative_ranges.len()); + let mut previous_end = 0usize; + for (index, range) in relative_ranges.iter().enumerate() { + anyhow::ensure!(range.length > 0, "feature range {index} has zero length"); + let relative_end = range + .offset + .checked_add(range.length) + .ok_or_else(|| anyhow::anyhow!("feature range {index} overflows usize"))?; + anyhow::ensure!( + relative_end <= replacement_length, + "feature range {index} ({}, {}) exceeds replacement length {replacement_length}", + range.offset, + range.length + ); + anyhow::ensure!( + index == 0 || range.offset >= previous_end, + "feature range {index} overlaps or is out of order" + ); + absolute_ranges.push(PlaceholderRange { + offset: replacement_offset + .checked_add(range.offset) + .ok_or_else(|| { + anyhow::anyhow!("feature range {index} absolute offset overflows") + })?, + length: range.length, + }); + previous_end = relative_end; + } + Ok(absolute_ranges) +} + #[cfg(test)] mod tests { use bytes::Bytes; @@ -707,6 +756,7 @@ mod tests { modality: Modality::Image, placeholder_token: "".to_string(), tokens: vec![50, 50, 50, 50], // Expand to 4 tokens + feature_ranges: None, structural_prefix: 0, }]; @@ -737,6 +787,7 @@ mod tests { modality: Modality::Video, placeholder_token: "<|video_pad|>".to_string(), tokens: vec![50, 50, 50], // expands to 3 video tokens + feature_ranges: None, structural_prefix: 1, }]; @@ -779,12 +830,14 @@ mod tests { modality: Modality::Image, placeholder_token: "".to_string(), tokens: vec![50, 50], // 2 tokens for first image + feature_ranges: None, structural_prefix: 0, }, PromptReplacement { modality: Modality::Image, placeholder_token: "".to_string(), tokens: vec![60, 60, 60], // 3 tokens for second image + feature_ranges: None, structural_prefix: 0, }, ]; @@ -814,6 +867,7 @@ mod tests { modality: Modality::Image, placeholder_token: "".to_string(), tokens: vec![88, 92, 92, 92, 93, 92, 92, 92, 89], // start + patches + sep + patches + end + feature_ranges: None, structural_prefix: 0, }]; @@ -837,6 +891,49 @@ mod tests { assert_eq!((patch[1].offset, patch[1].length), (6, 3)); } + #[test] + fn test_explicit_feature_range_excludes_same_id_structural_header() { + let replacements = vec![PromptReplacement::sequence( + Modality::Image, + "<|content_image|>", + vec![200005, 200005, 200005, 200005], + ) + .with_feature_span(1, 3)]; + let expansion = ModalityExpansion { + modality: Modality::Image, + search_token_id: Some(200005), + placeholder_token_id: Some(200005), + replacements: &replacements, + }; + + let result = expand_tokens_for_modalities(&[1, 200005, 2], &[expansion]).unwrap(); + + assert_eq!(result.token_ids, vec![1, 200005, 200005, 200005, 200005, 2]); + assert_eq!(result.bindings[0][0].structural.offset, 1); + assert_eq!(result.bindings[0][0].structural.length, 4); + assert_eq!(result.bindings[0][0].patches.len(), 1); + assert_eq!(result.bindings[0][0].patches[0].offset, 2); + assert_eq!(result.bindings[0][0].patches[0].length, 3); + } + + #[test] + fn test_explicit_feature_range_must_fit_replacement() { + let replacements = + vec![ + PromptReplacement::sequence(Modality::Image, "", vec![50, 50]) + .with_feature_span(1, 2), + ]; + let expansion = ModalityExpansion { + modality: Modality::Image, + search_token_id: Some(100), + placeholder_token_id: Some(50), + replacements: &replacements, + }; + + let error = expand_tokens_for_modalities(&[100], &[expansion]).unwrap_err(); + assert!(format!("{error:#}").contains("exceeds replacement length")); + } + #[test] fn test_expand_tokens_preserves_template_owned_audio_end() { // The template owns the audio end marker. Expansion preserves the @@ -854,6 +951,7 @@ mod tests { audio_placeholder as i32, audio_placeholder as i32, ], + feature_ranges: None, structural_prefix: 0, }]; @@ -914,6 +1012,7 @@ mod tests { audio_placeholder as i32, audio_placeholder as i32, ], + feature_ranges: None, structural_prefix: 0, }]; let image_replacements = vec![PromptReplacement { @@ -925,6 +1024,7 @@ mod tests { image_placeholder as i32, image_placeholder as i32, ], + feature_ranges: None, structural_prefix: 0, }]; let expansions = vec![ @@ -979,18 +1079,21 @@ mod tests { modality: Modality::Image, placeholder_token: "".to_string(), tokens: vec![110, 101, 111], + feature_ranges: None, structural_prefix: 0, }]; let video_replacements = vec![PromptReplacement { modality: Modality::Video, placeholder_token: "