给 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
|
## 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
|
||||||
|
|
||||||
|
|||||||
@@ -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-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
|
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) |
|
| **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 |
|
||||||
|
|||||||
@@ -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
@@ -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
@@ -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([
|
||||||
|
|||||||
@@ -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)?)
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -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,
|
||||||
|
|||||||
@@ -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
|
||||||
|
|||||||
+969
-33
File diff suppressed because it is too large
Load Diff
@@ -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;
|
||||||
|
|||||||
@@ -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.
|
/// Process image for vision encoder.
|
||||||
///
|
///
|
||||||
/// Matches Python's Glm46VImageProcessor._preprocess():
|
/// 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());
|
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 {
|
||||||
|
|||||||
Reference in New Issue
Block a user