add deepseek_ocr
This commit is contained in:
@@ -286,6 +286,9 @@ pub fn eager_attention_forward(
|
||||
Some(g) => repeat_kv(value_states.clone(), g)?.contiguous()?,
|
||||
None => value_states.clone(),
|
||||
};
|
||||
let query_states = query_states.contiguous()?;
|
||||
let key_states = key_states.contiguous()?;
|
||||
let value_states = value_states.contiguous()?;
|
||||
let attn_output = {
|
||||
#[cfg(not(feature = "flash-attn"))]
|
||||
{
|
||||
|
||||
@@ -1,15 +1,26 @@
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
pub struct DeepseekV2Config {
|
||||
pub bos_token_id: u32,
|
||||
pub eos_token_id: u32,
|
||||
pub first_k_dense_replace: u32,
|
||||
pub first_k_dense_replace: usize,
|
||||
pub hidden_size: usize,
|
||||
pub intermediate_size: usize,
|
||||
pub kv_lora_rank: Option<usize>,
|
||||
pub lm_head: bool,
|
||||
pub max_position_embeddings: usize,
|
||||
pub moe_intermediate_size: usize,
|
||||
#[serde(default = "default_moe_layer_freq")]
|
||||
pub moe_layer_freq: usize,
|
||||
#[serde(default = "default_routed_scaling_factor")]
|
||||
pub routed_scaling_factor: f64,
|
||||
#[serde(default = "default_scoring_func")]
|
||||
pub scoring_func: String,
|
||||
#[serde(default = "default_aux_loss_alpha")]
|
||||
pub aux_loss_alpha: f32,
|
||||
#[serde(default = "default_true")]
|
||||
pub seq_aux: bool,
|
||||
#[serde(default = "default_false")]
|
||||
pub norm_topk_prob: bool,
|
||||
pub n_group: usize,
|
||||
pub n_routed_experts: usize,
|
||||
pub n_shared_experts: usize,
|
||||
@@ -27,6 +38,30 @@ pub struct DeepseekV2Config {
|
||||
pub use_mla: bool,
|
||||
pub v_head_dim: usize,
|
||||
pub vocab_size: usize,
|
||||
#[serde(default = "default_rms_norm_eps")]
|
||||
pub rms_norm_eps: f64,
|
||||
}
|
||||
|
||||
fn default_moe_layer_freq() -> usize {
|
||||
1
|
||||
}
|
||||
fn default_routed_scaling_factor() -> f64 {
|
||||
1.0
|
||||
}
|
||||
fn default_scoring_func() -> String {
|
||||
"softmax".to_string()
|
||||
}
|
||||
fn default_aux_loss_alpha() -> f32 {
|
||||
0.001
|
||||
}
|
||||
fn default_true() -> bool {
|
||||
true
|
||||
}
|
||||
fn default_false() -> bool {
|
||||
false
|
||||
}
|
||||
fn default_rms_norm_eps() -> f64 {
|
||||
1e-6
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
@@ -43,13 +78,13 @@ pub struct ClipL14_224 {
|
||||
pub image_size: usize,
|
||||
pub layers: usize,
|
||||
pub patch_size: usize,
|
||||
pub width: usize
|
||||
pub width: usize,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
pub struct SamVitB {
|
||||
pub downsample_channels: Vec<usize>,
|
||||
pub global_attn_indexes: Vec<u32>,
|
||||
pub global_attn_indexes: Vec<usize>,
|
||||
pub heads: usize,
|
||||
pub layers: usize,
|
||||
pub width: usize,
|
||||
@@ -66,7 +101,7 @@ pub struct Width {
|
||||
pub struct DeepseekOCRVisionConfig {
|
||||
pub image_size: usize,
|
||||
pub mlp_ratio: f32,
|
||||
pub width: Width
|
||||
pub width: Width,
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
|
||||
@@ -100,4 +135,4 @@ pub struct DeepseekOCRConfig {
|
||||
pub use_mla: bool,
|
||||
pub v_head_dim: usize,
|
||||
pub vocab_size: usize,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,38 +1,168 @@
|
||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||
use anyhow::Result;
|
||||
use candle_core::{DType, Device};
|
||||
use aha_openai_dive::v1::resources::chat::{
|
||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||
};
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::{DType, Device, Tensor};
|
||||
use candle_nn::VarBuilder;
|
||||
use rocket::async_stream::stream;
|
||||
use rocket::futures::Stream;
|
||||
|
||||
use crate::{
|
||||
models::deepseek_ocr::{config::DeepseekOCRConfig, processor::DeepseekOCRProcessor},
|
||||
models::{
|
||||
GenerateModel,
|
||||
deepseek_ocr::{
|
||||
config::DeepseekOCRConfig, model::DeepseekOCRModel, processor::DeepseekOCRProcessor,
|
||||
},
|
||||
},
|
||||
tokenizer::TokenizerModel,
|
||||
utils::{get_device, get_dtype},
|
||||
utils::{
|
||||
build_completion_chunk_response, build_completion_response, find_type_files, get_device,
|
||||
get_dtype, get_logit_processor,
|
||||
},
|
||||
};
|
||||
|
||||
pub struct DeepseekOCRGenerateModel {
|
||||
tokenizer: TokenizerModel,
|
||||
processor: DeepseekOCRProcessor,
|
||||
deepseekocr_model: DeepseekOCRModel,
|
||||
bos_token_id: u32,
|
||||
eos_token_id: u32,
|
||||
device: Device,
|
||||
}
|
||||
|
||||
impl DeepseekOCRGenerateModel {
|
||||
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
|
||||
let tokenizer = TokenizerModel::init(path)?;
|
||||
let device = &get_device(device);
|
||||
let dtype = get_dtype(dtype, "bfloat16");
|
||||
let processor = DeepseekOCRProcessor::new(device, dtype)?;
|
||||
let config_path = path.to_string() + "/config.json";
|
||||
let cfg: DeepseekOCRConfig = serde_json::from_slice(&std::fs::read(config_path)?)?;
|
||||
|
||||
let cfg_dtype = cfg.language_config.torch_dtype.clone();
|
||||
let device = &get_device(device);
|
||||
let dtype = get_dtype(dtype, &cfg_dtype);
|
||||
let processor = DeepseekOCRProcessor::new(device, dtype)?;
|
||||
let eos_token_id = cfg.eos_token_id;
|
||||
let bos_token_id = cfg.bos_token_id;
|
||||
let model_list = find_type_files(path, "safetensors")?;
|
||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||
let deepseekocr_model = DeepseekOCRModel::new(vb, cfg)?;
|
||||
Ok(Self {
|
||||
tokenizer,
|
||||
processor,
|
||||
deepseekocr_model,
|
||||
bos_token_id,
|
||||
eos_token_id,
|
||||
device: device.clone(),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub fn generate(&mut self, mes: ChatCompletionParameters) -> Result<()> {
|
||||
let (input_ids, images_ori, image_crop, image_seq_mask, images_spatial_crop_t) = self
|
||||
impl GenerateModel for DeepseekOCRGenerateModel {
|
||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
|
||||
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)?;
|
||||
let mut images_ori = Some(&images_ori);
|
||||
let mut image_crop = Some(&image_crop);
|
||||
let mut images_seq_mask = Some(&images_seq_mask);
|
||||
let mut images_spatial_crop_t = Some(&images_spatial_crop_t);
|
||||
let mut seqlen_offset = 0;
|
||||
let mut seq_len = input_ids.dim(1)?;
|
||||
let mut generate = Vec::new();
|
||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||
for _ in 0..sample_len {
|
||||
let logits = self.deepseekocr_model.forward(
|
||||
&input_ids,
|
||||
images_ori,
|
||||
image_crop,
|
||||
images_seq_mask,
|
||||
images_spatial_crop_t,
|
||||
seqlen_offset,
|
||||
)?;
|
||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||
let next_token = logit_processor.sample(&logits)?;
|
||||
generate.push(next_token);
|
||||
if next_token == self.bos_token_id || next_token == self.eos_token_id {
|
||||
break;
|
||||
}
|
||||
seqlen_offset += seq_len;
|
||||
seq_len = 1;
|
||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||
images_ori = None;
|
||||
image_crop = None;
|
||||
images_seq_mask = None;
|
||||
images_spatial_crop_t = None;
|
||||
}
|
||||
let res = self.tokenizer.token_decode(generate)?;
|
||||
self.deepseekocr_model.clear_kv_cache();
|
||||
let response = build_completion_response(res, "deepseek_ocr");
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
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);
|
||||
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)?;
|
||||
|
||||
Ok(())
|
||||
let mut seqlen_offset = 0;
|
||||
let mut seq_len = input_ids.dim(1)?;
|
||||
let sample_len = mes.max_tokens.unwrap_or(1024);
|
||||
let stream = stream! {
|
||||
let mut error_tokens = Vec::new();
|
||||
let mut images_ori = Some(&images_ori);
|
||||
let mut image_crop = Some(&image_crop);
|
||||
let mut images_seq_mask = Some(&images_seq_mask);
|
||||
let mut images_spatial_crop_t = Some(&images_spatial_crop_t);
|
||||
for _ in 0..sample_len {
|
||||
let logits = self.deepseekocr_model.forward(
|
||||
&input_ids,
|
||||
images_ori,
|
||||
image_crop,
|
||||
images_seq_mask,
|
||||
images_spatial_crop_t,
|
||||
seqlen_offset,
|
||||
)?;
|
||||
let logits = logits.squeeze(0)?.squeeze(0)?.to_dtype(DType::F32)?;
|
||||
let next_token = logit_processor.sample(&logits)?;
|
||||
let mut decode_ids = Vec::new();
|
||||
if !error_tokens.is_empty() {
|
||||
decode_ids.extend_from_slice(&error_tokens);
|
||||
}
|
||||
decode_ids.push(next_token);
|
||||
let decoded_token = self.tokenizer.token_decode(decode_ids).map_err(|e| anyhow!(format!("stream decode error{}", e)))?;
|
||||
if decoded_token.contains("�") {
|
||||
error_tokens.push(next_token);
|
||||
if error_tokens.len() > 3 {
|
||||
error_tokens.clear();
|
||||
}
|
||||
seqlen_offset += seq_len;
|
||||
seq_len = 1;
|
||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||
images_ori = None;
|
||||
image_crop = None;
|
||||
images_seq_mask = None;
|
||||
images_spatial_crop_t = None;
|
||||
continue;
|
||||
}
|
||||
error_tokens.clear();
|
||||
let chunk = build_completion_chunk_response(decoded_token, "deepseek_ocr", None, None);
|
||||
yield Ok(chunk);
|
||||
if next_token == self.bos_token_id || next_token == self.eos_token_id {
|
||||
break;
|
||||
}
|
||||
seqlen_offset += seq_len;
|
||||
seq_len = 1;
|
||||
input_ids = Tensor::from_vec(vec![next_token], (1, 1), &self.device)?;
|
||||
images_ori = None;
|
||||
image_crop = None;
|
||||
images_seq_mask = None;
|
||||
images_spatial_crop_t = None;
|
||||
}
|
||||
self.deepseekocr_model.clear_kv_cache();
|
||||
};
|
||||
Ok(stream)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
pub mod processor;
|
||||
pub mod generate;
|
||||
pub mod config;
|
||||
pub mod model;
|
||||
pub mod generate;
|
||||
pub mod model;
|
||||
pub mod processor;
|
||||
|
||||
+1024
-53
File diff suppressed because it is too large
Load Diff
@@ -71,7 +71,7 @@ impl DeepseekOCRProcessor {
|
||||
let mut tokenized_id = vec![0u32];
|
||||
let mut images_spatial_crop = Vec::new();
|
||||
for (text_seq, image) in text_splits.iter().zip(imgs) {
|
||||
if text_seq.len() > 0 {
|
||||
if !text_seq.is_empty() {
|
||||
let token_ids = tokenizer.text_encode_vec(text_seq.to_string(), false)?;
|
||||
tokenized_id.extend_from_slice(&token_ids);
|
||||
let seq_mask = vec![0u32; token_ids.len()];
|
||||
@@ -143,8 +143,8 @@ impl DeepseekOCRProcessor {
|
||||
let seq_mask = vec![0u32; token_ids.len()];
|
||||
images_seq_mask.extend_from_slice(&seq_mask);
|
||||
let input_ids = Tensor::new(tokenized_id, &self.device)?.unsqueeze(0)?;
|
||||
let image_seq_mask = Tensor::new(images_seq_mask, &self.device)?;
|
||||
let (images_ori, images_spatial_crop_t, image_crop) = if images_list.len() == 0 {
|
||||
let image_seq_mask = Tensor::new(images_seq_mask, &self.device)?.unsqueeze(0)?;
|
||||
let (images_ori, images_spatial_crop_t, image_crop) = if images_list.is_empty() {
|
||||
let images_ori = Tensor::zeros(
|
||||
(1usize, 3usize, image_size as usize, image_size as usize),
|
||||
self.dtype,
|
||||
@@ -160,7 +160,7 @@ impl DeepseekOCRProcessor {
|
||||
} else {
|
||||
let images_ori = Tensor::stack(&images_list, 0)?;
|
||||
let images_spatial_crop_t = Tensor::new(images_spatial_crop, &self.device)?;
|
||||
let image_crop = if images_crop_list.len() > 0 {
|
||||
let image_crop = if !images_crop_list.is_empty() {
|
||||
Tensor::stack(&images_crop_list, 0)?
|
||||
} else {
|
||||
Tensor::zeros(
|
||||
|
||||
+1
-1
@@ -1,9 +1,9 @@
|
||||
pub mod common;
|
||||
pub mod deepseek_ocr;
|
||||
pub mod minicpm4;
|
||||
pub mod qwen2_5vl;
|
||||
pub mod qwen3vl;
|
||||
pub mod voxcpm;
|
||||
pub mod deepseek_ocr;
|
||||
|
||||
use aha_openai_dive::v1::resources::chat::{
|
||||
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
|
||||
|
||||
Reference in New Issue
Block a user