From 35c75506e2fd94e2285d8a3cb66208d518d5f992 Mon Sep 17 00:00:00 2001 From: sigoden Date: Mon, 27 Nov 2023 14:04:50 +0800 Subject: feat: support vision (#249) * feat: support vision * clippy * implement vision * resolve data url to local file * add model openai:gpt-4-vision-preview * use newline to concate embeded text files * set max_tokens for gpt-4-vision-preview --- src/client/common.rs | 19 +++++++------- src/client/ernie.rs | 8 +++--- src/client/message.rs | 69 ++++++++++++++++++++++++++++++++++++++++++++++++--- src/client/model.rs | 12 +++++++-- src/client/openai.rs | 11 ++++++-- src/client/palm.rs | 8 +++--- 6 files changed, 104 insertions(+), 23 deletions(-) (limited to 'src/client') diff --git a/src/client/common.rs b/src/client/common.rs index cf5ba9b..2716d87 100644 --- a/src/client/common.rs +++ b/src/client/common.rs @@ -1,7 +1,7 @@ use super::{openai::OpenAIConfig, ClientConfig, Message}; use crate::{ - config::GlobalConfig, + config::{GlobalConfig, Input}, render::ReplyHandler, utils::{ init_tokio_runtime, prompt_input_integer, prompt_input_string, tokenize, AbortSignal, @@ -50,7 +50,7 @@ macro_rules! register_client { } impl $client { - pub const NAME: &str = $name; + pub const NAME: &'static str = $name; pub fn init(global_config: &$crate::config::GlobalConfig) -> Option> { let model = global_config.read().model.clone(); @@ -186,22 +186,22 @@ pub trait Client { Ok(client) } - fn send_message(&self, content: &str) -> Result { + fn send_message(&self, input: Input) -> Result { init_tokio_runtime()?.block_on(async { let global_config = self.config().0; if global_config.read().dry_run { - let content = global_config.read().echo_messages(content); + let content = global_config.read().echo_messages(&input); return Ok(content); } let client = self.build_client()?; - let data = global_config.read().prepare_send_data(content, false)?; + let data = global_config.read().prepare_send_data(&input, false)?; self.send_message_inner(&client, data) .await .with_context(|| "Failed to get answer") }) } - fn send_message_streaming(&self, content: &str, handler: &mut ReplyHandler) -> Result<()> { + fn send_message_streaming(&self, input: &Input, handler: &mut ReplyHandler) -> Result<()> { async fn watch_abort(abort: AbortSignal) { loop { if abort.aborted() { @@ -211,12 +211,13 @@ pub trait Client { } } let abort = handler.get_abort(); - init_tokio_runtime()?.block_on(async { + let input = input.clone(); + init_tokio_runtime()?.block_on(async move { tokio::select! { ret = async { let global_config = self.config().0; if global_config.read().dry_run { - let content = global_config.read().echo_messages(content); + let content = global_config.read().echo_messages(&input); let tokens = tokenize(&content); for token in tokens { tokio::time::sleep(Duration::from_millis(10)).await; @@ -225,7 +226,7 @@ pub trait Client { return Ok(()); } let client = self.build_client()?; - let data = global_config.read().prepare_send_data(content, true)?; + let data = global_config.read().prepare_send_data(&input, true)?; self.send_message_streaming_inner(&client, handler, data).await } => { handler.done()?; diff --git a/src/client/ernie.rs b/src/client/ernie.rs index 200433c..4bb3435 100644 --- a/src/client/ernie.rs +++ b/src/client/ernie.rs @@ -1,4 +1,4 @@ -use super::{ErnieClient, Client, ExtraConfig, PromptType, SendData, Model}; +use super::{ErnieClient, Client, ExtraConfig, PromptType, SendData, Model, MessageContent}; use crate::{ config::GlobalConfig, @@ -198,8 +198,10 @@ fn build_body(data: SendData, _model: String) -> Value { if messages[0].role.is_system() { let system_message = messages.remove(0); - if let Some(message) = messages.get_mut(0) { - message.content = format!("{}\n\n{}", system_message.content, message.content) + if let (Some(message), MessageContent::Text(system_text)) = (messages.get_mut(0), system_message.content) { + if let MessageContent::Text(text) = message.content.clone() { + message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text)) + } } } diff --git a/src/client/message.rs b/src/client/message.rs index 55b2663..dc8c3e1 100644 --- a/src/client/message.rs +++ b/src/client/message.rs @@ -1,16 +1,18 @@ +use crate::config::Input; + use serde::{Deserialize, Serialize}; #[derive(Debug, Clone, Deserialize, Serialize)] pub struct Message { pub role: MessageRole, - pub content: String, + pub content: MessageContent, } impl Message { - pub fn new(content: &str) -> Self { + pub fn new(input: &Input) -> Self { Self { role: MessageRole::User, - content: content.to_string(), + content: input.to_message_content(), } } } @@ -38,6 +40,65 @@ impl MessageRole { } } +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(untagged)] +pub enum MessageContent { + Text(String), + Array(Vec), +} + +impl MessageContent { + pub fn render_input(&self, resolve_url_fn: impl Fn(&str) -> String) -> String { + match self { + MessageContent::Text(text) => text.to_string(), + MessageContent::Array(list) => { + let (mut concated_text, mut files) = (String::new(), vec![]); + for item in list { + match item { + MessageContentPart::Text { text } => { + concated_text = format!("{concated_text} {text}") + } + MessageContentPart::ImageUrl { image_url } => { + files.push(resolve_url_fn(&image_url.url)) + } + } + } + if !concated_text.is_empty() { + concated_text = format!(" -- {concated_text}") + } + format!(".file {}{}", files.join(" "), concated_text) + } + } + } + + pub fn merge_prompt(&mut self, replace_fn: impl Fn(&str) -> String) { + match self { + MessageContent::Text(text) => *text = replace_fn(text), + MessageContent::Array(list) => { + if list.is_empty() { + list.push(MessageContentPart::Text { + text: replace_fn(""), + }) + } else if let Some(MessageContentPart::Text { text }) = list.get_mut(0) { + *text = replace_fn(text) + } + } + } + } +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum MessageContentPart { + Text { text: String }, + ImageUrl { image_url: ImageUrl }, +} + +#[derive(Debug, Clone, Deserialize, Serialize)] +pub struct ImageUrl { + pub url: String, +} + #[cfg(test)] mod tests { use super::*; @@ -45,7 +106,7 @@ mod tests { #[test] fn test_serde() { assert_eq!( - serde_json::to_string(&Message::new("Hello World")).unwrap(), + serde_json::to_string(&Message::new(&Input::from_str("Hello World"))).unwrap(), "{\"role\":\"user\",\"content\":\"Hello World\"}" ); } diff --git a/src/client/model.rs b/src/client/model.rs index 16fe087..130489d 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -1,4 +1,4 @@ -use super::message::Message; +use super::message::{Message, MessageContent}; use crate::utils::count_tokens; @@ -79,7 +79,15 @@ impl Model { } pub fn messages_tokens(&self, messages: &[Message]) -> usize { - messages.iter().map(|v| count_tokens(&v.content)).sum() + messages + .iter() + .map(|v| { + match &v.content { + MessageContent::Text(text) => count_tokens(text), + MessageContent::Array(_) => 0, // TODO + } + }) + .sum() } pub fn total_tokens(&self, messages: &[Message]) -> usize { diff --git a/src/client/openai.rs b/src/client/openai.rs index dbe2661..928c6b3 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -19,13 +19,14 @@ use std::env; const API_BASE: &str = "https://api.openai.com/v1"; -const MODELS: [(&str, usize); 6] = [ +const MODELS: [(&str, usize); 7] = [ ("gpt-3.5-turbo", 4096), ("gpt-3.5-turbo-16k", 16385), ("gpt-3.5-turbo-1106", 16385), + ("gpt-4-1106-preview", 128000), + ("gpt-4-vision-preview", 128000), ("gpt-4", 8192), ("gpt-4-32k", 32768), - ("gpt-4-1106-preview", 128000), ]; pub const OPENAI_TOKENS_COUNT_FACTORS: TokensCountFactors = (5, 2); @@ -145,6 +146,12 @@ pub fn openai_build_body(data: SendData, model: String) -> Value { "model": model, "messages": messages, }); + + // The default max_tokens of gpt-4-vision-preview is only 16, we need to make it larger + if model == "gpt-4-vision-preview" { + body["max_tokens"] = json!(4096); + } + if let Some(v) = temperature { body["temperature"] = v.into(); } diff --git a/src/client/palm.rs b/src/client/palm.rs index a2aec8a..37ed0ae 100644 --- a/src/client/palm.rs +++ b/src/client/palm.rs @@ -1,4 +1,4 @@ -use super::{PaLMClient, Client, ExtraConfig, Model, PromptType, SendData, TokensCountFactors, send_message_as_streaming}; +use super::{PaLMClient, Client, ExtraConfig, Model, PromptType, SendData, TokensCountFactors, send_message_as_streaming, MessageContent}; use crate::{config::GlobalConfig, render::ReplyHandler, utils::PromptKind}; @@ -115,8 +115,10 @@ fn build_body(data: SendData, _model: String) -> Value { if messages[0].role.is_system() { let system_message = messages.remove(0); - if let Some(message) = messages.get_mut(0) { - message.content = format!("{}\n\n{}", system_message.content, message.content) + if let (Some(message), MessageContent::Text(system_text)) = (messages.get_mut(0), system_message.content) { + if let MessageContent::Text(text) = message.content.clone() { + message.content = MessageContent::Text(format!("{}\n\n{}", system_text, text)) + } } } -- cgit v1.2.3