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
10 changes: 7 additions & 3 deletions rust/src/chat/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -34,8 +34,8 @@ pub use parser::{ParserSelection, validate_parser_overrides};
pub use renderer::hf::ChatTemplateContentFormatOption;
pub use renderer::{
ChatRenderer, DeepSeekV4ChatRenderer, DeepSeekV32ChatRenderer, DeepSeekV41ChatRenderer,
DynChatRenderer, HarmonyChatRenderer, InklingChatRenderer, KimiK3ChatRenderer, RenderedPrompt,
RendererSelection,
DynChatRenderer, HarmonyChatRenderer, InklingChatRenderer, KimiK3ChatRenderer, MediaPartSource,
RenderedPrompt, RendererSelection,
};
pub use request::{
ChatContent, ChatContentPart, ChatMessage, ChatOptions, ChatRequest, ChatRole, ChatTool,
Expand Down Expand Up @@ -112,6 +112,10 @@ impl ChatRequestProcessor {
request: &ChatRequest,
rendered: RenderedPrompt,
) -> Result<(Prompt, Option<MmFeatures>)> {
let media_is_empty = match &rendered.media_order {
Some(media_order) => media_order.is_empty(),
None => !request.has_multimodal(),
};
match self.model_dtype {
Some(model_dtype) => {
multimodal::finalize_rendered_prompt(
Expand All @@ -122,7 +126,7 @@ impl ChatRequestProcessor {
)
.await
}
None if !request.has_multimodal() => Ok((rendered.prompt, None)),
None if media_is_empty => Ok((rendered.prompt, None)),
None => Err(Error::UnsupportedMultimodalRenderer),
}
}
Expand Down
205 changes: 152 additions & 53 deletions rust/src/chat/src/multimodal.rs
Original file line number Diff line number Diff line change
Expand Up @@ -16,7 +16,7 @@ use std::fs;
use std::path::Path;
use std::sync::{Arc, LazyLock};

use itertools::izip;
use itertools::{Either, izip};
use llm_multimodal::{
AsyncMultiModalTracker, AudioClip, AudioPreProcessor, EncoderFieldLayouts, FieldLayout,
ImageFrame, MediaConnector, MediaConnectorConfig, MediaContentPart, Modality, ModelMetadata,
Expand All @@ -33,7 +33,7 @@ use vllm_text::Prompt;
use vllm_text::tokenizer::{DynTokenizer, Tokenizer};

use crate::error::{Error, Result, bail_multimodal, multimodal};
use crate::renderer::RenderedPrompt;
use crate::renderer::{MediaPartSource, RenderedPrompt};
use crate::request::{ChatContent, ChatContentPart, ChatMessage, ChatRequest};

mod audio;
Expand Down Expand Up @@ -549,7 +549,8 @@ pub(crate) async fn finalize_rendered_prompt(
info: Option<&MultimodalModelInfo>,
model_dtype: ModelDtype,
) -> Result<(Prompt, Option<MmFeatures>)> {
if !request.has_multimodal() {
let media_parts = extract_media_parts(request, rendered.media_order.as_deref())?;
if media_parts.is_empty() {
return Ok((rendered.prompt, None));
}
let info = info.ok_or(Error::UnsupportedMultimodalRenderer)?;
Expand All @@ -561,63 +562,92 @@ pub(crate) async fn finalize_rendered_prompt(
.map_err(|error| multimodal!("{error}"))?,
Prompt::TokenIds(token_ids) => token_ids,
};
let media_parts = extract_media_parts(request)?;
let prepared = info.prepare_multimodal(media_parts, &mut prompt_token_ids, model_dtype).await?;

Ok((Prompt::TokenIds(prompt_token_ids), Some(prepared)))
}

/// Extract media parts from chat messages in message/content order.
///
/// Assistant history is skipped because generated assistant blocks are already
/// represented as text for prompt rendering in this crate.
fn extract_media_parts(request: &ChatRequest) -> Result<Vec<MediaContentPart>> {
let mut all_parts = Vec::new();
for message in &request.messages {
let content = match message {
ChatMessage::System { content }
| ChatMessage::Developer { content, .. }
| ChatMessage::User { content }
| ChatMessage::ToolResponse { content, .. } => content,
ChatMessage::Assistant { .. } => continue,
};
let ChatContent::Parts(parts) = content else {
continue;
};
for part in parts {
match part {
ChatContentPart::Text { .. } => {}
ChatContentPart::ImageUrl {
image_url,
detail,
uuid,
} => all_parts.push(MediaContentPart::ImageUrl {
url: image_url.clone(),
detail: *detail,
uuid: uuid.clone(),
}),
ChatContentPart::VideoUrl { video_url, uuid } => {
all_parts.push(MediaContentPart::VideoUrl {
url: video_url.clone(),
uuid: uuid.clone(),
})
}
ChatContentPart::InputAudio { data, format, uuid } => {
all_parts.push(MediaContentPart::AudioUrl {
url: input_audio_data_url(data, format.as_deref())?,
uuid: uuid.clone(),
})
}
ChatContentPart::AudioUrl { audio_url, uuid } => {
all_parts.push(MediaContentPart::AudioUrl {
url: audio_url.clone(),
uuid: uuid.clone(),
})
/// Resolve media parts in the placeholder order reported by the renderer.
fn extract_media_parts(
request: &ChatRequest,
media_order: Option<&[MediaPartSource]>,
) -> Result<Vec<MediaContentPart>> {
let parts = match media_order {
// Resolve exactly the sources selected by the renderer, in placeholder order.
Some(media_order) => Either::Left(media_order.iter().map(|source| -> Result<_> {
let message = request.messages.get(source.message_index).ok_or_else(|| {
multimodal!(
"renderer reported missing chat message {} for multimodal content",
source.message_index
)
})?;
let content = match message {
ChatMessage::System { content }
| ChatMessage::Developer { content, .. }
| ChatMessage::User { content }
| ChatMessage::ToolResponse { content, .. } => content,
ChatMessage::Assistant { .. } => {
bail_multimodal!("renderer reported multimodal assistant content")
}
};
let ChatContent::Parts(parts) = content else {
bail_multimodal!("renderer reported multimodal content in a text-only chat message")
};
let part = parts.get(source.content_part_index).ok_or_else(|| {
multimodal!(
"renderer reported missing content part {} in chat message {}",
source.content_part_index,
source.message_index
)
})?;
Ok(part)
})),
// Fall back to all media in request message and content-part order (e.g. Jinja).
None => Either::Right(request.messages.iter().flat_map(|message| {
let content = match message {
ChatMessage::System { content }
| ChatMessage::Developer { content, .. }
| ChatMessage::User { content }
| ChatMessage::ToolResponse { content, .. } => Some(content),
ChatMessage::Assistant { .. } => None,
};
match content {
Some(ChatContent::Parts(parts)) => parts.as_slice(),
_ => &[],
}
}
}
Ok(all_parts)
.iter()
.filter(|part| part.is_multimodal())
.map(Ok)
})),
};
parts
.map(|part| match part? {
ChatContentPart::Text { .. } => {
bail_multimodal!("renderer reported a text part as multimodal content")
}
ChatContentPart::ImageUrl {
image_url,
detail,
uuid,
} => Ok(MediaContentPart::ImageUrl {
url: image_url.clone(),
detail: *detail,
uuid: uuid.clone(),
}),
ChatContentPart::VideoUrl { video_url, uuid } => Ok(MediaContentPart::VideoUrl {
url: video_url.clone(),
uuid: uuid.clone(),
}),
ChatContentPart::InputAudio { data, format, uuid } => Ok(MediaContentPart::AudioUrl {
url: input_audio_data_url(data, format.as_deref())?,
uuid: uuid.clone(),
}),
ChatContentPart::AudioUrl { audio_url, uuid } => Ok(MediaContentPart::AudioUrl {
url: audio_url.clone(),
uuid: uuid.clone(),
}),
})
.collect()
}

/// Wrap OpenAI base64 audio in a data URL consumed by `MediaConnector`.
Expand Down Expand Up @@ -1043,6 +1073,75 @@ mod tests {
assert!(input_audio_data_url("AAAA", Some("flac")).is_err());
}

#[test]
fn extract_media_parts_uses_request_order_or_explicit_selection() {
let request = ChatRequest {
messages: vec![
ChatMessage::user(vec![
ChatContentPart::text("before image"),
ChatContentPart::ImageUrl {
image_url: "data:image/png;base64,image-b".to_string(),
detail: None,
uuid: Some("image-b".to_string()),
},
]),
ChatMessage::user(vec![ChatContentPart::ImageUrl {
image_url: "data:image/png;base64,image-a".to_string(),
detail: None,
uuid: Some("image-a".to_string()),
}]),
ChatMessage::assistant_text("answer"),
ChatMessage::user("follow-up"),
],
..ChatRequest::for_test()
};
let media_order = [
MediaPartSource {
message_index: 1,
content_part_index: 0,
},
MediaPartSource {
message_index: 0,
content_part_index: 1,
},
];

let selections = [None, Some(&media_order[..]), Some(&[][..])];
let uuids = selections.map(|order| {
extract_media_parts(&request, order)
.unwrap()
.into_iter()
.map(|part| match part {
MediaContentPart::ImageUrl { uuid, .. } => uuid,
_ => panic!("expected image URL"),
})
.collect::<Vec<_>>()
});

expect_test::expect![[r#"
[
[
Some(
"image-b",
),
Some(
"image-a",
),
],
[
Some(
"image-a",
),
Some(
"image-b",
),
],
[],
]
"#]]
.assert_debug_eq(&uuids);
}

fn image_url_part() -> MediaContentPart {
MediaContentPart::ImageUrl {
url: "https://example.com/image.png".to_string(),
Expand Down
Loading
Loading