add deepseek_ocr

This commit is contained in:
jhqxxx
2025-11-22 23:27:14 +08:00
parent c95b91ffc0
commit a14c35014a
20 changed files with 1715 additions and 151 deletions
+3
View File
@@ -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"))]
{
+41 -6
View File
@@ -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,
}
}
+142 -12
View File
@@ -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)
}
}
+3 -3
View File
@@ -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;
File diff suppressed because it is too large Load Diff
+4 -4
View File
@@ -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
View File
@@ -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,