给 glm-ocr 添加 gguf 和 onnx 格式推理
This commit is contained in:
@@ -27,6 +27,7 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
|
||||
## Changelog
|
||||
### v0.2.3 (2026-03-18)
|
||||
- add DeepSeek-OCR-2
|
||||
- add GLM-OCR gguf and onnx local loading
|
||||
|
||||
### 2026-03-17
|
||||
- add PaddleOCR-VL1.5 model
|
||||
@@ -106,6 +107,12 @@ aha run -m all-minilm-l6-v2 -i "Rust embedding test" --artifact-format gguf --gg
|
||||
# Run local all-MiniLM-L6-v2 embedding (ONNX)
|
||||
aha run -m all-minilm-l6-v2 -i "Rust embedding test" --artifact-format onnx --onnx-path D:\model_download\all-MiniLM-L6-v2\onnx --tokenizer-dir D:\model_download\all-MiniLM-L6-v2
|
||||
|
||||
# Run local GLM-OCR (GGUF)
|
||||
aha run -m glm-ocr -i .\assets\img\ocr_test1.png --artifact-format gguf --gguf-path D:\model_download\GLM-OCR-GGUF
|
||||
|
||||
# Run local GLM-OCR (ONNX)
|
||||
aha run -m glm-ocr -i .\assets\img\ocr_test1.png --artifact-format onnx --onnx-path D:\model_download\GLM-OCR-ONNX --tokenizer-dir D:\model_download\GLM-OCR-ONNX
|
||||
|
||||
# Start service only (model already downloaded)
|
||||
aha serv -m qwen3asr-0.6b -p 10100
|
||||
|
||||
|
||||
@@ -27,6 +27,7 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
|
||||
## 更新日志
|
||||
### v0.2.3 (2026-03-18)
|
||||
- 新增 DeepSeek-OCR-2
|
||||
- 增加 GLM-OCR 的 GGUF 和 ONNX 本地加载
|
||||
|
||||
### 2026-03-17
|
||||
- 新增 PaddleOCR-VL1.5 模型
|
||||
@@ -105,6 +106,12 @@ aha run -m all-minilm-l6-v2 -i "Rust embedding test" --artifact-format gguf --gg
|
||||
# 本地运行 all-MiniLM-L6-v2 向量模型(ONNX)
|
||||
aha run -m all-minilm-l6-v2 -i "Rust embedding test" --artifact-format onnx --onnx-path D:\model_download\all-MiniLM-L6-v2\onnx --tokenizer-dir D:\model_download\all-MiniLM-L6-v2
|
||||
|
||||
# 本地运行 GLM-OCR(GGUF)
|
||||
aha run -m glm-ocr -i .\assets\img\ocr_test1.png --artifact-format gguf --gguf-path D:\model_download\GLM-OCR-GGUF
|
||||
|
||||
# 本地运行 GLM-OCR(ONNX)
|
||||
aha run -m glm-ocr -i .\assets\img\ocr_test1.png --artifact-format onnx --onnx-path D:\model_download\GLM-OCR-ONNX --tokenizer-dir D:\model_download\GLM-OCR-ONNX
|
||||
|
||||
# 仅启动服务(模型已下载)
|
||||
aha serv -m qwen3asr-0.6b -p 10100
|
||||
|
||||
|
||||
@@ -63,6 +63,8 @@ aha supports a growing collection of state-of-the-art AI models across multiple
|
||||
| **DeepSeek-OCR** | Multi | Scene text | Natural images | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) |
|
||||
| **GLM-OCR** | 8 | Scene text | complex document | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) |
|
||||
|
||||
GLM-OCR local artifacts: `safetensors`, `gguf`, `onnx`
|
||||
|
||||
## Speech Recognition (ASR)
|
||||
|
||||
| Model | Parameters | Language | Real-time | Speed | License |
|
||||
|
||||
@@ -63,6 +63,8 @@ aha 支持多个领域的最先进 AI 模型集合。
|
||||
| **DeepSeek-OCR** | 多语言 | 场景文字 | 自然图像 | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) |
|
||||
| **GLM-OCR** | 8 | 场景文字 | 复杂文档 | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) |
|
||||
|
||||
GLM-OCR 本地制品格式:`safetensors`、`gguf`、`onnx`
|
||||
|
||||
## 语音识别 (ASR)
|
||||
|
||||
| 模型 | 参数量 | 语言 | 实时 | 速度 | 开源协议 |
|
||||
|
||||
+113
-41
@@ -1,56 +1,84 @@
|
||||
//! Glm-OCR exec implementation for CLI `run` subcommand
|
||||
use std::time::Instant;
|
||||
use std::{path::Path, time::Instant};
|
||||
|
||||
use anyhow::{Ok, Result};
|
||||
use anyhow::{Ok, Result, anyhow};
|
||||
use serde_json::json;
|
||||
|
||||
use crate::exec::ExecModel;
|
||||
use crate::models::{GenerateModel, glm_ocr::generate::GlmOcrGenerateModel};
|
||||
use crate::models::{GenerateModel, LoadSpec, glm_ocr::generate::GlmOcrGenerateModel};
|
||||
|
||||
pub struct GlmOcrExec;
|
||||
|
||||
impl ExecModel for GlmOcrExec {
|
||||
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
|
||||
let url = &input[0];
|
||||
let input_url = if url.starts_with("http://")
|
||||
|| url.starts_with("https://")
|
||||
|| url.starts_with("file://")
|
||||
{
|
||||
url.clone()
|
||||
} else {
|
||||
format!("file://{}", url)
|
||||
};
|
||||
fn resolve_input_url(input: &str) -> Result<String> {
|
||||
if input.starts_with("http://") || input.starts_with("https://") || input.starts_with("file://")
|
||||
{
|
||||
return Ok(input.to_string());
|
||||
}
|
||||
|
||||
let path = Path::new(input);
|
||||
let path = if path.is_absolute() {
|
||||
path.to_path_buf()
|
||||
} else {
|
||||
std::env::current_dir()?.join(path)
|
||||
};
|
||||
let path = path
|
||||
.canonicalize()
|
||||
.map_err(|e| anyhow!("failed to resolve input path {}: {e}", path.display()))?;
|
||||
url::Url::from_file_path(&path)
|
||||
.map(|url| url.to_string())
|
||||
.map_err(|_| {
|
||||
anyhow!(
|
||||
"failed to convert input path to file url: {}",
|
||||
path.display()
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
fn build_request(
|
||||
input_url: &str,
|
||||
max_tokens: Option<u32>,
|
||||
) -> Result<aha_openai_dive::v1::resources::chat::ChatCompletionParameters> {
|
||||
Ok(serde_json::from_value(json!({
|
||||
"model": "glm-ocr",
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "image_url",
|
||||
"image_url": {
|
||||
"url": input_url
|
||||
}
|
||||
},
|
||||
{
|
||||
"type": "text",
|
||||
"text": "Text Recognition:"
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"max_tokens": max_tokens.unwrap_or(256)
|
||||
}))?)
|
||||
}
|
||||
|
||||
impl GlmOcrExec {
|
||||
pub fn run_with_spec(
|
||||
input: &[String],
|
||||
output: Option<&str>,
|
||||
spec: &LoadSpec,
|
||||
max_tokens: Option<u32>,
|
||||
) -> Result<()> {
|
||||
let url = input
|
||||
.first()
|
||||
.ok_or_else(|| anyhow!("glm-ocr run requires an input image path or url"))?;
|
||||
let input_url = resolve_input_url(url)?;
|
||||
|
||||
let i_start = Instant::now();
|
||||
let mut model = GlmOcrGenerateModel::init(weight_path, None, None)?;
|
||||
let mut model = GlmOcrGenerateModel::init_from_spec(spec, None, None)?;
|
||||
let i_duration = i_start.elapsed();
|
||||
println!("Time elapsed in load model is: {:?}", i_duration);
|
||||
|
||||
let message = format!(
|
||||
r#"{{
|
||||
"model": "glm-ocr",
|
||||
"messages": [
|
||||
{{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{{
|
||||
"type": "image_url",
|
||||
"image_url": {{
|
||||
"url": "{}"
|
||||
}}
|
||||
}},
|
||||
{{
|
||||
"type": "text",
|
||||
"text": "Text Recognition:"
|
||||
}}
|
||||
]
|
||||
}}
|
||||
],
|
||||
"max_tokens": 1024
|
||||
}}"#,
|
||||
input_url
|
||||
);
|
||||
|
||||
let mes = serde_json::from_str(&message)?;
|
||||
let mes = build_request(&input_url, max_tokens)?;
|
||||
|
||||
let i_start = Instant::now();
|
||||
let result = model.generate(mes)?;
|
||||
@@ -67,3 +95,47 @@ impl ExecModel for GlmOcrExec {
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
impl ExecModel for GlmOcrExec {
|
||||
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> {
|
||||
let spec = LoadSpec::for_safetensors(crate::models::WhichModel::GlmOCR, weight_path);
|
||||
Self::run_with_spec(input, output, &spec, None)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::{build_request, resolve_input_url};
|
||||
use aha_openai_dive::v1::resources::chat::{
|
||||
ChatMessage, ChatMessageContent, ChatMessageContentPart,
|
||||
};
|
||||
use anyhow::Result;
|
||||
|
||||
#[test]
|
||||
fn resolve_input_url_converts_relative_windows_path_to_file_url() -> Result<()> {
|
||||
let url = resolve_input_url(r".\assets\img\ocr_test1.png")?;
|
||||
assert!(url.starts_with("file:///"));
|
||||
assert!(url.contains("assets/img/ocr_test1.png"));
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn build_request_accepts_file_url_without_json_escape_issues() -> Result<()> {
|
||||
let mes = build_request("file:///D:/model_download/ocr_test1.png", Some(32))?;
|
||||
let ChatMessage::User { content, .. } = &mes.messages[0] else {
|
||||
panic!("expected user message");
|
||||
};
|
||||
let ChatMessageContent::ContentPart(parts) = content else {
|
||||
panic!("expected multipart content");
|
||||
};
|
||||
let ChatMessageContentPart::Image(image) = &parts[0] else {
|
||||
panic!("expected image content part");
|
||||
};
|
||||
assert_eq!(
|
||||
image.image_url.url,
|
||||
"file:///D:/model_download/ocr_test1.png"
|
||||
);
|
||||
assert_eq!(mes.max_tokens, Some(32));
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
+72
@@ -236,6 +236,10 @@ struct RunArgs {
|
||||
#[arg(short, long)]
|
||||
output: Option<String>,
|
||||
|
||||
/// Maximum number of tokens to generate
|
||||
#[arg(long)]
|
||||
max_tokens: Option<u32>,
|
||||
|
||||
/// Local model weight path (defaults to ~/.aha/{model_id} if not specified)
|
||||
#[arg(long)]
|
||||
weight_path: Option<String>,
|
||||
@@ -467,6 +471,11 @@ fn run_target_model_with_spec(args: &RunArgs, spec: &LoadSpec) -> anyhow::Result
|
||||
Qwen3_5Exec::run_with_spec(&args.input, args.output.as_deref(), spec)?;
|
||||
Ok(true)
|
||||
}
|
||||
WhichModel::GlmOCR => {
|
||||
use aha::exec::glm_ocr::GlmOcrExec;
|
||||
GlmOcrExec::run_with_spec(&args.input, args.output.as_deref(), spec, args.max_tokens)?;
|
||||
Ok(true)
|
||||
}
|
||||
_ => Ok(false),
|
||||
}
|
||||
}
|
||||
@@ -1017,6 +1026,69 @@ mod tests {
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_glm_ocr_run_gguf_flags() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"aha",
|
||||
"run",
|
||||
"--model",
|
||||
"glm-ocr",
|
||||
"--input",
|
||||
"ocr.png",
|
||||
"--artifact-format",
|
||||
"gguf",
|
||||
"--gguf-path",
|
||||
"D:\\model_download\\GLM-OCR-GGUF",
|
||||
"--max-tokens",
|
||||
"8",
|
||||
])
|
||||
.expect("run args should parse");
|
||||
|
||||
let Some(Commands::Run(args)) = cli.command else {
|
||||
panic!("expected run subcommand");
|
||||
};
|
||||
assert!(matches!(args.artifact_format, Some(ArtifactArg::Gguf)));
|
||||
assert_eq!(args.model, WhichModel::GlmOCR);
|
||||
assert_eq!(
|
||||
args.gguf_path.as_deref(),
|
||||
Some("D:\\model_download\\GLM-OCR-GGUF")
|
||||
);
|
||||
assert_eq!(args.max_tokens, Some(8));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_glm_ocr_run_onnx_flags() {
|
||||
let cli = Cli::try_parse_from([
|
||||
"aha",
|
||||
"run",
|
||||
"--model",
|
||||
"glm-ocr",
|
||||
"--input",
|
||||
"ocr.png",
|
||||
"--artifact-format",
|
||||
"onnx",
|
||||
"--onnx-path",
|
||||
"D:\\model_download\\GLM-OCR-ONNX",
|
||||
"--tokenizer-dir",
|
||||
"D:\\model_download\\GLM-OCR-ONNX",
|
||||
])
|
||||
.expect("run args should parse");
|
||||
|
||||
let Some(Commands::Run(args)) = cli.command else {
|
||||
panic!("expected run subcommand");
|
||||
};
|
||||
assert!(matches!(args.artifact_format, Some(ArtifactArg::Onnx)));
|
||||
assert_eq!(args.model, WhichModel::GlmOCR);
|
||||
assert_eq!(
|
||||
args.onnx_path.as_deref(),
|
||||
Some("D:\\model_download\\GLM-OCR-ONNX")
|
||||
);
|
||||
assert_eq!(
|
||||
args.tokenizer_dir.as_deref(),
|
||||
Some("D:\\model_download\\GLM-OCR-ONNX")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_serv_onnx_flags() {
|
||||
let cli = Cli::try_parse_from([
|
||||
|
||||
@@ -86,6 +86,10 @@ impl<R: Read + Seek> Gguf<R> {
|
||||
&self.ct.metadata
|
||||
}
|
||||
|
||||
pub fn has_tensor(&self, name: &str) -> bool {
|
||||
self.ct.tensor_infos.contains_key(name)
|
||||
}
|
||||
|
||||
pub fn tensor(&mut self, name: &str) -> Result<QTensor> {
|
||||
Ok(self.ct.tensor(&mut self.reader, name, &self.device)?)
|
||||
}
|
||||
|
||||
@@ -33,6 +33,11 @@ pub fn supported_artifacts(model: WhichModel) -> &'static [ArtifactKind] {
|
||||
ArtifactKind::Gguf,
|
||||
ArtifactKind::Onnx,
|
||||
],
|
||||
WhichModel::GlmOCR => &[
|
||||
ArtifactKind::Safetensors,
|
||||
ArtifactKind::Gguf,
|
||||
ArtifactKind::Onnx,
|
||||
],
|
||||
WhichModel::Qwen3_0_6B => &[
|
||||
ArtifactKind::Safetensors,
|
||||
ArtifactKind::Gguf,
|
||||
|
||||
@@ -1,9 +1,10 @@
|
||||
use anyhow::Result;
|
||||
|
||||
use crate::models::{
|
||||
ModelInstance, WhichModel, all_minilm_l6_v2::generate::AllMiniLML6V2Model, load_model_legacy,
|
||||
qwen3::generate::Qwen3GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel,
|
||||
qwen3_embedding::generate::Qwen3EmbeddingModel, qwen3_reranker::generate::Qwen3RerankerModel,
|
||||
ModelInstance, WhichModel, all_minilm_l6_v2::generate::AllMiniLML6V2Model,
|
||||
glm_ocr::generate::GlmOcrGenerateModel, load_model_legacy, qwen3::generate::Qwen3GenerateModel,
|
||||
qwen3_5::generate::Qwen3_5GenerateModel, qwen3_embedding::generate::Qwen3EmbeddingModel,
|
||||
qwen3_reranker::generate::Qwen3RerankerModel,
|
||||
};
|
||||
|
||||
use super::artifact::LoadSpec;
|
||||
@@ -11,6 +12,7 @@ use super::artifact::LoadSpec;
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
enum ModelLoaderFamily {
|
||||
AllMiniLML6V2,
|
||||
GlmOCR,
|
||||
Qwen3,
|
||||
Qwen3Embedding,
|
||||
Qwen3Reranker,
|
||||
@@ -21,6 +23,7 @@ enum ModelLoaderFamily {
|
||||
fn resolve_model_loader_family(model: WhichModel) -> ModelLoaderFamily {
|
||||
match model {
|
||||
WhichModel::AllMiniLML6V2 => ModelLoaderFamily::AllMiniLML6V2,
|
||||
WhichModel::GlmOCR => ModelLoaderFamily::GlmOCR,
|
||||
WhichModel::Qwen3_0_6B => ModelLoaderFamily::Qwen3,
|
||||
WhichModel::Qwen3Embedding0_6B
|
||||
| WhichModel::Qwen3Embedding4B
|
||||
@@ -49,6 +52,9 @@ pub fn load_model_from_spec<'a>(spec: &LoadSpec) -> Result<ModelInstance<'a>> {
|
||||
ModelLoaderFamily::AllMiniLML6V2 => {
|
||||
ModelInstance::AllMiniLML6V2(AllMiniLML6V2Model::init_from_spec(spec, None, None)?)
|
||||
}
|
||||
ModelLoaderFamily::GlmOCR => {
|
||||
ModelInstance::GlmOCR(GlmOcrGenerateModel::init_from_spec(spec, None, None)?)
|
||||
}
|
||||
ModelLoaderFamily::Qwen3 => {
|
||||
ModelInstance::Qwen3(Qwen3GenerateModel::init_from_spec(spec, None, None)?)
|
||||
}
|
||||
@@ -86,6 +92,10 @@ mod tests {
|
||||
resolve_model_loader_family(WhichModel::Qwen3_0_6B),
|
||||
ModelLoaderFamily::Qwen3
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_model_loader_family(WhichModel::GlmOCR),
|
||||
ModelLoaderFamily::GlmOCR
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_model_loader_family(WhichModel::Qwen3Embedding0_6B),
|
||||
ModelLoaderFamily::Qwen3Embedding
|
||||
|
||||
+969
-33
File diff suppressed because it is too large
Load Diff
@@ -1,4 +1,5 @@
|
||||
pub mod config;
|
||||
pub mod generate;
|
||||
pub mod model;
|
||||
pub mod onnx;
|
||||
pub mod processor;
|
||||
|
||||
@@ -0,0 +1,847 @@
|
||||
use anyhow::{Result, anyhow};
|
||||
use candle_core::Tensor;
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
use candle_core::{DType, IndexOp};
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
use half::f16;
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
use ndarray::{Array, IxDyn};
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
use crate::models::common::onnx::create_session;
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
pub struct GlmOcrOnnxCacheEntry {
|
||||
pub name: String,
|
||||
pub dims: Vec<i64>,
|
||||
pub data: Vec<f16>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
#[derive(Clone)]
|
||||
struct OnnxInputDescriptor {
|
||||
name: String,
|
||||
shape: Vec<i64>,
|
||||
kind: Option<OnnxTensorKind>,
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
#[derive(Clone, Copy)]
|
||||
enum OnnxTensorKind {
|
||||
Bool,
|
||||
I32,
|
||||
I64,
|
||||
F16,
|
||||
F32,
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn map_tensor_kind(ty: ort::value::TensorElementType) -> Option<OnnxTensorKind> {
|
||||
match ty {
|
||||
ort::value::TensorElementType::Bool => Some(OnnxTensorKind::Bool),
|
||||
ort::value::TensorElementType::Int32 => Some(OnnxTensorKind::I32),
|
||||
ort::value::TensorElementType::Int64 => Some(OnnxTensorKind::I64),
|
||||
ort::value::TensorElementType::Float16 => Some(OnnxTensorKind::F16),
|
||||
ort::value::TensorElementType::Float32 => Some(OnnxTensorKind::F32),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
pub struct GlmOcrOnnxBackend {
|
||||
embed_session: ort::session::Session,
|
||||
decoder_session: ort::session::Session,
|
||||
vision_session: ort::session::Session,
|
||||
decoder_input_descriptors: Vec<OnnxInputDescriptor>,
|
||||
decoder_output_names: Vec<String>,
|
||||
vision_input_descriptors: Vec<OnnxInputDescriptor>,
|
||||
vision_output_names: Vec<String>,
|
||||
cache_values: Vec<GlmOcrOnnxCacheEntry>,
|
||||
spatial_merge_size: usize,
|
||||
next_mrope_pos: usize,
|
||||
prefill_seq_len: usize,
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn find_onnx_component_file(path: &str, marker: &str) -> Result<PathBuf> {
|
||||
let model_path = Path::new(path);
|
||||
if !model_path.exists() {
|
||||
return Err(anyhow!("onnx model path not found: {}", path));
|
||||
}
|
||||
|
||||
if model_path.is_file() {
|
||||
let file_name = model_path
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.unwrap_or_default();
|
||||
if file_name.contains(marker) && file_name.ends_with(".onnx") {
|
||||
return Ok(model_path.to_path_buf());
|
||||
}
|
||||
}
|
||||
|
||||
let search_root = if model_path.is_dir() {
|
||||
model_path.to_path_buf()
|
||||
} else {
|
||||
model_path
|
||||
.parent()
|
||||
.ok_or_else(|| anyhow!("onnx component parent directory not found for {}", path))?
|
||||
.to_path_buf()
|
||||
};
|
||||
|
||||
let mut stack = vec![search_root];
|
||||
let mut matches = Vec::new();
|
||||
while let Some(current) = stack.pop() {
|
||||
for entry in std::fs::read_dir(¤t)? {
|
||||
let entry = entry?;
|
||||
let entry_path = entry.path();
|
||||
if entry_path.is_dir() {
|
||||
stack.push(entry_path);
|
||||
continue;
|
||||
}
|
||||
let file_name = entry_path
|
||||
.file_name()
|
||||
.and_then(|name| name.to_str())
|
||||
.unwrap_or_default();
|
||||
if file_name.contains(marker) && file_name.ends_with(".onnx") {
|
||||
matches.push(entry_path);
|
||||
}
|
||||
}
|
||||
}
|
||||
matches.sort();
|
||||
matches.into_iter().next().ok_or_else(|| {
|
||||
anyhow!(
|
||||
"unable to locate onnx component {} under {}",
|
||||
marker,
|
||||
model_path.display()
|
||||
)
|
||||
})
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
impl GlmOcrOnnxBackend {
|
||||
pub fn load(onnx_path: &str, spatial_merge_size: usize) -> Result<Self> {
|
||||
let embed_file = find_onnx_component_file(onnx_path, "embed_tokens")?;
|
||||
let decoder_file = find_onnx_component_file(onnx_path, "decoder_model_merged")?;
|
||||
let vision_file = find_onnx_component_file(onnx_path, "vision_encoder")?;
|
||||
let embed_bundle = create_session(&embed_file.to_string_lossy(), None)?;
|
||||
let decoder_bundle = create_session(&decoder_file.to_string_lossy(), None)?;
|
||||
let vision_bundle = create_session(&vision_file.to_string_lossy(), None)?;
|
||||
|
||||
let decoder_input_descriptors = decoder_bundle
|
||||
.session
|
||||
.inputs()
|
||||
.iter()
|
||||
.map(|input| {
|
||||
let (shape, kind) = match input.dtype() {
|
||||
ort::value::ValueType::Tensor { ty, shape, .. } => (
|
||||
shape.iter().copied().collect::<Vec<_>>(),
|
||||
map_tensor_kind(*ty),
|
||||
),
|
||||
_ => (Vec::new(), None),
|
||||
};
|
||||
OnnxInputDescriptor {
|
||||
name: input.name().to_string(),
|
||||
shape,
|
||||
kind,
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
let vision_input_descriptors = vision_bundle
|
||||
.session
|
||||
.inputs()
|
||||
.iter()
|
||||
.map(|input| {
|
||||
let (shape, kind) = match input.dtype() {
|
||||
ort::value::ValueType::Tensor { ty, shape, .. } => (
|
||||
shape.iter().copied().collect::<Vec<_>>(),
|
||||
map_tensor_kind(*ty),
|
||||
),
|
||||
_ => (Vec::new(), None),
|
||||
};
|
||||
OnnxInputDescriptor {
|
||||
name: input.name().to_string(),
|
||||
shape,
|
||||
kind,
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
Ok(Self {
|
||||
embed_session: embed_bundle.session,
|
||||
decoder_session: decoder_bundle.session,
|
||||
vision_session: vision_bundle.session,
|
||||
decoder_input_descriptors,
|
||||
decoder_output_names: decoder_bundle.output_names,
|
||||
vision_input_descriptors,
|
||||
vision_output_names: vision_bundle.output_names,
|
||||
cache_values: Vec::new(),
|
||||
spatial_merge_size,
|
||||
next_mrope_pos: 0,
|
||||
prefill_seq_len: 0,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn clear_cache(&mut self) {
|
||||
self.cache_values.clear();
|
||||
self.next_mrope_pos = 0;
|
||||
self.prefill_seq_len = 0;
|
||||
}
|
||||
|
||||
pub fn forward_logits(
|
||||
&mut self,
|
||||
input_ids: &[u32],
|
||||
image_mask: Option<&Tensor>,
|
||||
pixel_values: Option<&Tensor>,
|
||||
image_grid_thw: Option<&Tensor>,
|
||||
position_start: usize,
|
||||
) -> Result<Vec<f32>> {
|
||||
if input_ids.is_empty() {
|
||||
return Err(anyhow!("glm-ocr onnx input_ids cannot be empty"));
|
||||
}
|
||||
|
||||
let (mut embed_data, hidden_size) = self.embed_input_ids(input_ids)?;
|
||||
if let Some(pixel_values) = pixel_values {
|
||||
let image_mask =
|
||||
image_mask.ok_or_else(|| anyhow!("glm-ocr onnx image_mask is required"))?;
|
||||
let image_grid_thw =
|
||||
image_grid_thw.ok_or_else(|| anyhow!("glm-ocr onnx image_grid_thw is required"))?;
|
||||
self.apply_vision_embeds(
|
||||
&mut embed_data,
|
||||
hidden_size,
|
||||
image_mask,
|
||||
pixel_values,
|
||||
image_grid_thw,
|
||||
)?;
|
||||
}
|
||||
|
||||
let position_ids = if self.cache_values.is_empty() {
|
||||
if let (Some(mask), Some(grid_thw)) = (image_mask, image_grid_thw) {
|
||||
self.compute_prefill_position_ids(mask, grid_thw, input_ids.len())?
|
||||
} else {
|
||||
self.next_mrope_pos = input_ids.len();
|
||||
self.prefill_seq_len = input_ids.len();
|
||||
build_text_position_ids(0, input_ids.len())
|
||||
}
|
||||
} else {
|
||||
let decode_pos = self
|
||||
.next_mrope_pos
|
||||
.saturating_add(position_start.saturating_sub(self.prefill_seq_len));
|
||||
build_text_position_ids(decode_pos, input_ids.len())
|
||||
};
|
||||
|
||||
let mut decoder_inputs = Vec::with_capacity(self.decoder_input_descriptors.len());
|
||||
for desc in &self.decoder_input_descriptors {
|
||||
let value = match desc.name.as_str() {
|
||||
"inputs_embeds" => build_float_input(
|
||||
desc,
|
||||
vec![1_i64, input_ids.len() as i64, hidden_size as i64],
|
||||
&embed_data,
|
||||
)?,
|
||||
"attention_mask" => {
|
||||
let attention_mask = self.build_attention_mask(input_ids.len())?;
|
||||
build_i64_like_input(
|
||||
desc,
|
||||
vec![1_i64, attention_mask.len() as i64],
|
||||
&attention_mask,
|
||||
)?
|
||||
}
|
||||
"position_ids" => build_i64_like_input(
|
||||
desc,
|
||||
vec![3_i64, 1_i64, input_ids.len() as i64],
|
||||
&position_ids,
|
||||
)?,
|
||||
"past_sequence_length" => build_i64_scalar_input(desc, self.past_seq_len() as i64)?,
|
||||
name if name.starts_with("past_key_values.") => {
|
||||
if let Some(cache) = self.cache_values.iter().find(|entry| entry.name == name) {
|
||||
ort::value::Tensor::from_array((cache.dims.clone(), cache.data.clone()))?
|
||||
.into_dyn()
|
||||
} else {
|
||||
build_zero_cache_input(desc)?
|
||||
}
|
||||
}
|
||||
_ => build_zero_input(desc)?,
|
||||
};
|
||||
decoder_inputs.push((desc.name.clone(), value));
|
||||
}
|
||||
|
||||
let decoder_outputs = self.decoder_session.run(decoder_inputs)?;
|
||||
let logits_value = decoder_outputs
|
||||
.get("logits")
|
||||
.or_else(|| {
|
||||
self.decoder_output_names
|
||||
.first()
|
||||
.and_then(|name| decoder_outputs.get(name))
|
||||
})
|
||||
.ok_or_else(|| anyhow!("glm-ocr onnx output logits not found"))?;
|
||||
let logits = extract_last_logits(logits_value)?;
|
||||
|
||||
let mut new_cache_values = Vec::new();
|
||||
for desc in &self.decoder_input_descriptors {
|
||||
let name = &desc.name;
|
||||
if !name.starts_with("past_key_values.") {
|
||||
continue;
|
||||
}
|
||||
let present_name = name.replace("past_key_values.", "present.");
|
||||
let value = decoder_outputs
|
||||
.get(&present_name)
|
||||
.ok_or_else(|| anyhow!("missing glm-ocr onnx output {}", present_name))?;
|
||||
let (shape, data) = value.try_extract_tensor::<f16>()?;
|
||||
new_cache_values.push(GlmOcrOnnxCacheEntry {
|
||||
name: name.clone(),
|
||||
dims: shape.iter().copied().collect::<Vec<_>>(),
|
||||
data: data.to_vec(),
|
||||
});
|
||||
}
|
||||
self.cache_values = new_cache_values;
|
||||
Ok(logits)
|
||||
}
|
||||
|
||||
fn embed_input_ids(&mut self, input_ids: &[u32]) -> Result<(Vec<f32>, usize)> {
|
||||
let outputs = self.embed_session.run(vec![(
|
||||
"input_ids".to_string(),
|
||||
ort::value::Tensor::from_array((
|
||||
vec![1_i64, input_ids.len() as i64],
|
||||
input_ids.iter().map(|id| *id as i64).collect::<Vec<_>>(),
|
||||
))?
|
||||
.into_dyn(),
|
||||
)])?;
|
||||
let embed_value = outputs
|
||||
.get("inputs_embeds")
|
||||
.ok_or_else(|| anyhow!("glm-ocr onnx output inputs_embeds not found"))?;
|
||||
|
||||
if let Ok((shape, embed_data)) = embed_value.try_extract_tensor::<f32>() {
|
||||
if shape.len() != 3 || shape[0] != 1 {
|
||||
return Err(anyhow!(
|
||||
"unexpected glm-ocr onnx inputs_embeds shape: {}",
|
||||
shape
|
||||
));
|
||||
}
|
||||
return Ok((embed_data.to_vec(), shape[2] as usize));
|
||||
}
|
||||
if let Ok((shape, embed_data)) = embed_value.try_extract_tensor::<f16>() {
|
||||
if shape.len() != 3 || shape[0] != 1 {
|
||||
return Err(anyhow!(
|
||||
"unexpected glm-ocr onnx inputs_embeds shape: {}",
|
||||
shape
|
||||
));
|
||||
}
|
||||
return Ok((
|
||||
embed_data
|
||||
.iter()
|
||||
.map(|value| value.to_f32())
|
||||
.collect::<Vec<_>>(),
|
||||
shape[2] as usize,
|
||||
));
|
||||
}
|
||||
Err(anyhow!(
|
||||
"glm-ocr onnx inputs_embeds output must be a f32/f16 tensor"
|
||||
))
|
||||
}
|
||||
|
||||
fn apply_vision_embeds(
|
||||
&mut self,
|
||||
embed_data: &mut [f32],
|
||||
hidden_size: usize,
|
||||
image_mask: &Tensor,
|
||||
pixel_values: &Tensor,
|
||||
image_grid_thw: &Tensor,
|
||||
) -> Result<()> {
|
||||
let (vision_embeds, vision_rows, vision_hidden) =
|
||||
self.run_vision_encoder(pixel_values, image_grid_thw)?;
|
||||
if vision_hidden != hidden_size {
|
||||
return Err(anyhow!(
|
||||
"glm-ocr onnx vision hidden size mismatch: vision={}, text={}",
|
||||
vision_hidden,
|
||||
hidden_size
|
||||
));
|
||||
}
|
||||
|
||||
let image_positions = image_mask
|
||||
.squeeze(0)?
|
||||
.to_dtype(DType::U8)?
|
||||
.to_vec1::<u8>()?
|
||||
.into_iter()
|
||||
.enumerate()
|
||||
.filter_map(|(idx, value)| (value == 1).then_some(idx))
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
if image_positions.len() != vision_rows {
|
||||
return Err(anyhow!(
|
||||
"glm-ocr onnx image token/vision embed mismatch: image_tokens={}, vision_embeds={}",
|
||||
image_positions.len(),
|
||||
vision_rows
|
||||
));
|
||||
}
|
||||
|
||||
for (row_idx, token_idx) in image_positions.into_iter().enumerate() {
|
||||
let src_start = row_idx * hidden_size;
|
||||
let dst_start = token_idx * hidden_size;
|
||||
let src_end = src_start + hidden_size;
|
||||
let dst_end = dst_start + hidden_size;
|
||||
embed_data[dst_start..dst_end].copy_from_slice(&vision_embeds[src_start..src_end]);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn run_vision_encoder(
|
||||
&mut self,
|
||||
pixel_values: &Tensor,
|
||||
image_grid_thw: &Tensor,
|
||||
) -> Result<(Vec<f32>, usize, usize)> {
|
||||
let raw_pixel_values_shape = pixel_values
|
||||
.dims()
|
||||
.iter()
|
||||
.map(|dim| *dim as i64)
|
||||
.collect::<Vec<_>>();
|
||||
let pixel_values_data = pixel_values
|
||||
.flatten_all()?
|
||||
.to_dtype(DType::F32)?
|
||||
.to_vec1::<f32>()?;
|
||||
let raw_image_grid_shape = image_grid_thw
|
||||
.dims()
|
||||
.iter()
|
||||
.map(|dim| *dim as i64)
|
||||
.collect::<Vec<_>>();
|
||||
let image_grid_data = image_grid_thw
|
||||
.flatten_all()?
|
||||
.to_vec1::<u32>()?
|
||||
.into_iter()
|
||||
.map(|value| value as i64)
|
||||
.collect::<Vec<_>>();
|
||||
|
||||
let mut vision_inputs = Vec::with_capacity(self.vision_input_descriptors.len());
|
||||
for desc in &self.vision_input_descriptors {
|
||||
let value = match desc.name.as_str() {
|
||||
"pixel_values" => {
|
||||
let pixel_values_shape =
|
||||
align_onnx_input_shape(desc, raw_pixel_values_shape.clone());
|
||||
build_float_input(desc, pixel_values_shape, &pixel_values_data)?
|
||||
}
|
||||
"image_grid_thw" => {
|
||||
let image_grid_shape =
|
||||
align_onnx_input_shape(desc, raw_image_grid_shape.clone());
|
||||
build_i64_like_input(desc, image_grid_shape, &image_grid_data)?
|
||||
}
|
||||
_ => build_zero_input(desc)?,
|
||||
};
|
||||
vision_inputs.push((desc.name.clone(), value));
|
||||
}
|
||||
|
||||
let outputs = self.vision_session.run(vision_inputs)?;
|
||||
let output_value = outputs
|
||||
.get("image_embeds")
|
||||
.or_else(|| outputs.get("vision_embeds"))
|
||||
.or_else(|| {
|
||||
self.vision_output_names
|
||||
.first()
|
||||
.and_then(|name| outputs.get(name))
|
||||
})
|
||||
.ok_or_else(|| anyhow!("glm-ocr onnx vision output not found"))?;
|
||||
|
||||
if let Ok((shape, values)) = output_value.try_extract_tensor::<f32>() {
|
||||
let shape_vec = shape.iter().copied().collect::<Vec<_>>();
|
||||
return extract_vision_output(shape_vec.as_slice(), values.to_vec());
|
||||
}
|
||||
if let Ok((shape, values)) = output_value.try_extract_tensor::<f16>() {
|
||||
let shape_vec = shape.iter().copied().collect::<Vec<_>>();
|
||||
let values = values
|
||||
.iter()
|
||||
.map(|value| value.to_f32())
|
||||
.collect::<Vec<_>>();
|
||||
return extract_vision_output(shape_vec.as_slice(), values);
|
||||
}
|
||||
Err(anyhow!(
|
||||
"glm-ocr onnx vision output must be a f32/f16 tensor"
|
||||
))
|
||||
}
|
||||
|
||||
fn past_seq_len(&self) -> usize {
|
||||
self.cache_values
|
||||
.iter()
|
||||
.find(|entry| entry.name.ends_with(".key"))
|
||||
.map(|entry| entry.dims.get(2).copied().unwrap_or_default() as usize)
|
||||
.unwrap_or(0)
|
||||
}
|
||||
|
||||
fn build_attention_mask(&self, current_seq_len: usize) -> Result<Vec<i64>> {
|
||||
let total_len = self.past_seq_len().saturating_add(current_seq_len);
|
||||
if total_len == 0 {
|
||||
return Err(anyhow!("glm-ocr onnx attention length cannot be zero"));
|
||||
}
|
||||
Ok(vec![1_i64; total_len])
|
||||
}
|
||||
|
||||
fn compute_prefill_position_ids(
|
||||
&mut self,
|
||||
image_mask: &Tensor,
|
||||
grid_thw: &Tensor,
|
||||
seq_len: usize,
|
||||
) -> Result<Vec<i64>> {
|
||||
let t_dim = grid_thw.i(0)?.to_dtype(DType::F32)?.to_scalar::<f32>()? as usize;
|
||||
let h_dim = grid_thw.i(1)?.to_dtype(DType::F32)?.to_scalar::<f32>()? as usize;
|
||||
let w_dim = grid_thw.i(2)?.to_dtype(DType::F32)?.to_scalar::<f32>()? as usize;
|
||||
|
||||
let llm_grid_t = t_dim;
|
||||
let llm_grid_h = h_dim / self.spatial_merge_size;
|
||||
let llm_grid_w = w_dim / self.spatial_merge_size;
|
||||
let num_image_tokens = llm_grid_t * llm_grid_h * llm_grid_w;
|
||||
|
||||
let mask_vec = image_mask
|
||||
.squeeze(0)?
|
||||
.to_dtype(DType::U8)?
|
||||
.to_vec1::<u8>()?;
|
||||
let mut t_ids = Vec::with_capacity(seq_len);
|
||||
let mut h_ids = Vec::with_capacity(seq_len);
|
||||
let mut w_ids = Vec::with_capacity(seq_len);
|
||||
let mut st_idx: i64 = 0;
|
||||
let mut i = 0usize;
|
||||
|
||||
while i < seq_len {
|
||||
let is_img = mask_vec[i] == 1;
|
||||
let start = i;
|
||||
while i < seq_len && (mask_vec[i] == 1) == is_img {
|
||||
i += 1;
|
||||
}
|
||||
let run_len = i - start;
|
||||
|
||||
if is_img {
|
||||
if run_len != num_image_tokens {
|
||||
return Err(anyhow!(
|
||||
"glm-ocr onnx image token count mismatch: mask={}, grid={}",
|
||||
run_len,
|
||||
num_image_tokens
|
||||
));
|
||||
}
|
||||
for ti in 0..llm_grid_t {
|
||||
for hi in 0..llm_grid_h {
|
||||
for wi in 0..llm_grid_w {
|
||||
t_ids.push(ti as i64 + st_idx);
|
||||
h_ids.push(hi as i64 + st_idx);
|
||||
w_ids.push(wi as i64 + st_idx);
|
||||
}
|
||||
}
|
||||
}
|
||||
let max_offset = (llm_grid_t as i64 - 1)
|
||||
.max(llm_grid_h as i64 - 1)
|
||||
.max(llm_grid_w as i64 - 1);
|
||||
st_idx += max_offset + 1;
|
||||
} else {
|
||||
for j in 0..run_len {
|
||||
let pos = st_idx + j as i64;
|
||||
t_ids.push(pos);
|
||||
h_ids.push(pos);
|
||||
w_ids.push(pos);
|
||||
}
|
||||
st_idx += run_len as i64;
|
||||
}
|
||||
}
|
||||
|
||||
self.next_mrope_pos = st_idx as usize;
|
||||
self.prefill_seq_len = seq_len;
|
||||
|
||||
let mut out = Vec::with_capacity(seq_len * 3);
|
||||
out.extend_from_slice(&t_ids);
|
||||
out.extend_from_slice(&h_ids);
|
||||
out.extend_from_slice(&w_ids);
|
||||
Ok(out)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn build_text_position_ids(position_start: usize, seq_len: usize) -> Vec<i64> {
|
||||
let positions = (position_start..position_start + seq_len)
|
||||
.map(|idx| idx as i64)
|
||||
.collect::<Vec<_>>();
|
||||
let mut ids = Vec::with_capacity(seq_len * 3);
|
||||
for _ in 0..3 {
|
||||
ids.extend_from_slice(&positions);
|
||||
}
|
||||
ids
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn align_onnx_input_shape(desc: &OnnxInputDescriptor, actual_shape: Vec<i64>) -> Vec<i64> {
|
||||
if desc.shape.is_empty() || desc.shape.len() <= actual_shape.len() {
|
||||
return actual_shape;
|
||||
}
|
||||
|
||||
let mut aligned = actual_shape;
|
||||
while aligned.len() < desc.shape.len() {
|
||||
aligned.insert(0, 1);
|
||||
}
|
||||
aligned
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn build_i64_scalar_input(desc: &OnnxInputDescriptor, value: i64) -> Result<ort::value::DynValue> {
|
||||
match desc.kind {
|
||||
Some(OnnxTensorKind::I32) => {
|
||||
Ok(ort::value::Tensor::from_array((vec![1_i64], vec![value as i32]))?.into_dyn())
|
||||
}
|
||||
Some(OnnxTensorKind::I64) | None => {
|
||||
Ok(ort::value::Tensor::from_array((vec![1_i64], vec![value]))?.into_dyn())
|
||||
}
|
||||
_ => Err(anyhow!(
|
||||
"unsupported glm-ocr onnx scalar input dtype for {}",
|
||||
desc.name
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn build_i64_like_input(
|
||||
desc: &OnnxInputDescriptor,
|
||||
shape: Vec<i64>,
|
||||
data: &[i64],
|
||||
) -> Result<ort::value::DynValue> {
|
||||
match desc.kind {
|
||||
Some(OnnxTensorKind::I32) => Ok(ort::value::Tensor::from_array((
|
||||
shape,
|
||||
data.iter().map(|value| *value as i32).collect::<Vec<_>>(),
|
||||
))?
|
||||
.into_dyn()),
|
||||
Some(OnnxTensorKind::I64) | None => {
|
||||
Ok(ort::value::Tensor::from_array((shape, data.to_vec()))?.into_dyn())
|
||||
}
|
||||
Some(OnnxTensorKind::Bool) => Ok(ort::value::Tensor::from_array((
|
||||
shape,
|
||||
data.iter().map(|value| *value != 0).collect::<Vec<_>>(),
|
||||
))?
|
||||
.into_dyn()),
|
||||
_ => Err(anyhow!(
|
||||
"unsupported glm-ocr onnx integer-like input dtype for {}",
|
||||
desc.name
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn build_float_input(
|
||||
desc: &OnnxInputDescriptor,
|
||||
shape: Vec<i64>,
|
||||
data: &[f32],
|
||||
) -> Result<ort::value::DynValue> {
|
||||
match desc.kind {
|
||||
Some(OnnxTensorKind::F16) => Ok(ort::value::Tensor::from_array((
|
||||
shape,
|
||||
data.iter()
|
||||
.map(|value| f16::from_f32(*value))
|
||||
.collect::<Vec<_>>(),
|
||||
))?
|
||||
.into_dyn()),
|
||||
Some(OnnxTensorKind::F32) | None => {
|
||||
Ok(ort::value::Tensor::from_array((shape, data.to_vec()))?.into_dyn())
|
||||
}
|
||||
_ => Err(anyhow!(
|
||||
"unsupported glm-ocr onnx float input dtype for {}",
|
||||
desc.name
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn build_zero_cache_input(desc: &OnnxInputDescriptor) -> Result<ort::value::DynValue> {
|
||||
let shape = desc
|
||||
.shape
|
||||
.iter()
|
||||
.enumerate()
|
||||
.map(|(idx, dim)| {
|
||||
if *dim >= 0 {
|
||||
*dim
|
||||
} else if idx == 2 {
|
||||
0
|
||||
} else {
|
||||
1
|
||||
}
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
build_zero_input_with_shape(desc, shape)
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn build_zero_input(desc: &OnnxInputDescriptor) -> Result<ort::value::DynValue> {
|
||||
let shape = desc
|
||||
.shape
|
||||
.iter()
|
||||
.map(|dim| if *dim < 0 { 1 } else { *dim })
|
||||
.collect::<Vec<_>>();
|
||||
build_zero_input_with_shape(desc, shape)
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn build_zero_input_with_shape(
|
||||
desc: &OnnxInputDescriptor,
|
||||
shape: Vec<i64>,
|
||||
) -> Result<ort::value::DynValue> {
|
||||
let kind = desc
|
||||
.kind
|
||||
.ok_or_else(|| anyhow!("unsupported glm-ocr onnx input dtype for {}", desc.name))?;
|
||||
let elem_count = shape.iter().try_fold(1_usize, |acc, dim| {
|
||||
if *dim < 0 {
|
||||
Err(anyhow!(
|
||||
"cannot resolve dynamic onnx shape for input {}: {:?}",
|
||||
desc.name,
|
||||
shape
|
||||
))
|
||||
} else {
|
||||
Ok(acc.saturating_mul(*dim as usize))
|
||||
}
|
||||
})?;
|
||||
let has_zero_dim = shape.contains(&0);
|
||||
let make_ndarray = || {
|
||||
let dims = shape.iter().map(|dim| *dim as usize).collect::<Vec<_>>();
|
||||
IxDyn(&dims)
|
||||
};
|
||||
let value = match kind {
|
||||
OnnxTensorKind::Bool => {
|
||||
if has_zero_dim {
|
||||
let arr = Array::from_shape_vec(make_ndarray(), vec![false; elem_count])?;
|
||||
ort::value::Tensor::from_array(arr)?.into_dyn()
|
||||
} else {
|
||||
ort::value::Tensor::from_array((shape, vec![false; elem_count]))?.into_dyn()
|
||||
}
|
||||
}
|
||||
OnnxTensorKind::I32 => {
|
||||
if has_zero_dim {
|
||||
let arr = Array::from_shape_vec(make_ndarray(), vec![0_i32; elem_count])?;
|
||||
ort::value::Tensor::from_array(arr)?.into_dyn()
|
||||
} else {
|
||||
ort::value::Tensor::from_array((shape, vec![0_i32; elem_count]))?.into_dyn()
|
||||
}
|
||||
}
|
||||
OnnxTensorKind::I64 => {
|
||||
if has_zero_dim {
|
||||
let arr = Array::from_shape_vec(make_ndarray(), vec![0_i64; elem_count])?;
|
||||
ort::value::Tensor::from_array(arr)?.into_dyn()
|
||||
} else {
|
||||
ort::value::Tensor::from_array((shape, vec![0_i64; elem_count]))?.into_dyn()
|
||||
}
|
||||
}
|
||||
OnnxTensorKind::F16 => {
|
||||
if has_zero_dim {
|
||||
let arr =
|
||||
Array::from_shape_vec(make_ndarray(), vec![f16::from_f32(0.0); elem_count])?;
|
||||
ort::value::Tensor::from_array(arr)?.into_dyn()
|
||||
} else {
|
||||
ort::value::Tensor::from_array((shape, vec![f16::from_f32(0.0); elem_count]))?
|
||||
.into_dyn()
|
||||
}
|
||||
}
|
||||
OnnxTensorKind::F32 => {
|
||||
if has_zero_dim {
|
||||
let arr = Array::from_shape_vec(make_ndarray(), vec![0_f32; elem_count])?;
|
||||
ort::value::Tensor::from_array(arr)?.into_dyn()
|
||||
} else {
|
||||
ort::value::Tensor::from_array((shape, vec![0_f32; elem_count]))?.into_dyn()
|
||||
}
|
||||
}
|
||||
};
|
||||
Ok(value)
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn extract_last_logits(value: &ort::value::DynValue) -> Result<Vec<f32>> {
|
||||
if let Ok((shape, values)) = value.try_extract_tensor::<f32>() {
|
||||
return extract_last_logits_from_shape(
|
||||
shape.iter().copied().collect::<Vec<_>>().as_slice(),
|
||||
values,
|
||||
);
|
||||
}
|
||||
if let Ok((shape, values)) = value.try_extract_tensor::<f16>() {
|
||||
let values = values
|
||||
.iter()
|
||||
.map(|value| value.to_f32())
|
||||
.collect::<Vec<_>>();
|
||||
return extract_last_logits_from_shape(
|
||||
shape.iter().copied().collect::<Vec<_>>().as_slice(),
|
||||
&values,
|
||||
);
|
||||
}
|
||||
Err(anyhow!(
|
||||
"glm-ocr onnx logits output must be a f32/f16 tensor"
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn extract_last_logits_from_shape(shape: &[i64], values: &[f32]) -> Result<Vec<f32>> {
|
||||
match shape {
|
||||
[1, seq_len, vocab_size] => {
|
||||
let seq_len = *seq_len as usize;
|
||||
let vocab_size = *vocab_size as usize;
|
||||
let start = seq_len
|
||||
.checked_sub(1)
|
||||
.ok_or_else(|| anyhow!("glm-ocr onnx logits sequence is empty"))?
|
||||
* vocab_size;
|
||||
Ok(values[start..start + vocab_size].to_vec())
|
||||
}
|
||||
[seq_len, vocab_size] => {
|
||||
let seq_len = *seq_len as usize;
|
||||
let vocab_size = *vocab_size as usize;
|
||||
let start = seq_len
|
||||
.checked_sub(1)
|
||||
.ok_or_else(|| anyhow!("glm-ocr onnx logits sequence is empty"))?
|
||||
* vocab_size;
|
||||
Ok(values[start..start + vocab_size].to_vec())
|
||||
}
|
||||
_ => Err(anyhow!("unexpected glm-ocr onnx logits shape: {:?}", shape)),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "onnx-runtime")]
|
||||
fn extract_vision_output(shape: &[i64], values: Vec<f32>) -> Result<(Vec<f32>, usize, usize)> {
|
||||
match shape {
|
||||
[rows, hidden] => Ok((values, *rows as usize, *hidden as usize)),
|
||||
[1, rows, hidden] => Ok((values, *rows as usize, *hidden as usize)),
|
||||
_ => Err(anyhow!(
|
||||
"unexpected glm-ocr onnx vision output shape: {:?}",
|
||||
shape
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(test, feature = "onnx-runtime"))]
|
||||
mod tests {
|
||||
use super::{OnnxInputDescriptor, OnnxTensorKind, align_onnx_input_shape};
|
||||
|
||||
#[test]
|
||||
fn align_onnx_input_shape_prepends_batch_dimension_when_descriptor_has_higher_rank() {
|
||||
let desc = OnnxInputDescriptor {
|
||||
name: "image_grid_thw".to_string(),
|
||||
shape: vec![-1, 3],
|
||||
kind: Some(OnnxTensorKind::I64),
|
||||
};
|
||||
assert_eq!(align_onnx_input_shape(&desc, vec![3]), vec![1, 3]);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "onnx-runtime"))]
|
||||
pub struct GlmOcrOnnxBackend;
|
||||
|
||||
#[cfg(not(feature = "onnx-runtime"))]
|
||||
impl GlmOcrOnnxBackend {
|
||||
pub fn load(_onnx_path: &str, _spatial_merge_size: usize) -> Result<Self> {
|
||||
Err(anyhow!(
|
||||
"onnx runtime support is not enabled; rebuild with --features onnx-runtime"
|
||||
))
|
||||
}
|
||||
|
||||
pub fn clear_cache(&mut self) {}
|
||||
|
||||
pub fn forward_logits(
|
||||
&mut self,
|
||||
_input_ids: &[u32],
|
||||
_image_mask: Option<&Tensor>,
|
||||
_pixel_values: Option<&Tensor>,
|
||||
_image_grid_thw: Option<&Tensor>,
|
||||
_position_start: usize,
|
||||
) -> Result<Vec<f32>> {
|
||||
Err(anyhow!(
|
||||
"onnx runtime support is not enabled; rebuild with --features onnx-runtime"
|
||||
))
|
||||
}
|
||||
}
|
||||
@@ -63,6 +63,30 @@ impl GlmOcrProcessor {
|
||||
})
|
||||
}
|
||||
|
||||
pub fn from_params(
|
||||
image_mean: Vec<f32>,
|
||||
image_std: Vec<f32>,
|
||||
shortest_edge: usize,
|
||||
longest_edge: usize,
|
||||
patch_size: usize,
|
||||
merge_size: usize,
|
||||
temporal_patch_size: usize,
|
||||
device: &Device,
|
||||
dtype: DType,
|
||||
) -> Self {
|
||||
Self {
|
||||
image_mean,
|
||||
image_std,
|
||||
shortest_edge,
|
||||
longest_edge,
|
||||
patch_size,
|
||||
merge_size,
|
||||
temporal_patch_size,
|
||||
device: device.clone(),
|
||||
dtype,
|
||||
}
|
||||
}
|
||||
|
||||
/// Process image for vision encoder.
|
||||
///
|
||||
/// Matches Python's Glm46VImageProcessor._preprocess():
|
||||
|
||||
@@ -0,0 +1,220 @@
|
||||
use std::path::{Path, PathBuf};
|
||||
|
||||
use aha::models::{
|
||||
ArtifactKind, GenerateModel, LoadSpec, ModelPaths, WhichModel,
|
||||
glm_ocr::generate::GlmOcrGenerateModel,
|
||||
};
|
||||
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
|
||||
use anyhow::{Context, Result, anyhow};
|
||||
|
||||
const GLM_OCR_GGUF_DIR: &str = r"D:\model_download\GLM-OCR-GGUF";
|
||||
const GLM_OCR_ONNX_DIR: &str = r"D:\model_download\GLM-OCR-ONNX";
|
||||
|
||||
fn require_existing_dir(path: &str) -> Result<()> {
|
||||
let dir = Path::new(path);
|
||||
if !dir.exists() {
|
||||
return Err(anyhow!("model dir not found: {}", path));
|
||||
}
|
||||
if !dir.is_dir() {
|
||||
return Err(anyhow!("path is not a directory: {}", path));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn first_file_with_extension(dir: &str, extension: &str) -> Result<PathBuf> {
|
||||
require_existing_dir(dir)?;
|
||||
let mut matches = std::fs::read_dir(dir)?
|
||||
.flatten()
|
||||
.map(|entry| entry.path())
|
||||
.filter(|path| {
|
||||
path.is_file()
|
||||
&& path
|
||||
.extension()
|
||||
.is_some_and(|ext| ext.eq_ignore_ascii_case(extension))
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
matches.sort();
|
||||
matches
|
||||
.into_iter()
|
||||
.next()
|
||||
.ok_or_else(|| anyhow!("no .{} file found in {}", extension, dir))
|
||||
}
|
||||
|
||||
fn first_file_with_extension_recursive(dir: &str, extension: &str) -> Result<PathBuf> {
|
||||
require_existing_dir(dir)?;
|
||||
|
||||
let mut stack = vec![PathBuf::from(dir)];
|
||||
let mut matches = Vec::new();
|
||||
while let Some(current) = stack.pop() {
|
||||
for entry in std::fs::read_dir(¤t)? {
|
||||
let entry = entry?;
|
||||
let path = entry.path();
|
||||
if path.is_dir() {
|
||||
stack.push(path);
|
||||
continue;
|
||||
}
|
||||
if path
|
||||
.extension()
|
||||
.is_some_and(|ext| ext.eq_ignore_ascii_case(extension))
|
||||
{
|
||||
matches.push(path);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
matches.sort();
|
||||
matches
|
||||
.into_iter()
|
||||
.next()
|
||||
.ok_or_else(|| anyhow!("no .{} file found (recursive) in {}", extension, dir))
|
||||
}
|
||||
|
||||
fn ocr_test_image_url() -> Result<String> {
|
||||
let path = std::env::current_dir()?
|
||||
.join("assets")
|
||||
.join("img")
|
||||
.join("ocr_test1.png");
|
||||
let path = path.canonicalize()?;
|
||||
url::Url::from_file_path(&path)
|
||||
.map(|url| url.to_string())
|
||||
.map_err(|_| anyhow!("failed to build file url for {}", path.display()))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn glm_ocr_gguf_files_can_load() -> Result<()> {
|
||||
let gguf_path = first_file_with_extension(GLM_OCR_GGUF_DIR, "gguf")?;
|
||||
let metadata = std::fs::metadata(&gguf_path)
|
||||
.with_context(|| format!("failed to read gguf metadata: {}", gguf_path.display()))?;
|
||||
if metadata.len() == 0 {
|
||||
return Err(anyhow!("gguf file is empty: {}", gguf_path.display()));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn glm_ocr_onnx_files_can_load() -> Result<()> {
|
||||
let onnx_path = first_file_with_extension_recursive(GLM_OCR_ONNX_DIR, "onnx")?;
|
||||
let metadata = std::fs::metadata(&onnx_path)
|
||||
.with_context(|| format!("failed to read onnx metadata: {}", onnx_path.display()))?;
|
||||
if metadata.len() == 0 {
|
||||
return Err(anyhow!("onnx file is empty: {}", onnx_path.display()));
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn glm_ocr_load_spec_accepts_gguf() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::GlmOCR,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some(GLM_OCR_GGUF_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
spec.validate()
|
||||
.expect("glm-ocr should accept gguf artifact");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn glm_ocr_load_spec_accepts_onnx() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::GlmOCR,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some(GLM_OCR_ONNX_DIR.to_string()),
|
||||
tokenizer_dir: Some(GLM_OCR_ONNX_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
spec.validate()
|
||||
.expect("glm-ocr should accept onnx artifact");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn glm_ocr_gguf_init_from_spec_can_init() -> Result<()> {
|
||||
if let Err(err) = require_existing_dir(GLM_OCR_GGUF_DIR) {
|
||||
println!("skip glm-ocr gguf init test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::GlmOCR,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some(GLM_OCR_GGUF_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
GlmOcrGenerateModel::init_from_spec(&spec, None, None)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn glm_ocr_onnx_init_from_spec_can_init() -> Result<()> {
|
||||
if let Err(err) = require_existing_dir(GLM_OCR_ONNX_DIR) {
|
||||
println!("skip glm-ocr onnx init test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::GlmOCR,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some(GLM_OCR_ONNX_DIR.to_string()),
|
||||
tokenizer_dir: Some(GLM_OCR_ONNX_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
GlmOcrGenerateModel::init_from_spec(&spec, None, None)?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[ignore = "manual real-model inference smoke test"]
|
||||
fn glm_ocr_gguf_generate_smoke() -> Result<()> {
|
||||
if let Err(err) = require_existing_dir(GLM_OCR_GGUF_DIR) {
|
||||
println!("skip glm-ocr gguf generate smoke test: {err}");
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::GlmOCR,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some(GLM_OCR_GGUF_DIR.to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
let ocr_test_image = ocr_test_image_url()?;
|
||||
let message = format!(
|
||||
r#"{{
|
||||
"model": "glm-ocr",
|
||||
"messages": [
|
||||
{{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{{
|
||||
"type": "image",
|
||||
"image_url": {{
|
||||
"url": "{ocr_test_image}"
|
||||
}}
|
||||
}},
|
||||
{{
|
||||
"type": "text",
|
||||
"text": "Text Recognition:"
|
||||
}}
|
||||
]
|
||||
}}
|
||||
],
|
||||
"max_tokens": 1
|
||||
}}"#
|
||||
);
|
||||
let mes: ChatCompletionParameters = serde_json::from_str(&message)?;
|
||||
let mut model = GlmOcrGenerateModel::init_from_spec(&spec, None, None)?;
|
||||
let res = model.generate(mes)?;
|
||||
if let Some(usage) = res.usage {
|
||||
assert!(usage.total_tokens > 0);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
@@ -48,6 +48,35 @@ fn load_spec_all_minilm_accepts_gguf() {
|
||||
assert!(spec.validate().is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn load_spec_glm_ocr_accepts_gguf() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::GlmOCR,
|
||||
artifact: ArtifactKind::Gguf,
|
||||
paths: ModelPaths {
|
||||
gguf_path: Some("D:/model_download/GLM-OCR-GGUF".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
assert!(spec.validate().is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn load_spec_glm_ocr_accepts_onnx() {
|
||||
let spec = LoadSpec {
|
||||
model: WhichModel::GlmOCR,
|
||||
artifact: ArtifactKind::Onnx,
|
||||
paths: ModelPaths {
|
||||
onnx_path: Some("D:/model_download/GLM-OCR-ONNX".to_string()),
|
||||
tokenizer_dir: Some("D:/model_download/GLM-OCR-ONNX".to_string()),
|
||||
..Default::default()
|
||||
},
|
||||
};
|
||||
|
||||
assert!(spec.validate().is_ok());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn load_spec_gguf_requires_gguf_path() {
|
||||
let spec = LoadSpec {
|
||||
|
||||
Reference in New Issue
Block a user