给 glm-ocr 添加 gguf 和 onnx 格式推理

This commit is contained in:
273265088@qq.com
2026-03-26 17:42:58 +08:00
parent b6b970fd25
commit aec47e47af
15 changed files with 2315 additions and 77 deletions
+7
View File
@@ -27,6 +27,7 @@ aha is a high-performance, cross-platform AI inference engine built with Rust an
## Changelog ## Changelog
### v0.2.3 (2026-03-18) ### v0.2.3 (2026-03-18)
- add DeepSeek-OCR-2 - add DeepSeek-OCR-2
- add GLM-OCR gguf and onnx local loading
### 2026-03-17 ### 2026-03-17
- add PaddleOCR-VL1.5 model - 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) # 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 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) # Start service only (model already downloaded)
aha serv -m qwen3asr-0.6b -p 10100 aha serv -m qwen3asr-0.6b -p 10100
+7
View File
@@ -27,6 +27,7 @@ aha 是一款基于 Rust 和 Candle 框架构建的高性能跨平台 AI 推理
## 更新日志 ## 更新日志
### v0.2.3 (2026-03-18) ### v0.2.3 (2026-03-18)
- 新增 DeepSeek-OCR-2 - 新增 DeepSeek-OCR-2
- 增加 GLM-OCR 的 GGUF 和 ONNX 本地加载
### 2026-03-17 ### 2026-03-17
- 新增 PaddleOCR-VL1.5 模型 - 新增 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 # 本地运行 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 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-OCRGGUF
aha run -m glm-ocr -i .\assets\img\ocr_test1.png --artifact-format gguf --gguf-path D:\model_download\GLM-OCR-GGUF
# 本地运行 GLM-OCRONNX
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 aha serv -m qwen3asr-0.6b -p 10100
+2
View File
@@ -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) | | **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** | 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) ## Speech Recognition (ASR)
| Model | Parameters | Language | Real-time | Speed | License | | Model | Parameters | Language | Real-time | Speed | License |
+2
View File
@@ -63,6 +63,8 @@ aha 支持多个领域的最先进 AI 模型集合。
| **DeepSeek-OCR** | 多语言 | 场景文字 | 自然图像 | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) | | **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** | 8 | 场景文字 | 复杂文档 | [MIT](https://huggingface.co/datasets/choosealicense/licenses/blob/main/markdown/mit.md) |
GLM-OCR 本地制品格式:`safetensors``gguf``onnx`
## 语音识别 (ASR) ## 语音识别 (ASR)
| 模型 | 参数量 | 语言 | 实时 | 速度 | 开源协议 | | 模型 | 参数量 | 语言 | 实时 | 速度 | 开源协议 |
+113 -41
View File
@@ -1,56 +1,84 @@
//! Glm-OCR exec implementation for CLI `run` subcommand //! 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::exec::ExecModel;
use crate::models::{GenerateModel, glm_ocr::generate::GlmOcrGenerateModel}; use crate::models::{GenerateModel, LoadSpec, glm_ocr::generate::GlmOcrGenerateModel};
pub struct GlmOcrExec; pub struct GlmOcrExec;
impl ExecModel for GlmOcrExec { fn resolve_input_url(input: &str) -> Result<String> {
fn run(input: &[String], output: Option<&str>, weight_path: &str) -> Result<()> { if input.starts_with("http://") || input.starts_with("https://") || input.starts_with("file://")
let url = &input[0]; {
let input_url = if url.starts_with("http://") return Ok(input.to_string());
|| url.starts_with("https://") }
|| url.starts_with("file://")
{ let path = Path::new(input);
url.clone() let path = if path.is_absolute() {
} else { path.to_path_buf()
format!("file://{}", url) } 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 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(); let i_duration = i_start.elapsed();
println!("Time elapsed in load model is: {:?}", i_duration); println!("Time elapsed in load model is: {:?}", i_duration);
let message = format!( let mes = build_request(&input_url, max_tokens)?;
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 i_start = Instant::now(); let i_start = Instant::now();
let result = model.generate(mes)?; let result = model.generate(mes)?;
@@ -67,3 +95,47 @@ impl ExecModel for GlmOcrExec {
Ok(()) 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
View File
@@ -236,6 +236,10 @@ struct RunArgs {
#[arg(short, long)] #[arg(short, long)]
output: Option<String>, 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) /// Local model weight path (defaults to ~/.aha/{model_id} if not specified)
#[arg(long)] #[arg(long)]
weight_path: Option<String>, 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)?; Qwen3_5Exec::run_with_spec(&args.input, args.output.as_deref(), spec)?;
Ok(true) 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), _ => 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] #[test]
fn parse_serv_onnx_flags() { fn parse_serv_onnx_flags() {
let cli = Cli::try_parse_from([ let cli = Cli::try_parse_from([
+4
View File
@@ -86,6 +86,10 @@ impl<R: Read + Seek> Gguf<R> {
&self.ct.metadata &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> { pub fn tensor(&mut self, name: &str) -> Result<QTensor> {
Ok(self.ct.tensor(&mut self.reader, name, &self.device)?) Ok(self.ct.tensor(&mut self.reader, name, &self.device)?)
} }
+5
View File
@@ -33,6 +33,11 @@ pub fn supported_artifacts(model: WhichModel) -> &'static [ArtifactKind] {
ArtifactKind::Gguf, ArtifactKind::Gguf,
ArtifactKind::Onnx, ArtifactKind::Onnx,
], ],
WhichModel::GlmOCR => &[
ArtifactKind::Safetensors,
ArtifactKind::Gguf,
ArtifactKind::Onnx,
],
WhichModel::Qwen3_0_6B => &[ WhichModel::Qwen3_0_6B => &[
ArtifactKind::Safetensors, ArtifactKind::Safetensors,
ArtifactKind::Gguf, ArtifactKind::Gguf,
+13 -3
View File
@@ -1,9 +1,10 @@
use anyhow::Result; use anyhow::Result;
use crate::models::{ use crate::models::{
ModelInstance, WhichModel, all_minilm_l6_v2::generate::AllMiniLML6V2Model, load_model_legacy, ModelInstance, WhichModel, all_minilm_l6_v2::generate::AllMiniLML6V2Model,
qwen3::generate::Qwen3GenerateModel, qwen3_5::generate::Qwen3_5GenerateModel, glm_ocr::generate::GlmOcrGenerateModel, load_model_legacy, qwen3::generate::Qwen3GenerateModel,
qwen3_embedding::generate::Qwen3EmbeddingModel, qwen3_reranker::generate::Qwen3RerankerModel, qwen3_5::generate::Qwen3_5GenerateModel, qwen3_embedding::generate::Qwen3EmbeddingModel,
qwen3_reranker::generate::Qwen3RerankerModel,
}; };
use super::artifact::LoadSpec; use super::artifact::LoadSpec;
@@ -11,6 +12,7 @@ use super::artifact::LoadSpec;
#[derive(Debug, Clone, Copy, PartialEq, Eq)] #[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum ModelLoaderFamily { enum ModelLoaderFamily {
AllMiniLML6V2, AllMiniLML6V2,
GlmOCR,
Qwen3, Qwen3,
Qwen3Embedding, Qwen3Embedding,
Qwen3Reranker, Qwen3Reranker,
@@ -21,6 +23,7 @@ enum ModelLoaderFamily {
fn resolve_model_loader_family(model: WhichModel) -> ModelLoaderFamily { fn resolve_model_loader_family(model: WhichModel) -> ModelLoaderFamily {
match model { match model {
WhichModel::AllMiniLML6V2 => ModelLoaderFamily::AllMiniLML6V2, WhichModel::AllMiniLML6V2 => ModelLoaderFamily::AllMiniLML6V2,
WhichModel::GlmOCR => ModelLoaderFamily::GlmOCR,
WhichModel::Qwen3_0_6B => ModelLoaderFamily::Qwen3, WhichModel::Qwen3_0_6B => ModelLoaderFamily::Qwen3,
WhichModel::Qwen3Embedding0_6B WhichModel::Qwen3Embedding0_6B
| WhichModel::Qwen3Embedding4B | WhichModel::Qwen3Embedding4B
@@ -49,6 +52,9 @@ pub fn load_model_from_spec<'a>(spec: &LoadSpec) -> Result<ModelInstance<'a>> {
ModelLoaderFamily::AllMiniLML6V2 => { ModelLoaderFamily::AllMiniLML6V2 => {
ModelInstance::AllMiniLML6V2(AllMiniLML6V2Model::init_from_spec(spec, None, None)?) ModelInstance::AllMiniLML6V2(AllMiniLML6V2Model::init_from_spec(spec, None, None)?)
} }
ModelLoaderFamily::GlmOCR => {
ModelInstance::GlmOCR(GlmOcrGenerateModel::init_from_spec(spec, None, None)?)
}
ModelLoaderFamily::Qwen3 => { ModelLoaderFamily::Qwen3 => {
ModelInstance::Qwen3(Qwen3GenerateModel::init_from_spec(spec, None, None)?) ModelInstance::Qwen3(Qwen3GenerateModel::init_from_spec(spec, None, None)?)
} }
@@ -86,6 +92,10 @@ mod tests {
resolve_model_loader_family(WhichModel::Qwen3_0_6B), resolve_model_loader_family(WhichModel::Qwen3_0_6B),
ModelLoaderFamily::Qwen3 ModelLoaderFamily::Qwen3
); );
assert_eq!(
resolve_model_loader_family(WhichModel::GlmOCR),
ModelLoaderFamily::GlmOCR
);
assert_eq!( assert_eq!(
resolve_model_loader_family(WhichModel::Qwen3Embedding0_6B), resolve_model_loader_family(WhichModel::Qwen3Embedding0_6B),
ModelLoaderFamily::Qwen3Embedding ModelLoaderFamily::Qwen3Embedding
File diff suppressed because it is too large Load Diff
+1
View File
@@ -1,4 +1,5 @@
pub mod config; pub mod config;
pub mod generate; pub mod generate;
pub mod model; pub mod model;
pub mod onnx;
pub mod processor; pub mod processor;
+847
View File
@@ -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(&current)? {
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"
))
}
}
+24
View File
@@ -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. /// Process image for vision encoder.
/// ///
/// Matches Python's Glm46VImageProcessor._preprocess(): /// Matches Python's Glm46VImageProcessor._preprocess():
+220
View File
@@ -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(&current)? {
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(())
}
+29
View File
@@ -48,6 +48,35 @@ fn load_spec_all_minilm_accepts_gguf() {
assert!(spec.validate().is_ok()); 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] #[test]
fn load_spec_gguf_requires_gguf_path() { fn load_spec_gguf_requires_gguf_path() {
let spec = LoadSpec { let spec = LoadSpec {