merge deploy
This commit is contained in:
@@ -86,7 +86,11 @@ impl GenerateModel for DeepseekOCRGenerateModel {
|
||||
} else {
|
||||
false
|
||||
};
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
||||
let (mut input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
|
||||
.processor
|
||||
.process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?;
|
||||
@@ -130,11 +134,48 @@ impl GenerateModel for DeepseekOCRGenerateModel {
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ Send
|
||||
+ Unpin
|
||||
+ '_,
|
||||
>,
|
||||
> {
|
||||
let base_size = if let Some(map) = &mes.metadata
|
||||
&& map.contains_key("base_size")
|
||||
{
|
||||
let size = map.get("base_size").unwrap();
|
||||
let size = size.parse::<u32>().unwrap_or(640);
|
||||
if self.size.contains(&size) { size } else { 640 }
|
||||
} else {
|
||||
640
|
||||
};
|
||||
let image_size = if let Some(map) = &mes.metadata
|
||||
&& map.contains_key("image_size")
|
||||
{
|
||||
let size = map.get("image_size").unwrap();
|
||||
let size = size.parse::<u32>().unwrap_or(640);
|
||||
if self.size.contains(&size) { size } else { 640 }
|
||||
} else {
|
||||
640
|
||||
};
|
||||
let crop_mode = if let Some(map) = &mes.metadata
|
||||
&& map.contains_key("crop_mode")
|
||||
{
|
||||
let size = map.get("crop_mode").unwrap();
|
||||
size.parse::<bool>().unwrap_or(false)
|
||||
} else {
|
||||
false
|
||||
};
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
||||
let (mut input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
|
||||
.processor
|
||||
.process_info(&mes, &self.tokenizer, 640, 640, true)?;
|
||||
.process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?;
|
||||
|
||||
let mut seqlen_offset = 0;
|
||||
let mut seq_len = input_ids.dim(1)?;
|
||||
@@ -192,6 +233,6 @@ impl GenerateModel for DeepseekOCRGenerateModel {
|
||||
}
|
||||
self.deepseekocr_model.clear_kv_cache();
|
||||
};
|
||||
Ok(stream)
|
||||
Ok(Box::new(Box::pin(stream)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -610,7 +610,7 @@ impl CLIPVisionEmbeddings {
|
||||
}
|
||||
|
||||
fn get_abs_pos(&self, tgt_size: usize) -> Result<Tensor> {
|
||||
println!("self.pos_embeds: {:?}", self.pos_embeds);
|
||||
// println!("self.pos_embeds: {:?}", self.pos_embeds);
|
||||
let abs_pos_new = self.pos_embeds.clone();
|
||||
let (len, dim) = abs_pos_new.dims2()?;
|
||||
let src_size = ((len - 1) as f32).sqrt() as usize;
|
||||
|
||||
@@ -53,7 +53,11 @@ impl<'a> MiniCPMGenerateModel<'a> {
|
||||
|
||||
impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||
let mut seq_len = input_ids.dim(1)?;
|
||||
@@ -80,8 +84,19 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ Send
|
||||
+ Unpin
|
||||
+ '_,
|
||||
>,
|
||||
> {
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
let mut input_ids = self.tokenizer.text_encode(mes_render, &self.device)?;
|
||||
let mut seq_len = input_ids.dim(1)?;
|
||||
@@ -125,6 +140,6 @@ impl<'a> GenerateModel for MiniCPMGenerateModel<'a> {
|
||||
}
|
||||
self.minicpm.clear_kv_cache();
|
||||
};
|
||||
Ok(stream)
|
||||
Ok(Box::new(Box::pin(stream)))
|
||||
}
|
||||
}
|
||||
|
||||
+107
-3
@@ -11,12 +11,116 @@ use aha_openai_dive::v1::resources::chat::{
|
||||
use anyhow::Result;
|
||||
use rocket::futures::Stream;
|
||||
|
||||
use crate::models::{
|
||||
deepseek_ocr::generate::DeepseekOCRGenerateModel, minicpm4::generate::MiniCPMGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel, qwen3vl::generate::Qwen3VLGenerateModel
|
||||
};
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)]
|
||||
pub enum WhichModel {
|
||||
#[value(name = "minicpm4-0.5b")]
|
||||
MiniCPM4_0_5B,
|
||||
#[value(name = "qwen2.5vl-3b")]
|
||||
Qwen2_5vl3B,
|
||||
#[value(name = "qwen2.5vl-7b")]
|
||||
Qwen2_5vl7B,
|
||||
#[value(name = "qwen3vl-2b")]
|
||||
Qwen3vl2B,
|
||||
#[value(name = "qwen3vl-4b")]
|
||||
Qwen3vl4B,
|
||||
#[value(name = "qwen3vl-8b")]
|
||||
Qwen3vl8B,
|
||||
#[value(name = "qwen3vl-32b")]
|
||||
Qwen3vl32B,
|
||||
#[value(name = "deepseek-ocr")]
|
||||
DeepSeekOCR,
|
||||
}
|
||||
|
||||
pub trait GenerateModel {
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse>;
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>>
|
||||
where
|
||||
Self: Sized;
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ Send
|
||||
+ Unpin
|
||||
+ '_,
|
||||
>,
|
||||
>;
|
||||
}
|
||||
|
||||
pub enum ModelInstance<'a> {
|
||||
MiniCPM4(MiniCPMGenerateModel<'a>),
|
||||
Qwen2_5VL(Qwen2_5VLGenerateModel<'a>),
|
||||
Qwen3VL(Qwen3VLGenerateModel<'a>),
|
||||
DeepSeekOCR(DeepseekOCRGenerateModel),
|
||||
}
|
||||
|
||||
impl<'a> GenerateModel for ModelInstance<'a> {
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||
match self {
|
||||
ModelInstance::MiniCPM4(model) => model.generate(mes),
|
||||
ModelInstance::Qwen2_5VL(model) => model.generate(mes),
|
||||
ModelInstance::Qwen3VL(model) => model.generate(mes),
|
||||
ModelInstance::DeepSeekOCR(model) => model.generate(mes),
|
||||
}
|
||||
}
|
||||
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ Send
|
||||
+ Unpin
|
||||
+ '_,
|
||||
>,
|
||||
> {
|
||||
match self {
|
||||
ModelInstance::MiniCPM4(model) => model.generate_stream(mes),
|
||||
ModelInstance::Qwen2_5VL(model) => model.generate_stream(mes),
|
||||
ModelInstance::Qwen3VL(model) => model.generate_stream(mes),
|
||||
ModelInstance::DeepSeekOCR(model) => model.generate_stream(mes),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn load_model(model_type: WhichModel, path: &str) -> Result<ModelInstance<'_>> {
|
||||
let model = match model_type {
|
||||
WhichModel::MiniCPM4_0_5B => {
|
||||
let model = MiniCPMGenerateModel::init(path, None, None)?;
|
||||
ModelInstance::MiniCPM4(model)
|
||||
}
|
||||
WhichModel::Qwen2_5vl3B => {
|
||||
let model = Qwen2_5VLGenerateModel::init(path, None, None)?;
|
||||
ModelInstance::Qwen2_5VL(model)
|
||||
}
|
||||
WhichModel::Qwen2_5vl7B => {
|
||||
let model = Qwen2_5VLGenerateModel::init(path, None, None)?;
|
||||
ModelInstance::Qwen2_5VL(model)
|
||||
}
|
||||
WhichModel::Qwen3vl2B => {
|
||||
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
||||
ModelInstance::Qwen3VL(model)
|
||||
}
|
||||
WhichModel::Qwen3vl4B => {
|
||||
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
||||
ModelInstance::Qwen3VL(model)
|
||||
}
|
||||
WhichModel::Qwen3vl8B => {
|
||||
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
||||
ModelInstance::Qwen3VL(model)
|
||||
}
|
||||
WhichModel::Qwen3vl32B => {
|
||||
let model = Qwen3VLGenerateModel::init(path, None, None)?;
|
||||
ModelInstance::Qwen3VL(model)
|
||||
}
|
||||
WhichModel::DeepSeekOCR => {
|
||||
let model = DeepseekOCRGenerateModel::init(path, None, None)?;
|
||||
ModelInstance::DeepSeekOCR(model)
|
||||
}
|
||||
};
|
||||
Ok(model)
|
||||
}
|
||||
|
||||
@@ -63,7 +63,11 @@ impl<'a> Qwen2_5VLGenerateModel<'a> {
|
||||
|
||||
impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||
let mut input_ids = self
|
||||
@@ -122,8 +126,19 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ Send
|
||||
+ Unpin
|
||||
+ '_,
|
||||
>,
|
||||
> {
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None, seed);
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||
let mut input_ids = self
|
||||
@@ -203,6 +218,6 @@ impl<'a> GenerateModel for Qwen2_5VLGenerateModel<'a> {
|
||||
}
|
||||
self.qwen2_5_vl.clear_kv_cache();
|
||||
};
|
||||
Ok(stream)
|
||||
Ok(Box::new(Box::pin(stream)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -42,7 +42,6 @@ pub struct Qwen3VLTextConfig {
|
||||
pub rms_norm_eps: f64,
|
||||
pub rope_scaling: RopeScaling,
|
||||
pub rope_theta: f32,
|
||||
pub tie_word_embeddings: bool,
|
||||
pub use_cache: bool,
|
||||
pub vocab_size: usize,
|
||||
}
|
||||
|
||||
@@ -47,7 +47,6 @@ impl<'a> Qwen3VLGenerateModel<'a> {
|
||||
let pre_processor = Qwen3VLProcessor::new(path, &device, dtype)?;
|
||||
let model_list = find_type_files(path, "safetensors")?;
|
||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, &device)? };
|
||||
let vb = vb.pp("model");
|
||||
let qwen3_vl = Qwen3VLModel::new(cfg, vb)?;
|
||||
let generation_config_path = path.to_string() + "/generation_config.json";
|
||||
let generation_config: Qwen3VLGenerationConfig =
|
||||
@@ -76,7 +75,12 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
||||
Some(top_p) => top_p,
|
||||
};
|
||||
let top_k = self.generation_config.top_k;
|
||||
let mut logit_processor = get_logit_processor(Some(temperature), Some(top_p), Some(top_k));
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor =
|
||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||
let mut input_ids = self
|
||||
@@ -123,7 +127,14 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
||||
fn generate_stream(
|
||||
&mut self,
|
||||
mes: ChatCompletionParameters,
|
||||
) -> Result<impl Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>> {
|
||||
) -> Result<
|
||||
Box<
|
||||
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
|
||||
+ Send
|
||||
+ Unpin
|
||||
+ '_,
|
||||
>,
|
||||
> {
|
||||
let temperature = match mes.temperature {
|
||||
None => self.generation_config.temperature,
|
||||
Some(tem) => tem,
|
||||
@@ -133,7 +144,12 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
||||
Some(top_p) => top_p,
|
||||
};
|
||||
let top_k = self.generation_config.top_k;
|
||||
let mut logit_processor = get_logit_processor(Some(temperature), Some(top_p), Some(top_k));
|
||||
let seed = match mes.seed {
|
||||
None => 34562u64,
|
||||
Some(s) => s as u64,
|
||||
};
|
||||
let mut logit_processor =
|
||||
get_logit_processor(Some(temperature), Some(top_p), Some(top_k), seed);
|
||||
let mes_render = self.chat_template.apply_chat_template(&mes)?;
|
||||
let input = self.pre_processor.process_info(&mes, &mes_render)?;
|
||||
let mut input_ids = self
|
||||
@@ -199,6 +215,6 @@ impl<'a> GenerateModel for Qwen3VLGenerateModel<'a> {
|
||||
}
|
||||
self.qwen3_vl.clear_kv_cache();
|
||||
};
|
||||
Ok(stream)
|
||||
Ok(Box::new(Box::pin(stream)))
|
||||
}
|
||||
}
|
||||
|
||||
@@ -559,7 +559,6 @@ pub struct Qwen3VLTextAttention {
|
||||
num_key_value_heads: usize,
|
||||
num_kv_groups: usize,
|
||||
head_dim: usize,
|
||||
hidden_size: usize,
|
||||
scaling: f64,
|
||||
kv_cache: Option<(Tensor, Tensor)>,
|
||||
}
|
||||
@@ -568,7 +567,7 @@ impl Qwen3VLTextAttention {
|
||||
pub fn new(config: Qwen3VLTextConfig, vb: VarBuilder) -> Result<Self> {
|
||||
let hidden_size = config.hidden_size;
|
||||
let num_attention_heads = config.num_attention_heads;
|
||||
let head_dim = hidden_size / num_attention_heads;
|
||||
let head_dim = config.head_dim;
|
||||
let num_key_value_heads = config.num_key_value_heads;
|
||||
let num_kv_groups = num_attention_heads / num_key_value_heads;
|
||||
let scaling = 1f64 / f64::sqrt(head_dim as f64);
|
||||
@@ -576,7 +575,7 @@ impl Qwen3VLTextAttention {
|
||||
let q_proj = linear(hidden_size, num_attention_heads * head_dim, vb.pp("q_proj"))?;
|
||||
let k_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
|
||||
let v_proj = linear(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
|
||||
let o_proj = linear(hidden_size, hidden_size, vb.pp("o_proj"))?;
|
||||
let o_proj = linear(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
|
||||
(q_proj, k_proj, v_proj, o_proj)
|
||||
} else {
|
||||
let q_proj =
|
||||
@@ -585,7 +584,8 @@ impl Qwen3VLTextAttention {
|
||||
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("k_proj"))?;
|
||||
let v_proj =
|
||||
linear_no_bias(hidden_size, num_key_value_heads * head_dim, vb.pp("v_proj"))?;
|
||||
let o_proj = linear_no_bias(hidden_size, hidden_size, vb.pp("o_proj"))?;
|
||||
let o_proj =
|
||||
linear_no_bias(num_attention_heads * head_dim, hidden_size, vb.pp("o_proj"))?;
|
||||
(q_proj, k_proj, v_proj, o_proj)
|
||||
};
|
||||
let q_norm = rms_norm(head_dim, config.rms_norm_eps, vb.pp("q_norm"))?;
|
||||
@@ -601,7 +601,6 @@ impl Qwen3VLTextAttention {
|
||||
num_key_value_heads,
|
||||
num_kv_groups,
|
||||
head_dim,
|
||||
hidden_size,
|
||||
scaling,
|
||||
kv_cache: None,
|
||||
})
|
||||
@@ -652,7 +651,8 @@ impl Qwen3VLTextAttention {
|
||||
attention_mask,
|
||||
self.scaling,
|
||||
)?;
|
||||
let attn_output = attn_output.reshape((b_sz, q_len, self.hidden_size))?;
|
||||
let attn_output =
|
||||
attn_output.reshape((b_sz, q_len, self.num_attention_heads * self.head_dim))?;
|
||||
let attn_output = attn_output.apply(&self.o_proj)?;
|
||||
Ok(attn_output)
|
||||
}
|
||||
@@ -738,7 +738,7 @@ impl Qwen3VLTextModel {
|
||||
layers.push(layer)
|
||||
}
|
||||
let norm = rms_norm(config.hidden_size, config.rms_norm_eps, vb.pp("norm"))?;
|
||||
let head_dim = config.hidden_size / config.num_attention_heads;
|
||||
let head_dim = config.head_dim;
|
||||
let rotary_emb = Qwen3VLTextRotaryEmbedding::new(head_dim, config.rope_theta);
|
||||
let mrope_section = config.rope_scaling.mrope_section.clone();
|
||||
Ok(Self {
|
||||
@@ -823,10 +823,11 @@ pub struct Qwen3VLModel {
|
||||
|
||||
impl Qwen3VLModel {
|
||||
pub fn new(config: Qwen3VLConfig, vb: VarBuilder) -> Result<Self> {
|
||||
let vb_m = vb.pp("model");
|
||||
let config = config.clone();
|
||||
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb.pp("visual"))?;
|
||||
let visual = Qwen3VLVisionModel::new(config.vision_config.clone(), vb_m.pp("visual"))?;
|
||||
let language_model =
|
||||
Qwen3VLTextModel::new(config.text_config.clone(), vb.pp("language_model"))?;
|
||||
Qwen3VLTextModel::new(config.text_config.clone(), vb_m.pp("language_model"))?;
|
||||
let lm_head = if config.tie_word_embeddings {
|
||||
Linear::new(language_model.embed_tokens.embeddings().clone(), None)
|
||||
} else {
|
||||
|
||||
Reference in New Issue
Block a user