deepseek_ocr mes metadata add base_size,image_size,crop_mode
This commit is contained in:
@@ -28,6 +28,7 @@ pub struct DeepseekOCRGenerateModel {
|
|||||||
bos_token_id: u32,
|
bos_token_id: u32,
|
||||||
eos_token_id: u32,
|
eos_token_id: u32,
|
||||||
device: Device,
|
device: Device,
|
||||||
|
size: Vec<u32>,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl DeepseekOCRGenerateModel {
|
impl DeepseekOCRGenerateModel {
|
||||||
@@ -44,6 +45,7 @@ impl DeepseekOCRGenerateModel {
|
|||||||
let model_list = find_type_files(path, "safetensors")?;
|
let model_list = find_type_files(path, "safetensors")?;
|
||||||
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
let vb = unsafe { VarBuilder::from_mmaped_safetensors(&model_list, dtype, device)? };
|
||||||
let deepseekocr_model = DeepseekOCRModel::new(vb, cfg)?;
|
let deepseekocr_model = DeepseekOCRModel::new(vb, cfg)?;
|
||||||
|
let size = vec![512u32, 640, 1024, 1280];
|
||||||
Ok(Self {
|
Ok(Self {
|
||||||
tokenizer,
|
tokenizer,
|
||||||
processor,
|
processor,
|
||||||
@@ -51,16 +53,47 @@ impl DeepseekOCRGenerateModel {
|
|||||||
bos_token_id,
|
bos_token_id,
|
||||||
eos_token_id,
|
eos_token_id,
|
||||||
device: device.clone(),
|
device: device.clone(),
|
||||||
|
size,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
impl GenerateModel for DeepseekOCRGenerateModel {
|
impl GenerateModel for DeepseekOCRGenerateModel {
|
||||||
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
|
||||||
|
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
|
||||||
|
};
|
||||||
|
println!(
|
||||||
|
"base_size: {}, image_size: {}, crop_mode: {}",
|
||||||
|
base_size, image_size, crop_mode
|
||||||
|
);
|
||||||
let mut logit_processor = get_logit_processor(mes.temperature, mes.top_p, None);
|
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
|
let (mut input_ids, images_ori, image_crop, images_seq_mask, images_spatial_crop_t) = self
|
||||||
.processor
|
.processor
|
||||||
.process_info(&mes, &self.tokenizer, 640, 640, true)?;
|
.process_info(&mes, &self.tokenizer, base_size, image_size, crop_mode)?;
|
||||||
let mut images_ori = Some(&images_ori);
|
let mut images_ori = Some(&images_ori);
|
||||||
let mut image_crop = Some(&image_crop);
|
let mut image_crop = Some(&image_crop);
|
||||||
let mut images_seq_mask = Some(&images_seq_mask);
|
let mut images_seq_mask = Some(&images_seq_mask);
|
||||||
|
|||||||
@@ -27,12 +27,9 @@ fn deepseek_ocr_generate() -> Result<()> {
|
|||||||
"text": "<image>\n<|grounding|>Convert the document to markdown. "
|
"text": "<image>\n<|grounding|>Convert the document to markdown. "
|
||||||
}
|
}
|
||||||
]
|
]
|
||||||
},
|
|
||||||
{
|
|
||||||
"role": "assistant",
|
|
||||||
"content": ""
|
|
||||||
}
|
}
|
||||||
]
|
],
|
||||||
|
"metadata": {"base_size": "640", "image_size": "640", "crop_mode": "false"}
|
||||||
}
|
}
|
||||||
"#;
|
"#;
|
||||||
let model_path = "/home/jhq/huggingface_model/deepseek-ai/DeepSeek-OCR/";
|
let model_path = "/home/jhq/huggingface_model/deepseek-ai/DeepSeek-OCR/";
|
||||||
@@ -77,7 +74,8 @@ async fn deepseek_ocr_stream() -> Result<()> {
|
|||||||
"role": "assistant",
|
"role": "assistant",
|
||||||
"content": ""
|
"content": ""
|
||||||
}
|
}
|
||||||
]
|
],
|
||||||
|
"metadata": {"base_size": "640", "image_size": "640", "crop_mode": "false"}
|
||||||
}
|
}
|
||||||
"#;
|
"#;
|
||||||
let model_path = "/home/jhq/huggingface_model/deepseek-ai/DeepSeek-OCR/";
|
let model_path = "/home/jhq/huggingface_model/deepseek-ai/DeepSeek-OCR/";
|
||||||
|
|||||||
Reference in New Issue
Block a user