Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion bindings/golang/src/client.rs
Original file line number Diff line number Diff line change
Expand Up @@ -155,7 +155,7 @@ pub unsafe extern "C" fn sgl_client_chat_completion_stream(
};

// Process messages and apply chat template
let processed_messages = match process_chat_messages(&chat_request, tokenizer.as_ref()) {
let processed_messages = match process_chat_messages(&chat_request, tokenizer.as_ref(), None) {
Ok(msgs) => msgs,
Err(e) => {
set_error_message(error_out, &format!("Failed to process messages: {e}"));
Expand Down
2 changes: 1 addition & 1 deletion bindings/golang/src/policy.rs
Original file line number Diff line number Diff line change
Expand Up @@ -499,7 +499,7 @@ pub unsafe extern "C" fn sgl_multi_client_chat_completion_stream(
};

// Process messages and apply chat template
let processed_messages = match process_chat_messages(&chat_request, tokenizer.as_ref()) {
let processed_messages = match process_chat_messages(&chat_request, tokenizer.as_ref(), None) {
Ok(msgs) => msgs,
Err(e) => {
set_error_message(error_out, &format!("Failed to process messages: {e}"));
Expand Down
2 changes: 1 addition & 1 deletion bindings/golang/src/preprocessor.rs
Original file line number Diff line number Diff line change
Expand Up @@ -37,7 +37,7 @@ fn preprocess_impl(
tokenizer: &dyn Tokenizer,
) -> Result<PreprocessResult, (SglErrorCode, String)> {
// Process chat messages (apply chat_template)
let processed_messages = process_chat_messages(chat_request, tokenizer).map_err(|e| {
let processed_messages = process_chat_messages(chat_request, tokenizer, None).map_err(|e| {
(
SglErrorCode::ParsingError,
format!("Failed to process chat messages: {e}"),
Expand Down
26 changes: 18 additions & 8 deletions crates/multimodal/src/registry/phi3_v.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ use serde_json::{json, Value};

use crate::{
registry::{ModelMetadata, ModelProcessorSpec, RegistryResult},
types::{Modality, PromptReplacement, TokenId},
types::{FieldLayout, Modality, PromptReplacement, TokenId},
vision::image_processor::PreprocessedImages,
};

Expand Down Expand Up @@ -32,11 +32,11 @@ impl ModelProcessorSpec for Phi3VisionSpec {
}

fn placeholder_token(&self, _metadata: &ModelMetadata) -> RegistryResult<String> {
Ok("<image>".to_owned())
Ok("<|image|>".to_owned())
}

fn placeholder_token_id(&self, metadata: &ModelMetadata) -> RegistryResult<TokenId> {
metadata.token_id("<image>")
metadata.token_id("<|image|>")
}

fn modality_limits(
Expand All @@ -50,18 +50,28 @@ impl ModelProcessorSpec for Phi3VisionSpec {
Ok(json!({}))
}

fn field_layouts(&self) -> HashMap<String, FieldLayout> {
HashMap::from([
("pixel_values".to_string(), FieldLayout::Batched),
("image_sizes".to_string(), FieldLayout::Batched),
])
}

fn prompt_replacements(
&self,
metadata: &ModelMetadata,
preprocessed: &PreprocessedImages,
) -> RegistryResult<Vec<PromptReplacement>> {
let token_id = self.placeholder_token_id(metadata)?;
let token = self.placeholder_token(metadata)?;
let count = Self::tokens_per_image(metadata);
let fallback = Self::tokens_per_image(metadata);
Ok(preprocessed
.image_sizes
.num_img_tokens
.iter()
.map(|_| PromptReplacement::repeated(Modality::Image, &token, token_id, count))
.map(|&count| {
let n = if count > 0 { count } else { fallback };
PromptReplacement::repeated(Modality::Image, &token, token_id, n)
})
.collect())
}
}
Expand All @@ -77,7 +87,7 @@ mod tests {

#[test]
fn phi3_uses_num_img_tokens() {
let tokenizer = TestTokenizer::new(&[("<image>", 555)]);
let tokenizer = TestTokenizer::new(&[("<|image|>", 555)]);
let config = json!({
"model_type": "phi3_v",
"img_processor": {"num_img_tokens": 144}
Expand All @@ -98,7 +108,7 @@ mod tests {

#[test]
fn phi3_matches_alias_via_model_type() {
let tokenizer = TestTokenizer::new(&[("<image>", 555)]);
let tokenizer = TestTokenizer::new(&[("<|image|>", 555)]);
let config = json!({
"model_type": "phi3_v",
"img_processor": {"num_img_tokens": 144}
Expand Down
28 changes: 28 additions & 0 deletions model_gateway/src/routers/grpc/multimodal.rs
Original file line number Diff line number Diff line change
Expand Up @@ -156,6 +156,34 @@ pub(crate) struct MultimodalIntermediate {
pub keep_on_cpu_keys: Vec<String>,
}

/// Resolve the placeholder token string for a multimodal model.
///
/// Loads the model config and looks up the model spec to get the placeholder
/// token (e.g. `"<|image|>"` for Phi-3-vision). Returns `None` if the model
/// is not recognized as multimodal.
pub(crate) async fn resolve_placeholder_token(
model_id: &str,
tokenizer: &dyn TokenizerTrait,
components: &MultimodalComponents,
tokenizer_source: &str,
) -> Result<Option<String>> {
let model_config = components
.get_or_load_config(model_id, tokenizer_source)
.await?;
let metadata = ModelMetadata {
model_id,
tokenizer,
config: &model_config.config,
};
let spec = match components.model_registry.lookup(&metadata) {
Some(s) => s,
None => return Ok(None),
};
Ok(Some(spec.placeholder_token(&metadata).map_err(|e| {
anyhow::anyhow!("Failed to get placeholder token: {e}")
})?))
}

/// Check if any messages in the request contain multimodal content (images).
pub(crate) fn has_multimodal_content(messages: &[ChatMessage]) -> bool {
messages.iter().any(|msg| {
Expand Down
130 changes: 79 additions & 51 deletions model_gateway/src/routers/grpc/regular/stages/chat/preparation.rs
Original file line number Diff line number Diff line change
Expand Up @@ -48,32 +48,12 @@ impl ChatPreparationStage {
// Step 1: Filter tools if needed
let body_ref = utils::filter_chat_request_by_tool_choice(request);

// Step 2: Process messages and apply chat template
let processed_messages = match utils::process_chat_messages(&body_ref, &*tokenizer) {
Ok(msgs) => msgs,
Err(e) => {
error!(function = "ChatPreparationStage::execute", error = %e, "Failed to process chat messages");
return Err(error::bad_request("process_messages_failed", e));
}
};

// Step 3: Tokenize the processed text (no special tokens - chat template already handles them)
let encoding = match tokenizer.encode(&processed_messages.text, false) {
Ok(encoding) => encoding,
Err(e) => {
error!(function = "ChatPreparationStage::execute", error = %e, "Tokenization failed");
return Err(error::internal_error(
"tokenization_failed",
format!("Tokenization failed: {e}"),
));
}
};

let mut token_ids = encoding.token_ids().to_vec();

// Step 3.5: Full multimodal processing (fetch + preprocess + expand tokens + hash)
let mut multimodal_intermediate = None;
if multimodal::has_multimodal_content(&request.messages) {
// Resolve multimodal context once: placeholder token, model_id, tokenizer_source.
// The placeholder is passed to process_chat_messages so that string-format chat
// templates insert it per image instead of stripping image parts. The remaining
// fields are reused by process_multimodal to avoid duplicate lookups.
let is_multimodal = multimodal::has_multimodal_content(&request.messages);
let (image_placeholder, mm_context) = if is_multimodal {
if let Some(mm_components) = ctx.components.multimodal.as_ref() {
let model_id = ctx.input.model_id.as_str();
let tokenizer_source = ctx
Expand All @@ -96,37 +76,20 @@ impl ChatPreparationStage {
));
}

match multimodal::process_multimodal(
&request.messages,
let placeholder = multimodal::resolve_placeholder_token(
model_id,
&*tokenizer,
token_ids,
mm_components,
&tokenizer_source,
)
.await
{
Ok(output) => {
debug!(
function = "ChatPreparationStage::execute",
expanded_tokens = output.expanded_token_ids.len(),
"Multimodal processing complete"
);
token_ids = output.expanded_token_ids;
multimodal_intermediate = Some(output.intermediate);
}
Err(e) => {
error!(
function = "ChatPreparationStage::execute",
error = %e,
"Multimodal processing failed"
);
return Err(error::bad_request(
"multimodal_processing_failed",
format!("Multimodal processing failed: {e}"),
));
}
}
.ok()
.flatten();
Comment thread
CatherineSue marked this conversation as resolved.

(
placeholder,
Some((mm_components, model_id, tokenizer_source)),
)
} else {
error!(
function = "ChatPreparationStage::execute",
Expand All @@ -137,6 +100,71 @@ impl ChatPreparationStage {
"Multimodal content detected but multimodal processing is not available",
));
}
} else {
(None, None)
};

// Step 2: Process messages and apply chat template
let processed_messages = match utils::process_chat_messages(
&body_ref,
&*tokenizer,
image_placeholder.as_deref(),
) {
Ok(msgs) => msgs,
Err(e) => {
error!(function = "ChatPreparationStage::execute", error = %e, "Failed to process chat messages");
return Err(error::bad_request("process_messages_failed", e));
}
};

// Step 3: Tokenize the processed text (no special tokens - chat template already handles them)
let encoding = match tokenizer.encode(&processed_messages.text, false) {
Ok(encoding) => encoding,
Err(e) => {
error!(function = "ChatPreparationStage::execute", error = %e, "Tokenization failed");
return Err(error::internal_error(
"tokenization_failed",
format!("Tokenization failed: {e}"),
));
}
};

let mut token_ids = encoding.token_ids().to_vec();

// Step 4: Full multimodal processing (fetch + preprocess + expand tokens + hash)
let mut multimodal_intermediate = None;
if let Some((mm_components, model_id, tokenizer_source)) = mm_context {
match multimodal::process_multimodal(
&request.messages,
model_id,
&*tokenizer,
token_ids,
mm_components,
&tokenizer_source,
)
.await
{
Ok(output) => {
debug!(
function = "ChatPreparationStage::execute",
expanded_tokens = output.expanded_token_ids.len(),
"Multimodal processing complete"
);
token_ids = output.expanded_token_ids;
multimodal_intermediate = Some(output.intermediate);
}
Err(e) => {
error!(
function = "ChatPreparationStage::execute",
error = %e,
"Multimodal processing failed"
);
return Err(error::bad_request(
"multimodal_processing_failed",
format!("Multimodal processing failed: {e}"),
));
}
}
}

// Step 4: Build tool constraints if needed
Expand Down
Loading
Loading