pub mod batch_scheduler;
pub mod calib;
pub mod connector;
pub mod decoder;
pub mod decoder_qwen2;
pub mod got;
pub mod model_arch;
pub mod moe;
pub mod nn;
pub mod onechart;
pub mod postprocess;
pub mod progress;
pub mod rswa;
pub mod sampler;
pub mod smolvlm2;
pub(crate) mod spec;
pub mod tensor;
pub mod token_compress;
pub mod tromr;
pub(crate) mod unlimited_ocr_census;
pub mod vision_bridge;
pub mod vision_clip;
pub mod vision_sam;
pub mod vision_siglip;
pub mod weights;
use std::ffi::OsStr;
use std::io::Read;
use std::path::{Component, Path, PathBuf};
use std::sync::{Arc, Condvar, Mutex, MutexGuard, OnceLock, Weak};
#[cfg(not(target_arch = "wasm32"))]
use std::time::Instant;
#[cfg(target_arch = "wasm32")]
use web_time::Instant;
use crate::error::{FocrError, FocrResult};
use crate::preprocess::{self, Preprocessed};
use crate::quant::recipe::{Recipe, is_truthy};
use sampler::{DecodeOutput, DecodeParams};
use tensor::Mat;
use weights::{DType, Weights};
struct ExclusiveGate {
occupied: std::sync::atomic::AtomicBool,
waiters: Mutex<()>,
wake: Condvar,
}
impl ExclusiveGate {
const fn new() -> Self {
Self {
occupied: std::sync::atomic::AtomicBool::new(false),
waiters: Mutex::new(()),
wake: Condvar::new(),
}
}
fn acquire(&self) -> ExclusivePermit<'_> {
use std::sync::atomic::Ordering;
if self
.occupied
.compare_exchange(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
{
return ExclusivePermit { gate: self };
}
let mut waiter = self
.waiters
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
loop {
if self
.occupied
.compare_exchange(false, true, Ordering::Acquire, Ordering::Relaxed)
.is_ok()
{
drop(waiter);
return ExclusivePermit { gate: self };
}
waiter = self
.wake
.wait(waiter)
.unwrap_or_else(std::sync::PoisonError::into_inner);
}
}
fn release(&self) {
use std::sync::atomic::Ordering;
let _waiter = self
.waiters
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
self.occupied.store(false, Ordering::Release);
self.wake.notify_one();
}
}
struct ExclusivePermit<'a> {
gate: &'a ExclusiveGate,
}
impl Drop for ExclusivePermit<'_> {
fn drop(&mut self) {
self.gate.release();
}
}
struct FallibleOnce<T> {
value: OnceLock<T>,
init: ExclusiveGate,
}
impl<T> FallibleOnce<T> {
const fn new() -> Self {
Self {
value: OnceLock::new(),
init: ExclusiveGate::new(),
}
}
fn get_or_try_init<E, F>(&self, initialize: F) -> Result<&T, E>
where
F: FnOnce() -> Result<T, E>,
{
if let Some(value) = self.value.get() {
return Ok(value);
}
let _permit = self.init.acquire();
if let Some(value) = self.value.get() {
return Ok(value);
}
let value = initialize()?;
let set = self.value.set(value);
debug_assert!(set.is_ok(), "single initializer set raced unexpectedly");
Ok(self.value.get().expect("single initializer just set"))
}
}
enum RetainedOrOwned<'a, T> {
Retained(&'a T),
Owned(T),
}
impl<T> AsRef<T> for RetainedOrOwned<'_, T> {
fn as_ref(&self) -> &T {
match self {
Self::Retained(value) => value,
Self::Owned(value) => value,
}
}
}
fn retained_or_owned<'a, T, E, F>(
retain: bool,
cache: &'a FallibleOnce<T>,
initialize: F,
) -> Result<RetainedOrOwned<'a, T>, E>
where
F: FnOnce() -> Result<T, E>,
{
if retain {
cache
.get_or_try_init(initialize)
.map(RetainedOrOwned::Retained)
} else {
initialize().map(RetainedOrOwned::Owned)
}
}
#[derive(Default)]
pub struct SidecarBundle {
pub tokenizer_json: Option<Vec<u8>>,
pub qwen_tiktoken: Option<Vec<u8>>,
pub music_tables: Option<[String; 4]>,
pub opt_triple: Option<[Vec<u8>; 3]>,
}
pub struct OcrModel {
path: PathBuf,
sidecars: SidecarBundle,
weights: Weights,
decode_params: DecodeParams,
decoder_cache_i8: FallibleOnce<decoder::DecoderWeightCacheI8>,
decoder_cache: FallibleOnce<decoder::DecoderWeightCache>,
clip_cache: FallibleOnce<vision_clip::ClipWeights>,
unlimited_vision: FallibleOnce<UnlimitedVisionStatics>,
got_statics: FallibleOnce<got::GotStatics>,
onechart_statics: FallibleOnce<onechart::OnechartStatics>,
smol_statics: FallibleOnce<smolvlm2::SmolStatics>,
tokenizer: FallibleOnce<crate::tokenizer::Tokenizer>,
got_tokenizer: FallibleOnce<crate::tokenizer::tiktoken::Tiktoken>,
tromr_tokenizer: FallibleOnce<crate::tokenizer::music::MusicTokenizer>,
music_meta: std::sync::Mutex<Option<MusicPageMeta>>,
}
struct UnlimitedVisionStatics {
sam: vision_sam::SamWeights,
projector: vision_bridge::ProjectorWeights,
}
type ModelCacheEntry = Option<(PathBuf, Weak<OcrModel>)>;
type ModelCache = Mutex<ModelCacheEntry>;
const RAW_SAFETENSORS_SHARD_NAME: &str = "model-00001-of-000001.safetensors";
const MODEL_HEADER_VALIDATION_MAX_BYTES: usize = 8 * 1024 * 1024;
const MODEL_DIR_ENV: &str = "FOCR_MODEL_DIR";
const MODEL_QUANT_ENV: &str = "FOCR_QUANT";
const BASE_PROMPT_TEXT: &str = "document parsing.";
const MULTI_PAGE_PROMPT_TEXT: &str = "Multi page parsing.";
const MULTI_PAGE_BASE_SIZE: usize = 640;
const MAX_POSITION_EMBEDDINGS: usize = 32768;
type PageSink<'a> = &'a mut dyn FnMut(usize, &str);
const MAX_NEW_TOKENS_ENV: &str = "FOCR_MAX_NEW_TOKENS";
const DECODE_STATELESS_ENV: &str = "FOCR_DECODE_STATELESS";
const UNLIMITED_VISION_CACHE_ENV: &str = "FOCR_UNLIMITED_VISION_CACHE";
const STREAMED_VIEW_CHUNK: usize = 8;
fn unlimited_vision_cache_enabled_for(value: Option<&str>) -> bool {
!matches!(
value
.map(|value| value.trim().to_ascii_lowercase())
.as_deref(),
Some("0" | "off" | "false" | "no")
)
}
fn unlimited_vision_cache_enabled() -> bool {
static ENABLED: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*ENABLED.get_or_init(|| {
let value = std::env::var(UNLIMITED_VISION_CACHE_ENV).ok();
unlimited_vision_cache_enabled_for(value.as_deref())
})
}
fn stream_vision_tower() -> bool {
static STREAM: std::sync::OnceLock<bool> = std::sync::OnceLock::new();
*STREAM.get_or_init(|| match std::env::var("FOCR_STREAM_VISION").ok() {
Some(value) => matches!(
value.trim().to_ascii_lowercase().as_str(),
"1" | "on" | "true" | "yes"
),
None => cfg!(target_os = "ios"),
})
}
const DECODE_INT8_ENV: &str = "FOCR_DECODE_INT8";
static FORCE_INT8_DECODE: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
pub fn force_int8_decode(on: bool) -> FocrResult<()> {
if on {
require_experimental_full_int8_recipe()?;
}
FORCE_INT8_DECODE.store(on, std::sync::atomic::Ordering::Relaxed);
Ok(())
}
fn int8_decode_requested() -> bool {
FORCE_INT8_DECODE.load(std::sync::atomic::Ordering::Relaxed)
|| std::env::var(DECODE_INT8_ENV).is_ok_and(|value| is_truthy(&value))
}
#[must_use]
pub fn experimental_full_int8_decode_requested() -> bool {
int8_decode_requested()
}
fn require_experimental_full_int8_recipe() -> FocrResult<()> {
let recipe = Recipe::from_env();
if recipe.attn_int8() && recipe.lmhead_int8() {
return Ok(());
}
Err(FocrError::Usage(format!(
"{DECODE_INT8_ENV}=1 requests the experimental all-int8 decoder; it also requires \
{}=1 and {}=1 because attention q/k/v/o and lm_head are outside the validated \
default recipe",
crate::quant::recipe::FOCR_INT8_ATTN_ENV,
crate::quant::recipe::FOCR_INT8_LMHEAD_ENV,
)))
}
pub fn validate_experimental_full_int8_decode() -> FocrResult<()> {
if int8_decode_requested() {
require_experimental_full_int8_recipe()?;
}
Ok(())
}
const GOT_FORMAT_ENV: &str = "FOCR_GOT_FORMAT";
static FORCE_GOT_FORMAT: std::sync::atomic::AtomicBool = std::sync::atomic::AtomicBool::new(false);
pub fn force_got_format(on: bool) {
FORCE_GOT_FORMAT.store(on, std::sync::atomic::Ordering::Relaxed);
}
fn got_format_requested() -> bool {
FORCE_GOT_FORMAT.load(std::sync::atomic::Ordering::Relaxed)
|| std::env::var_os(GOT_FORMAT_ENV).is_some()
}
const SMOLVLM2_QUESTION_ENV: &str = "FOCR_SMOLVLM2_QUESTION";
static SMOLVLM2_QUESTION: std::sync::Mutex<Option<String>> = std::sync::Mutex::new(None);
pub fn set_smolvlm2_question(question: Option<String>) {
*SMOLVLM2_QUESTION
.lock()
.expect("smolvlm2 question mutex poisoned") = question;
}
fn smolvlm2_question() -> String {
if let Some(q) = SMOLVLM2_QUESTION
.lock()
.expect("smolvlm2 question mutex poisoned")
.clone()
{
return q;
}
std::env::var(SMOLVLM2_QUESTION_ENV).unwrap_or_else(|_| smolvlm2::DESCRIBE_QUESTION.to_string())
}
static FORWARDS_LIVE: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
static FORWARDS_MAX: std::sync::atomic::AtomicUsize = std::sync::atomic::AtomicUsize::new(0);
static FORWARD_UNDER_GUARD: std::sync::atomic::AtomicBool =
std::sync::atomic::AtomicBool::new(false);
thread_local! {
static CACHE_GUARD_HELD_HERE: std::cell::Cell<usize> = const { std::cell::Cell::new(0) };
}
pub(crate) struct ForwardPass(());
impl Drop for ForwardPass {
fn drop(&mut self) {
FORWARDS_LIVE.fetch_sub(1, std::sync::atomic::Ordering::SeqCst);
}
}
pub(crate) fn enter_forward() -> ForwardPass {
let live = FORWARDS_LIVE.fetch_add(1, std::sync::atomic::Ordering::SeqCst) + 1;
FORWARDS_MAX.fetch_max(live, std::sync::atomic::Ordering::SeqCst);
if CACHE_GUARD_HELD_HERE.with(std::cell::Cell::get) > 0 {
FORWARD_UNDER_GUARD.store(true, std::sync::atomic::Ordering::SeqCst);
}
ForwardPass(())
}
#[doc(hidden)]
pub fn forward_gauge_take() -> (usize, bool) {
let max = FORWARDS_MAX.swap(0, std::sync::atomic::Ordering::SeqCst);
let under = FORWARD_UNDER_GUARD.swap(false, std::sync::atomic::Ordering::SeqCst);
(max, under)
}
const BATCH_VISION_ENV: &str = "FOCR_BATCH_VISION";
fn batch_vision_enabled() -> bool {
!matches!(
std::env::var(BATCH_VISION_ENV)
.ok()
.map(|v| v.trim().to_ascii_lowercase())
.as_deref(),
Some("0" | "off" | "false" | "no")
)
}
#[derive(Clone, Copy, Debug, Default, PartialEq)]
pub struct DecodeOverrides {
pub max_length: Option<usize>,
pub temperature: Option<f32>,
pub no_repeat_ngram: Option<usize>,
pub ngram_window: Option<usize>,
}
#[derive(Clone, Copy, Debug, Default, PartialEq)]
pub struct PreprocessOverrides {
pub base_size: Option<usize>,
pub image_size: Option<usize>,
pub gundam: Option<bool>,
}
static PREPROCESS_OVERRIDES: Mutex<PreprocessOverrides> = Mutex::new(PreprocessOverrides {
base_size: None,
image_size: None,
gundam: None,
});
pub fn set_preprocess_overrides(overrides: PreprocessOverrides) {
*PREPROCESS_OVERRIDES
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = overrides;
}
fn resolve_preprocess_mode(o: PreprocessOverrides) -> preprocess::PreprocessMode {
let base_size = o.base_size.unwrap_or(1024);
if o.gundam == Some(true) {
preprocess::PreprocessMode::Gundam {
base_size,
tile_size: o.image_size.unwrap_or(640),
}
} else {
preprocess::PreprocessMode::Base { base_size }
}
}
fn preprocess_mode() -> preprocess::PreprocessMode {
resolve_preprocess_mode(
*PREPROCESS_OVERRIDES
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner),
)
}
static DECODE_OVERRIDES: Mutex<DecodeOverrides> = Mutex::new(DecodeOverrides {
max_length: None,
temperature: None,
no_repeat_ngram: None,
ngram_window: None,
});
pub fn set_decode_overrides(overrides: DecodeOverrides) {
*DECODE_OVERRIDES
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner) = overrides;
}
fn decode_overrides() -> DecodeOverrides {
*DECODE_OVERRIDES
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
}
fn apply_decode_overrides(params: &mut DecodeParams, overrides: DecodeOverrides) {
if let Some(n) = overrides.max_length {
params.max_length = n;
}
if let Some(t) = overrides.temperature {
params.temperature = t;
}
if let Some(n) = overrides.no_repeat_ngram {
params.no_repeat_ngram_size = n;
}
if let Some(w) = overrides.ngram_window {
params.ngram_window = w;
}
}
const SPEC_DECODE_ENV: &str = "FOCR_SPEC_DECODE";
fn spec_decode_enabled() -> bool {
static FLAG: OnceLock<bool> = OnceLock::new();
*FLAG.get_or_init(|| std::env::var_os(SPEC_DECODE_ENV).is_some())
}
pub(crate) fn timing_log(msg: &str) {
if std::env::var_os("FOCR_TIMING").is_some() {
crate::progress::stderr_message(format_args!("[focr-timing] {msg}"));
}
}
fn decode_params_from_env() -> DecodeParams {
let mut p = DecodeParams::single_image();
if let Some(raw) = std::env::var_os(MAX_NEW_TOKENS_ENV)
&& let Some(s) = raw.to_str()
&& let Ok(n) = s.trim().parse::<usize>()
&& n > 0
{
p.max_length = n;
}
apply_decode_overrides(&mut p, decode_overrides());
p
}
fn model_cache() -> &'static ModelCache {
static CACHE: OnceLock<ModelCache> = OnceLock::new();
CACHE.get_or_init(|| Mutex::new(None))
}
fn model_load_admission() -> &'static ExclusiveGate {
static ADMISSION: ExclusiveGate = ExclusiveGate::new();
&ADMISSION
}
fn forward_admission() -> &'static ExclusiveGate {
static ADMISSION: ExclusiveGate = ExclusiveGate::new();
&ADMISSION
}
fn model_cache_guard() -> FocrResult<TrackedCacheGuard> {
let guard = model_cache()
.lock()
.map_err(|_| FocrError::Other(anyhow::anyhow!("model cache mutex poisoned")))?;
CACHE_GUARD_HELD_HERE.with(|c| c.set(c.get() + 1));
Ok(TrackedCacheGuard(guard))
}
struct TrackedCacheGuard(MutexGuard<'static, ModelCacheEntry>);
impl std::ops::Deref for TrackedCacheGuard {
type Target = ModelCacheEntry;
fn deref(&self) -> &ModelCacheEntry {
&self.0
}
}
impl std::ops::DerefMut for TrackedCacheGuard {
fn deref_mut(&mut self) -> &mut ModelCacheEntry {
&mut self.0
}
}
impl Drop for TrackedCacheGuard {
fn drop(&mut self) {
CACHE_GUARD_HELD_HERE.with(|c| c.set(c.get() - 1));
}
}
fn resolve_existing_model_artifact(path: &Path) -> Option<PathBuf> {
if path.is_dir() {
let shard = path.join(RAW_SAFETENSORS_SHARD_NAME);
shard.is_file().then_some(shard)
} else if path.exists() {
Some(path.to_path_buf())
} else {
None
}
}
fn is_short_model_spec(path: &Path) -> bool {
let mut components = path.components();
matches!(components.next(), Some(Component::Normal(_))) && components.next().is_none()
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum ModelQuantPreference {
Int8,
Int4,
}
impl ModelQuantPreference {
const ALL: [ModelQuantPreference; 2] = [Self::Int8, Self::Int4];
fn as_str(self) -> &'static str {
match self {
Self::Int8 => "int8",
Self::Int4 => "int4",
}
}
}
fn model_quant_preference_from_os(raw: Option<&OsStr>) -> Option<ModelQuantPreference> {
let value = raw?.to_str()?.trim().to_ascii_lowercase();
match value.as_str() {
"int8" | "q8" => Some(ModelQuantPreference::Int8),
"int4" | "q4" => Some(ModelQuantPreference::Int4),
_ => None,
}
}
fn model_quant_preference() -> Option<ModelQuantPreference> {
let raw = std::env::var_os(MODEL_QUANT_ENV);
model_quant_preference_from_os(raw.as_deref())
}
#[cfg(feature = "native")]
fn dist_cache_root() -> Option<PathBuf> {
crate::dist::cache_root()
}
#[cfg(not(feature = "native"))]
fn dist_cache_root() -> Option<PathBuf> {
None
}
fn model_search_dirs() -> Vec<PathBuf> {
let mut dirs = Vec::new();
if let Some(raw) = std::env::var_os(MODEL_DIR_ENV)
&& !raw.is_empty()
{
dirs.extend(std::env::split_paths(&raw));
}
if let Some(root) = dist_cache_root() {
let models = root.join("models");
let mut subdirs: Vec<PathBuf> = std::fs::read_dir(&models)
.into_iter()
.flatten()
.flatten()
.map(|e| e.path())
.filter(|p| p.is_dir())
.collect();
subdirs.sort();
dirs.push(models);
dirs.extend(subdirs);
}
dirs
}
#[must_use]
pub fn model_resolution_search_dirs() -> Vec<PathBuf> {
model_search_dirs()
}
#[derive(Debug, Clone)]
pub struct MusicPageMeta {
pub staves: Vec<(usize, tromr::StaffBBox)>,
pub skips: Vec<tromr::StaffSkip>,
pub warnings: Vec<tromr::MusicWarning>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct LayoutSpan {
pub label: String,
pub boxes: Vec<[i64; 4]>,
}
#[derive(Debug, Clone)]
pub struct RecognizedDocument {
pub markdown: String,
pub layout: Vec<LayoutSpan>,
}
#[derive(Debug, Clone)]
pub struct ExtractedFigure {
pub index: usize,
pub label: String,
pub bbox: [i64; 4],
pub markdown_ref: String,
pub image: image::DynamicImage,
}
fn push_unique_path(paths: &mut Vec<PathBuf>, path: PathBuf) {
if !paths.iter().any(|p| p == &path) {
paths.push(path);
}
}
fn versioned_quant_path(direct: &Path, quant: ModelQuantPreference) -> PathBuf {
direct.with_extension(format!(
"v{}.{}.focrq",
crate::UNLIMITED_OCR_ARTIFACT_VERSION,
quant.as_str()
))
}
fn short_model_candidates(
search_dir: &Path,
spec: &Path,
quant: Option<ModelQuantPreference>,
) -> Vec<PathBuf> {
let direct = search_dir.join(spec);
let mut candidates = Vec::new();
let is_focrq_or_bare = match spec.extension() {
None => true,
Some(ext) => ext.eq_ignore_ascii_case("focrq"),
};
if is_focrq_or_bare {
if let Some(quant) = quant {
push_unique_path(&mut candidates, versioned_quant_path(&direct, quant));
push_unique_path(
&mut candidates,
direct.with_extension(format!("{}.focrq", quant.as_str())),
);
}
push_unique_path(&mut candidates, direct.clone());
push_unique_path(&mut candidates, direct.with_extension("focrq"));
for quant in ModelQuantPreference::ALL {
push_unique_path(&mut candidates, versioned_quant_path(&direct, quant));
}
for quant in ModelQuantPreference::ALL {
push_unique_path(
&mut candidates,
direct.with_extension(format!("{}.focrq", quant.as_str())),
);
}
} else {
push_unique_path(&mut candidates, direct);
}
candidates
}
fn is_searchable_model_spec(path: &Path) -> bool {
path.is_relative() && !path.components().any(|c| matches!(c, Component::ParentDir))
}
fn model_search_specs(spec: &Path) -> Vec<PathBuf> {
let mut specs = vec![spec.to_path_buf()];
if !is_short_model_spec(spec)
&& let Some(file_name) = spec.file_name()
{
let basename = PathBuf::from(file_name);
if basename.as_path() != spec {
specs.push(basename);
}
}
specs
}
fn focrq_declares_wasm_recipe(path: &Path) -> bool {
const FOCRQ_PREAMBLE_LEN: usize = 6 + 4 + 1 + 32 + 8;
let Ok(mut file) = std::fs::File::open(path) else {
return false;
};
let mut preamble = [0u8; FOCRQ_PREAMBLE_LEN];
if file.read_exact(&mut preamble).is_err()
|| &preamble[..weights::FOCRQ_MAGIC.len()] != weights::FOCRQ_MAGIC
{
return false;
}
let header_len = u64::from_le_bytes(
preamble[43..51]
.try_into()
.expect("fixed .focrq header length slice"),
);
let Ok(header_len) = usize::try_from(header_len) else {
return false;
};
if header_len == 0 || header_len > MODEL_HEADER_VALIDATION_MAX_BYTES {
return false;
}
let mut header_bytes = vec![0u8; header_len];
if file.read_exact(&mut header_bytes).is_err() {
return false;
}
#[derive(serde::Deserialize)]
struct RecipeProbe {
#[serde(default)]
packing_manifest: Option<FocrqPackingManifestProbe>,
}
let Ok(probe) = serde_json::from_slice::<RecipeProbe>(&header_bytes) else {
return false;
};
probe
.packing_manifest
.and_then(|manifest| manifest.quant_recipe)
.as_deref()
== Some(crate::quant::convert::UNLIMITED_OCR_WASM_INT4_RECIPE_ID)
}
fn resolve_model_from_search_dirs_with_quant(
spec: &Path,
search_dirs: &[PathBuf],
quant: Option<ModelQuantPreference>,
) -> FocrResult<PathBuf> {
let specs = model_search_specs(spec);
for dir in search_dirs {
if let Some(resolved) = resolve_existing_model_artifact(dir)
&& !focrq_declares_wasm_recipe(&resolved)
{
return Ok(resolved);
}
for search_spec in &specs {
for candidate in short_model_candidates(dir, search_spec, quant) {
if let Some(resolved) = resolve_existing_model_artifact(&candidate)
&& !focrq_declares_wasm_recipe(&resolved)
{
return Ok(resolved);
}
}
}
}
let searched = if search_dirs.is_empty() {
"<none>".into()
} else {
search_dirs
.iter()
.map(|p| p.display().to_string())
.collect::<Vec<_>>()
.join(", ")
};
Err(FocrError::ModelNotFound(format!(
"no model artifact named {} (searched directories: {searched}; set {MODEL_DIR_ENV} \
or pass an explicit path)",
spec.display()
)))
}
#[cfg(test)]
fn resolve_model_from_search_dirs(spec: &Path, search_dirs: &[PathBuf]) -> FocrResult<PathBuf> {
resolve_model_from_search_dirs_with_quant(spec, search_dirs, None)
}
#[derive(serde::Deserialize)]
struct FocrqHeaderProbe {
tensors: std::collections::BTreeMap<String, weights::TensorRecord>,
#[serde(default)]
arch_target: u8,
#[serde(default)]
format_version: Option<u32>,
#[serde(default)]
source_sha256: String,
#[serde(default)]
license_notice: String,
#[serde(default)]
model_id: String,
#[serde(default)]
packing_manifest: Option<FocrqPackingManifestProbe>,
}
#[derive(serde::Deserialize)]
struct FocrqPackingManifestProbe {
#[serde(default)]
quant_recipe: Option<String>,
}
fn compatible_model_id(declared: &str) -> FocrResult<&'static str> {
if declared.is_empty() {
return Ok(model_arch::default_arch().id());
}
model_arch::arch_by_id(declared)
.map(|arch| arch.id())
.ok_or_else(|| FocrError::FormatMismatch(format!("unknown model_id {declared:?}")))
}
fn validate_header_license(notice: &str, model_id: &str) -> FocrResult<()> {
if model_id == model_arch::default_arch().id() {
if notice == crate::FOCR_MODEL_LICENSE_NOTICE
|| (notice.contains("Copyright (c) 2026 Baidu") && notice.contains("MIT License"))
{
return Ok(());
}
return Err(FocrError::FormatMismatch(
"Unlimited-OCR .focrq has an incompatible license notice".into(),
));
}
match model_arch::arch_by_id(model_id) {
Some(arch) if notice == arch.license_notice() => Ok(()),
_ => Err(FocrError::FormatMismatch(format!(
".focrq license notice does not match model {model_id:?}"
))),
}
}
fn validate_header_source_sha256(source_sha256: &str) -> FocrResult<()> {
if source_sha256.len() == 64
&& source_sha256
.bytes()
.all(|byte| matches!(byte, b'0'..=b'9' | b'a'..=b'f'))
{
return Ok(());
}
Err(FocrError::FormatMismatch(
".focrq source_sha256 is not 64 lowercase hex characters".into(),
))
}
fn lowercase_hex(bytes: &[u8]) -> String {
const HEX: &[u8; 16] = b"0123456789abcdef";
let mut encoded = String::with_capacity(bytes.len() * 2);
for &byte in bytes {
encoded.push(char::from(HEX[usize::from(byte >> 4)]));
encoded.push(char::from(HEX[usize::from(byte & 0x0f)]));
}
encoded
}
fn header_record_expected_bytes(
name: &str,
record: &weights::TensorRecord,
arch_target: u8,
) -> FocrResult<usize> {
let numel = record.shape.iter().copied().try_fold(1usize, |acc, dim| {
acc.checked_mul(dim).ok_or_else(|| {
FocrError::FormatMismatch(format!("tensor {name:?} shape element count overflows"))
})
})?;
if arch_target == 1 && record.dtype == DType::QInt8PerChan && record.shape.len() == 2 {
return record.shape[0]
.div_ceil(2)
.checked_mul(record.shape[1].div_ceil(8))
.and_then(|panels| panels.checked_mul(16))
.ok_or_else(|| {
FocrError::FormatMismatch(format!(
"tensor {name:?} SMMLA packed byte length overflows"
))
});
}
match record.dtype {
DType::F32 => numel.checked_mul(4),
DType::F16 | DType::BF16 => numel.checked_mul(2),
DType::QInt8PerChan => Some(numel),
DType::QInt4PerGroup if numel.is_multiple_of(2) => Some(numel / 2),
DType::QInt4PerGroup => {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?} has an odd int4 element count"
)));
}
}
.ok_or_else(|| FocrError::FormatMismatch(format!("tensor {name:?} byte length overflows")))
}
fn validate_header_quant_metadata(name: &str, record: &weights::TensorRecord) -> FocrResult<()> {
match record.dtype {
DType::F32 | DType::F16 | DType::BF16 => {
if record.scales_offset != 0
|| record.scales_len != 0
|| record.group_size != 0
|| record.tier != 0
{
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?}: non-quantized dtype has stray quantization metadata"
)));
}
}
DType::QInt8PerChan => {
let [rows, _cols] = record.shape.as_slice() else {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?}: QInt8 shape must be rank-2 [rows, cols]"
)));
};
if record.group_size != 0 || record.tier != 0 {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?}: QInt8 group_size and tier must be zero"
)));
}
let expected_scales = rows.checked_mul(4).ok_or_else(|| {
FocrError::FormatMismatch(format!(
"tensor {name:?}: QInt8 scale byte length overflows"
))
})?;
if record.scales_len != expected_scales {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?}: QInt8 scales_len {} != rows*f32 {expected_scales}",
record.scales_len
)));
}
}
DType::QInt4PerGroup => {
let [rows, cols] = record.shape.as_slice() else {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?}: QInt4 shape must be rank-2 [rows, cols]"
)));
};
if !crate::quant::int4::VALID_GROUP_SIZES.contains(&record.group_size)
|| !cols.is_multiple_of(record.group_size)
{
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?}: QInt4 group_size {} is invalid for {cols} columns",
record.group_size
)));
}
let expected_scales = rows
.checked_mul(cols / record.group_size)
.and_then(|groups| groups.checked_mul(4))
.ok_or_else(|| {
FocrError::FormatMismatch(format!(
"tensor {name:?}: QInt4 scale byte length overflows"
))
})?;
if record.scales_len != expected_scales {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?}: QInt4 scales_len {} != rows*groups*f32 {expected_scales}",
record.scales_len
)));
}
}
}
Ok(())
}
fn validate_header_directory(
directory: &std::collections::BTreeMap<String, weights::TensorRecord>,
payload_len: usize,
arch_target: u8,
) -> FocrResult<()> {
for (name, record) in directory {
validate_header_quant_metadata(name, record)?;
let end = record
.byte_offset
.checked_add(record.byte_len)
.ok_or_else(|| FocrError::FormatMismatch(format!("tensor {name:?} range overflows")))?;
if end > payload_len {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?} overruns the model payload"
)));
}
let expected = header_record_expected_bytes(name, record, arch_target)?;
if record.byte_len != expected {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?} byte_len {} != expected {expected}",
record.byte_len
)));
}
let scales_end = record
.scales_offset
.checked_add(record.scales_len)
.ok_or_else(|| {
FocrError::FormatMismatch(format!("tensor {name:?} scales range overflows"))
})?;
if scales_end > payload_len {
return Err(FocrError::FormatMismatch(format!(
"tensor {name:?} scales overrun the model payload"
)));
}
}
weights::validate_non_overlapping_ranges(directory)?;
Ok(())
}
fn validate_safetensors_header(header: &[u8], payload_len: usize) -> FocrResult<()> {
#[derive(serde::Deserialize)]
struct Entry {
dtype: String,
shape: Vec<usize>,
data_offsets: [usize; 2],
}
let raw: serde_json::Map<String, serde_json::Value> =
serde_json::from_slice(header).map_err(|error| {
FocrError::FormatMismatch(format!("safetensors header JSON invalid: {error}"))
})?;
let mut directory = std::collections::BTreeMap::new();
for (name, value) in raw {
if name == "__metadata__" {
continue;
}
let entry: Entry = serde_json::from_value(value).map_err(|error| {
FocrError::FormatMismatch(format!("safetensors entry {name:?} invalid: {error}"))
})?;
let [start, end] = entry.data_offsets;
if end < start {
return Err(FocrError::FormatMismatch(format!(
"safetensors entry {name:?} has reversed offsets"
)));
}
let dtype = match entry.dtype.as_str() {
"F32" => DType::F32,
"F16" => DType::F16,
"BF16" => DType::BF16,
other => {
return Err(FocrError::FormatMismatch(format!(
"unsupported safetensors dtype {other:?}"
)));
}
};
directory.insert(
name,
weights::TensorRecord {
dtype,
shape: entry.shape,
byte_offset: start,
byte_len: end - start,
scales_offset: 0,
scales_len: 0,
group_size: 0,
tier: 0,
},
);
}
if directory.is_empty() {
return Err(FocrError::FormatMismatch(
"safetensors tensor directory is empty; no native forward can use this artifact".into(),
));
}
validate_header_directory(&directory, payload_len, 0)?;
unlimited_ocr_census::validate_source_header(&directory)
}
fn validate_model_header_from_reader(mut reader: impl Read, file_len: u64) -> FocrResult<()> {
const FOCRQ_PREAMBLE_LEN: usize = 6 + 4 + 1 + 32 + 8;
let file_len = usize::try_from(file_len).map_err(|_| {
FocrError::FormatMismatch("model file exceeds this platform's addressable size".into())
})?;
let mut prefix = [0u8; 8];
reader
.read_exact(&mut prefix)
.map_err(|error| FocrError::FormatMismatch(format!("model header truncated: {error}")))?;
if &prefix[..weights::FOCRQ_MAGIC.len()] == weights::FOCRQ_MAGIC {
let mut preamble = [0u8; FOCRQ_PREAMBLE_LEN];
preamble[..prefix.len()].copy_from_slice(&prefix);
reader
.read_exact(&mut preamble[prefix.len()..])
.map_err(|error| {
FocrError::FormatMismatch(format!(".focrq preamble truncated: {error}"))
})?;
let version = u32::from_le_bytes(
preamble[6..10]
.try_into()
.expect("fixed .focrq version slice"),
);
if version != weights::FOCRQ_FORMAT_VERSION {
return Err(FocrError::FormatMismatch(format!(
".focrq format version {version} is unsupported; this binary requires exactly {}",
weights::FOCRQ_FORMAT_VERSION,
)));
}
let header_len_u64 = u64::from_le_bytes(
preamble[43..51]
.try_into()
.expect("fixed .focrq header length slice"),
);
let header_len = usize::try_from(header_len_u64)
.map_err(|_| FocrError::FormatMismatch(".focrq header length exceeds usize".into()))?;
if header_len == 0 || header_len > MODEL_HEADER_VALIDATION_MAX_BYTES {
return Err(FocrError::FormatMismatch(format!(
".focrq header length {header_len} is outside the bounded validation limit"
)));
}
let payload_base = FOCRQ_PREAMBLE_LEN
.checked_add(header_len)
.ok_or_else(|| FocrError::FormatMismatch(".focrq header length overflows".into()))?;
if payload_base > file_len {
return Err(FocrError::FormatMismatch(
".focrq header overruns the file".into(),
));
}
let mut header_bytes = vec![0u8; header_len];
reader.read_exact(&mut header_bytes).map_err(|error| {
FocrError::FormatMismatch(format!(".focrq header truncated: {error}"))
})?;
let header: FocrqHeaderProbe = serde_json::from_slice(&header_bytes).map_err(|error| {
FocrError::FormatMismatch(format!(".focrq header JSON invalid: {error}"))
})?;
match header.format_version {
Some(header_version) if header_version == version => {}
Some(header_version) => {
return Err(FocrError::FormatMismatch(format!(
".focrq header format_version {header_version} disagrees with preamble {version}"
)));
}
None => {
return Err(FocrError::FormatMismatch(
".focrq header is missing required format_version".into(),
));
}
}
let preamble_arch_target = preamble[10];
if header.arch_target != preamble_arch_target {
return Err(FocrError::FormatMismatch(format!(
".focrq header arch_target {} disagrees with preamble {preamble_arch_target}",
header.arch_target
)));
}
weights::validate_arch_target(header.arch_target)?;
validate_header_source_sha256(&header.source_sha256)?;
let preamble_source_sha256 = lowercase_hex(&preamble[11..43]);
if header.source_sha256 != preamble_source_sha256 {
return Err(FocrError::FormatMismatch(format!(
".focrq header source_sha256 {} disagrees with preamble {preamble_source_sha256}",
header.source_sha256
)));
}
if header.tensors.is_empty() {
return Err(FocrError::FormatMismatch(
".focrq tensor directory is empty; no native forward can use this artifact".into(),
));
}
let model_id = compatible_model_id(&header.model_id)?;
validate_header_license(&header.license_notice, model_id)?;
let declared_recipe = header
.packing_manifest
.as_ref()
.and_then(|manifest| manifest.quant_recipe.as_deref());
validate_header_directory(&header.tensors, file_len - payload_base, header.arch_target)?;
if model_id == model_arch::default_arch().id() {
unlimited_ocr_census::validate_focrq_header(
&header.tensors,
&header.source_sha256,
declared_recipe,
)?;
}
validate_unlimited_ocr_quant_records(
true,
model_id,
declared_recipe,
header
.tensors
.iter()
.map(|(name, record)| (name.as_str(), record.dtype)),
)?;
return Ok(());
}
let header_len_u64 = u64::from_le_bytes(prefix);
let header_len = usize::try_from(header_len_u64)
.map_err(|_| FocrError::FormatMismatch("safetensors header length exceeds usize".into()))?;
if header_len == 0 || header_len > MODEL_HEADER_VALIDATION_MAX_BYTES {
return Err(FocrError::FormatMismatch(format!(
"safetensors header length {header_len} is outside the bounded validation limit"
)));
}
let payload_base = 8usize
.checked_add(header_len)
.ok_or_else(|| FocrError::FormatMismatch("safetensors header length overflows".into()))?;
if payload_base > file_len {
return Err(FocrError::FormatMismatch(
"safetensors header overruns the file".into(),
));
}
let mut header = vec![0u8; header_len];
reader.read_exact(&mut header).map_err(|error| {
FocrError::FormatMismatch(format!("safetensors header truncated: {error}"))
})?;
validate_safetensors_header(&header, file_len - payload_base)
}
fn open_model_file(path: &Path) -> FocrResult<(std::fs::File, u64)> {
let file = std::fs::File::open(path).map_err(|error| {
FocrError::ModelNotFound(format!(
"cannot open model header at {}: {error}",
path.display()
))
})?;
let file_len = file.metadata().map_err(|error| {
FocrError::ModelNotFound(format!(
"cannot stat model artifact at {}: {error}",
path.display()
))
})?;
if !file_len.is_file() {
return Err(FocrError::FormatMismatch(format!(
"model artifact at {} is not a regular file",
path.display()
)));
}
Ok((file, file_len.len()))
}
fn validate_model_header_file(path: &Path) -> FocrResult<()> {
let (file, file_len) = open_model_file(path)?;
validate_model_header_from_reader(&file, file_len)
}
#[must_use]
pub fn native_model_available(path: &Path) -> bool {
let Ok(resolved) = OcrModel::resolve_model(path) else {
return false;
};
validate_model_header_file(&resolved).is_ok()
}
fn load_weights_from_resolved_model(resolved: &Path) -> FocrResult<Weights> {
load_weights_from_resolved_model_after_validation(resolved, || {})
}
fn load_weights_from_resolved_model_after_validation(
resolved: &Path,
after_validation: impl FnOnce(),
) -> FocrResult<Weights> {
let (file, file_len) = open_model_file(resolved)?;
validate_model_header_from_reader(&file, file_len)?;
after_validation();
Weights::load_opened(file, resolved).and_then(|weights| {
validate_unlimited_ocr_quant_recipe(&weights)?;
Ok(weights)
})
}
fn validate_unlimited_ocr_quant_recipe(weights: &Weights) -> FocrResult<()> {
validate_unlimited_ocr_quant_records(
weights.is_focrq(),
weights.model_id(),
weights.quant_recipe(),
weights.names().map(|name| {
let dtype = weights
.record(name)
.expect("name iterator and tensor directory must agree")
.dtype;
(name, dtype)
}),
)?;
unlimited_ocr_census::validate_loaded_weights(weights)
}
fn validate_unlimited_ocr_quant_records<'a>(
is_focrq: bool,
model_id: &str,
declared_recipe: Option<&str>,
records: impl IntoIterator<Item = (&'a str, DType)>,
) -> FocrResult<()> {
if !is_focrq || model_id != model_arch::default_arch().id() {
return Ok(());
}
let wasm_recipe =
declared_recipe == Some(crate::quant::convert::UNLIMITED_OCR_WASM_INT4_RECIPE_ID);
let recipe = Recipe::validated_default();
let mut violations = Vec::new();
for (name, dtype) in records {
let expected: (&str, bool) = if wasm_recipe {
match crate::quant::recipe::classify_wasm_experts_int4(name) {
crate::quant::recipe::WasmInt4Policy::ExpertInt4 => {
("QInt4PerGroup", dtype == DType::QInt4PerGroup)
}
crate::quant::recipe::WasmInt4Policy::Int8 => {
("QInt8PerChan", dtype == DType::QInt8PerChan)
}
crate::quant::recipe::WasmInt4Policy::KeepHighPrecision => {
("BF16 or F32", matches!(dtype, DType::BF16 | DType::F32))
}
}
} else if recipe.is_quantized(name) {
("QInt8PerChan", dtype == DType::QInt8PerChan)
} else {
("BF16 or F32", matches!(dtype, DType::BF16 | DType::F32))
};
let (expected, valid) = expected;
if !valid {
violations.push(format!("{name} is {dtype:?}, expected {expected}"));
}
}
if violations.is_empty() {
return Ok(());
}
let shown = violations
.iter()
.take(8)
.cloned()
.collect::<Vec<_>>()
.join("; ");
let omitted = violations.len().saturating_sub(8);
let suffix = if omitted == 0 {
String::new()
} else {
format!("; and {omitted} more")
};
let recipe_id = if wasm_recipe {
crate::quant::convert::UNLIMITED_OCR_WASM_INT4_RECIPE_ID
} else {
crate::quant::convert::UNLIMITED_OCR_INT8_RECIPE_ID
};
Err(FocrError::FormatMismatch(format!(
"Unlimited-OCR .focrq violates quant recipe {recipe_id}: {shown}{suffix}. Re-convert the \
original safetensors with this build; legacy full-int8 artifacts are not accepted",
)))
}
pub enum EmbedTable<'w> {
F32(Mat),
QInt8 {
data: &'w [u8],
scales: Vec<f32>,
rows: usize,
cols: usize,
},
}
impl EmbedTable<'_> {
pub fn rows(&self) -> usize {
match self {
EmbedTable::F32(m) => m.rows,
EmbedTable::QInt8 { rows, .. } => *rows,
}
}
pub fn cols(&self) -> usize {
match self {
EmbedTable::F32(m) => m.cols,
EmbedTable::QInt8 { cols, .. } => *cols,
}
}
#[must_use]
pub fn row_f32(&self, idx: usize) -> Vec<f32> {
match self {
EmbedTable::F32(m) => m.data[idx * m.cols..(idx + 1) * m.cols].to_vec(),
EmbedTable::QInt8 {
data, scales, cols, ..
} => {
let scale = scales[idx];
data[idx * cols..(idx + 1) * cols]
.iter()
.map(|&b| f32::from(b as i8) * scale)
.collect()
}
}
}
pub fn embed_ids(&self, ids: &[u32]) -> FocrResult<Mat> {
match self {
EmbedTable::F32(m) => decoder::embed_tokens(&m.data, m.rows, m.cols, ids),
EmbedTable::QInt8 { rows, cols, .. } => {
let mut out = Vec::with_capacity(ids.len() * cols);
for &id in ids {
let row = id as usize;
if row >= *rows {
return Err(FocrError::Other(anyhow::anyhow!(
"embed_tokens: id {row} out of range (vocab {rows})"
)));
}
out.extend_from_slice(&self.row_f32(row));
}
Ok(Mat::from_vec(ids.len(), *cols, out))
}
}
}
}
const EMBED_TOKENS: &str = "model.embed_tokens.weight";
fn embed_table_from(weights: &Weights) -> FocrResult<EmbedTable<'_>> {
const NAME: &str = EMBED_TOKENS;
if weights.arch_target() == 0
&& matches!(
weights.record(NAME).map(|rec| rec.dtype),
Some(DType::QInt8PerChan)
)
{
let view = weights.tensor(NAME)?;
let [rows, cols] = view.shape else {
return Err(FocrError::FormatMismatch(format!(
"tensor {NAME:?} has rank {}; expected 2 ([vocab, hidden])",
view.shape.len()
)));
};
let (rows, cols) = (*rows, *cols);
let expected = rows
.checked_mul(cols)
.ok_or_else(|| FocrError::FormatMismatch(format!("tensor {NAME:?} shape overflows")))?;
if view.data.len() != expected {
return Err(FocrError::FormatMismatch(format!(
"QInt8 tensor {NAME:?}: {} payload bytes != vocab*hidden {expected}",
view.data.len()
)));
}
if view.scales.len() != rows * 4 {
return Err(FocrError::FormatMismatch(format!(
"QInt8 tensor {NAME:?}: {} scale bytes != vocab*f32 {}",
view.scales.len(),
rows * 4
)));
}
let scales: Vec<f32> = view
.scales
.as_chunks::<4>()
.0
.iter()
.map(|c| f32::from_le_bytes(*c))
.collect();
return Ok(EmbedTable::QInt8 {
data: view.data,
scales,
rows,
cols,
});
}
Ok(EmbedTable::F32(weights.mat(NAME)?))
}
struct PagePrefill {
prefill_len: usize,
prompt_ids: Vec<u32>,
last_hidden: Mat,
caches: Vec<rswa::RingCache>,
image_w: u32,
image_h: u32,
}
impl OcrModel {
pub fn resolve_model(path: &Path) -> FocrResult<PathBuf> {
if let Some(resolved) = resolve_existing_model_artifact(path) {
return Ok(resolved);
}
if path.is_dir() {
return Err(FocrError::ModelNotFound(format!(
"no model artifact at {} (expected {RAW_SAFETENSORS_SHARD_NAME} inside \
safetensors directory; resolver lands in Phase 0/1, bd-223.7)",
path.display()
)));
}
if is_searchable_model_spec(path) {
resolve_model_from_search_dirs_with_quant(
path,
&model_search_dirs(),
model_quant_preference(),
)
} else {
Err(FocrError::ModelNotFound(format!(
"no model artifact at {} (resolver lands in Phase 0/1, bd-223.7)",
path.display()
)))
}
}
pub fn load(path: &Path) -> FocrResult<Arc<Self>> {
let resolved = Self::resolve_model(path)?;
let _load_permit = model_load_admission().acquire();
{
let guard = model_cache_guard()?;
if let Some((cached_path, weak)) = guard.as_ref()
&& *cached_path == resolved
&& let Some(strong) = weak.upgrade()
{
return Ok(strong);
}
}
let weights = load_weights_from_resolved_model(&resolved)?;
let model = Arc::new(Self {
path: resolved.clone(),
sidecars: SidecarBundle::default(),
weights,
decode_params: decode_params_from_env(),
decoder_cache_i8: FallibleOnce::new(),
decoder_cache: FallibleOnce::new(),
clip_cache: FallibleOnce::new(),
unlimited_vision: FallibleOnce::new(),
got_statics: FallibleOnce::new(),
onechart_statics: FallibleOnce::new(),
smol_statics: FallibleOnce::new(),
tokenizer: FallibleOnce::new(),
got_tokenizer: FallibleOnce::new(),
tromr_tokenizer: FallibleOnce::new(),
music_meta: std::sync::Mutex::new(None),
});
let mut guard = model_cache_guard()?;
*guard = Some((resolved, Arc::downgrade(&model)));
Ok(model)
}
pub fn from_weights(weights: Weights, sidecars: SidecarBundle) -> FocrResult<Arc<Self>> {
validate_unlimited_ocr_quant_recipe(&weights)?;
Ok(Arc::new(Self {
path: PathBuf::from("<in-memory>"),
sidecars,
weights,
decode_params: decode_params_from_env(),
decoder_cache_i8: FallibleOnce::new(),
decoder_cache: FallibleOnce::new(),
clip_cache: FallibleOnce::new(),
unlimited_vision: FallibleOnce::new(),
got_statics: FallibleOnce::new(),
onechart_statics: FallibleOnce::new(),
smol_statics: FallibleOnce::new(),
tokenizer: FallibleOnce::new(),
got_tokenizer: FallibleOnce::new(),
tromr_tokenizer: FallibleOnce::new(),
music_meta: std::sync::Mutex::new(None),
}))
}
#[must_use]
pub fn decode_params(&self) -> &DecodeParams {
&self.decode_params
}
#[must_use]
pub fn path(&self) -> &Path {
&self.path
}
pub fn forward(&self, image_path: &Path) -> FocrResult<(String, u32, u32)> {
let _admission = forward_admission().acquire();
if self.arch().id() == "got-ocr2" {
return self.forward_got(&preprocess::decode_path(image_path)?);
}
if self.arch().id() == "smolvlm2" {
return self.forward_smolvlm2(&preprocess::decode_path(image_path)?);
}
if self.arch().id() == "onechart" {
return self.forward_onechart(&preprocess::decode_path(image_path)?);
}
if self.arch().id() == "tromr" {
return self.forward_tromr(&preprocess::decode_path(image_path)?);
}
let t = Instant::now();
progress::emit("preprocess", 0, 0);
let pre = preprocess::preprocess_image(image_path, preprocess_mode())?;
progress::emit("preprocess", 1, 1);
timing_log(&format!("preprocess {:.2}s", t.elapsed().as_secs_f64()));
self.forward_pre(pre)
}
pub fn forward_dynamic(&self, img: image::DynamicImage) -> FocrResult<(String, u32, u32)> {
let _admission = forward_admission().acquire();
if self.arch().id() == "got-ocr2" {
return self.forward_got(&img);
}
if self.arch().id() == "smolvlm2" {
return self.forward_smolvlm2(&img);
}
if self.arch().id() == "onechart" {
return self.forward_onechart(&img);
}
if self.arch().id() == "tromr" {
return self.forward_tromr(&img);
}
let t = Instant::now();
progress::emit("preprocess", 0, 0);
let pre = preprocess::preprocess_dynamic(img, preprocess_mode())?;
progress::emit("preprocess", 1, 1);
timing_log(&format!("preprocess {:.2}s", t.elapsed().as_secs_f64()));
self.forward_pre(pre)
}
fn forward_got(&self, img: &image::DynamicImage) -> FocrResult<(String, u32, u32)> {
use image::GenericImageView;
let t = Instant::now();
let (w, h) = img.dimensions();
let tk = self.got_tokenizer()?;
let format = got_format_requested();
let max_new = self.decode_params.max_length.min(got::MAX_NEW_TOKENS);
let text = got::recognize(&self.weights, self.got_statics()?, tk, img, max_new, format)?;
timing_log(&format!("got forward {:.2}s", t.elapsed().as_secs_f64()));
Ok((text, w, h))
}
fn forward_smolvlm2(&self, img: &image::DynamicImage) -> FocrResult<(String, u32, u32)> {
use image::GenericImageView;
let t = Instant::now();
let (w, h) = img.dimensions();
let tk = self.tokenizer()?;
let question = smolvlm2_question();
let max_new = self.decode_params.max_length;
let text = smolvlm2::recognize(
&self.weights,
self.smol_statics()?,
tk,
img,
&question,
max_new,
)?;
timing_log(&format!(
"smolvlm2 forward {:.2}s",
t.elapsed().as_secs_f64()
));
Ok((text, w, h))
}
fn forward_onechart(&self, img: &image::DynamicImage) -> FocrResult<(String, u32, u32)> {
use image::GenericImageView;
let t = Instant::now();
let (w, h) = img.dimensions();
let tk = self.tokenizer()?;
let max_new = self.decode_params.max_length;
let res = onechart::recognize(&self.weights, self.onechart_statics()?, tk, img, max_new)?;
timing_log(&format!(
"onechart forward {:.2}s (reliable_distance {:?} reliable {:?})",
t.elapsed().as_secs_f64(),
res.reliable_distance,
res.reliable
));
Ok((res.json_text, w, h))
}
pub fn take_music_meta(&self) -> Option<MusicPageMeta> {
self.music_meta.lock().ok().and_then(|mut slot| slot.take())
}
fn forward_tromr(&self, img: &image::DynamicImage) -> FocrResult<(String, u32, u32)> {
use image::GenericImageView;
let t = Instant::now();
let (w, h) = img.dimensions();
let tk = self.tromr_tokenizer()?;
let page = tromr::recognize_page(&self.weights, tk, img)?;
let total = page.staves.len() + page.skips.len();
for skip in &page.skips {
crate::progress::stderr_message(format_args!(
"[focr] staff {}/{} skipped: {} (bbox x{} y{} w{} h{})",
skip.index + 1,
total,
skip.reason,
skip.bbox.0,
skip.bbox.1,
skip.bbox.2,
skip.bbox.3
));
}
timing_log(&format!(
"tromr forward {:.2}s ({}/{} staves recognized, semantic {} chars total)",
t.elapsed().as_secs_f64(),
page.staves.len(),
total,
page.staves
.iter()
.map(|(_, r, _)| r.semantic.len())
.sum::<usize>()
));
let semantics: Vec<String> = page
.staves
.iter()
.map(|(_, r, _)| r.semantic.clone())
.collect();
let warnings = tromr::sanity_warnings(&semantics);
for w in &warnings {
timing_log(&format!(
" tromr.sanity {} part {} measure {}: {}",
w.kind, w.part, w.measure, w.detail
));
}
if let Ok(mut slot) = self.music_meta.lock() {
*slot = Some(MusicPageMeta {
staves: page.staves.iter().map(|(i, _, b)| (*i, *b)).collect(),
skips: page.skips.clone(),
warnings,
});
}
let xml = if page.staves.len() == 1 {
page.staves
.into_iter()
.next()
.expect("one staff")
.1
.musicxml
} else {
let semantics: Vec<String> = page
.staves
.into_iter()
.map(|(_, r, _)| r.semantic)
.collect();
tromr::staves_to_musicxml(&semantics)?
};
Ok((xml, w, h))
}
fn tromr_tokenizer(&self) -> FocrResult<&crate::tokenizer::music::MusicTokenizer> {
self.tromr_tokenizer.get_or_try_init(|| {
if let Some([rhythm, pitch, lift, note]) = &self.sidecars.music_tables {
return crate::tokenizer::music::MusicTokenizer::from_json_tables(
[rhythm, pitch, lift, note],
[
"tokenizer_rhythm.json (in-memory)",
"tokenizer_pitch.json (in-memory)",
"tokenizer_lift.json (in-memory)",
"tokenizer_note.json (in-memory)",
],
);
}
let dir = self.path.parent().unwrap_or_else(|| Path::new("."));
crate::tokenizer::music::MusicTokenizer::from_dir(dir)
})
}
fn got_tokenizer(&self) -> FocrResult<&crate::tokenizer::tiktoken::Tiktoken> {
self.got_tokenizer.get_or_try_init(|| {
if let Some(bytes) = &self.sidecars.qwen_tiktoken {
return crate::tokenizer::tiktoken::Tiktoken::from_qwen_tiktoken(bytes);
}
let dir = self.path.parent().unwrap_or_else(|| Path::new("."));
let path = dir.join("qwen.tiktoken");
let bytes = std::fs::read(&path).map_err(|e| {
FocrError::ModelNotFound(format!(
"GOT-OCR2 tokenizer qwen.tiktoken not found beside the model at {}: {e}",
path.display()
))
})?;
crate::tokenizer::tiktoken::Tiktoken::from_qwen_tiktoken(&bytes)
})
}
#[must_use]
pub fn arch(&self) -> &'static dyn model_arch::ModelArch {
model_arch::arch_by_id(self.weights.model_id()).unwrap_or_else(model_arch::default_arch)
}
fn ensure_arch_implemented(arch: &dyn model_arch::ModelArch) -> FocrResult<()> {
if arch.implemented() {
Ok(())
} else {
Err(FocrError::NotImplemented(format!(
"model architecture '{}' ({}) forward is not yet implemented \
(franken_ocr model zoo, epic bd-3jo6)",
arch.id(),
arch.display_name()
)))
}
}
fn forward_pre(&self, pre: Preprocessed) -> FocrResult<(String, u32, u32)> {
Self::ensure_arch_implemented(self.arch())?;
let (image_w, image_h) = Self::image_dims(&pre);
let tv = Instant::now();
let vision_features = self.vision_tower(&pre)?;
timing_log(&format!("vision_tower {:.2}s", tv.elapsed().as_secs_f64()));
let (inputs_embeds, prompt_ids) = self.build_inputs_embeds(&pre, &vision_features)?;
let generated = self.generate(inputs_embeds, &prompt_ids)?;
progress::emit("postprocess", 0, 0);
let decoded = self.tokenizer()?.decode(&generated)?;
Ok((decoded, image_w, image_h))
}
pub fn recognize(&self, image_path: &Path) -> FocrResult<String> {
let (decoded, image_w, image_h) = self.forward(image_path)?;
postprocess::finalize(&decoded, image_w, image_h)
}
pub fn recognize_dynamic(&self, img: image::DynamicImage) -> FocrResult<String> {
let (decoded, image_w, image_h) = self.forward_dynamic(img)?;
postprocess::finalize(&decoded, image_w, image_h)
}
pub fn recognize_with_layout(&self, image_path: &Path) -> FocrResult<RecognizedDocument> {
let (decoded, image_w, image_h) = self.forward(image_path)?;
Self::finalize_document(&decoded, image_w, image_h)
}
pub fn recognize_dynamic_with_layout(
&self,
img: image::DynamicImage,
) -> FocrResult<RecognizedDocument> {
let (decoded, image_w, image_h) = self.forward_dynamic(img)?;
Self::finalize_document(&decoded, image_w, image_h)
}
fn finalize_document(
decoded: &str,
image_w: u32,
image_h: u32,
) -> FocrResult<RecognizedDocument> {
let markdown = postprocess::finalize(decoded, image_w, image_h)?;
let layout = postprocess::parse_layout(decoded, image_w, image_h)
.into_iter()
.map(|(label, boxes)| LayoutSpan { label, boxes })
.collect();
Ok(RecognizedDocument { markdown, layout })
}
pub fn recognize_with_figures(
&self,
image_path: &Path,
) -> FocrResult<(RecognizedDocument, Vec<ExtractedFigure>)> {
let (decoded, image_w, image_h) = self.forward(image_path)?;
let document = Self::finalize_document(&decoded, image_w, image_h)?;
let figures = if postprocess::figure_refs(&decoded, image_w, image_h, "").is_empty() {
Vec::new()
} else {
let source = preprocess::decode_path(image_path)?;
Self::crop_figures(&decoded, &source, image_w, image_h, "")
};
Ok((document, figures))
}
pub fn recognize_dynamic_with_figures(
&self,
img: image::DynamicImage,
) -> FocrResult<(RecognizedDocument, Vec<ExtractedFigure>)> {
let source = img.clone();
let (decoded, image_w, image_h) = self.forward_dynamic(img)?;
let document = Self::finalize_document(&decoded, image_w, image_h)?;
let figures = Self::crop_figures(&decoded, &source, image_w, image_h, "");
Ok((document, figures))
}
fn crop_figures(
decoded: &str,
source: &image::DynamicImage,
image_w: u32,
image_h: u32,
img_base: &str,
) -> Vec<ExtractedFigure> {
let iw = i64::from(source.width());
let ih = i64::from(source.height());
postprocess::figure_refs(decoded, image_w, image_h, img_base)
.into_iter()
.filter_map(|fr| {
let &[x1, y1, x2, y2] = fr.boxes.first()?;
let cx1 = x1.min(x2).clamp(0, iw);
let cy1 = y1.min(y2).clamp(0, ih);
let cx2 = x1.max(x2).clamp(0, iw);
let cy2 = y1.max(y2).clamp(0, ih);
let w = u32::try_from(cx2 - cx1).unwrap_or(0);
let h = u32::try_from(cy2 - cy1).unwrap_or(0);
if w == 0 || h == 0 {
return None; }
#[allow(clippy::cast_sign_loss)] let image = source.crop_imm(cx1 as u32, cy1 as u32, w, h);
Some(ExtractedFigure {
index: fr.index,
label: fr.label,
bbox: [cx1, cy1, cx2, cy2],
markdown_ref: fr.markdown_ref,
image,
})
})
.collect()
}
pub fn recognize_multi_page(&self, image_paths: &[&Path]) -> FocrResult<String> {
if image_paths.is_empty() {
return Err(FocrError::Other(anyhow::anyhow!(
"recognize_multi_page: image_paths must be non-empty"
)));
}
if self.arch().id() != model_arch::default_arch().id() {
return Err(FocrError::NotImplemented(format!(
"multi-page cross-page parsing (infer_multi) is the Unlimited-OCR \
contract; model {:?} parses pages independently (use the plain \
multi-input/batch path)",
self.arch().id()
)));
}
Self::ensure_arch_implemented(self.arch())?;
let mut pres = Vec::with_capacity(image_paths.len());
for path in image_paths {
let img = image::open(path)
.map_err(|e| FocrError::InputDecode(format!("{}: {e}", path.display())))?;
pres.push(preprocess::preprocess_dynamic_squash(
img,
MULTI_PAGE_BASE_SIZE,
)?);
}
self.recognize_multi_page_pres(pres, None)
}
pub fn recognize_multi_page_dynamic(
&self,
images: Vec<image::DynamicImage>,
) -> FocrResult<String> {
if images.is_empty() {
return Err(FocrError::Other(anyhow::anyhow!(
"recognize_multi_page_dynamic: images must be non-empty"
)));
}
if self.arch().id() != model_arch::default_arch().id() {
return Err(FocrError::NotImplemented(format!(
"multi-page cross-page parsing (infer_multi) is the Unlimited-OCR \
contract; model {:?} parses pages independently (use the plain \
multi-input/batch path)",
self.arch().id()
)));
}
Self::ensure_arch_implemented(self.arch())?;
let mut pres = Vec::with_capacity(images.len());
for img in images {
pres.push(preprocess::preprocess_dynamic_squash(
img,
MULTI_PAGE_BASE_SIZE,
)?);
}
self.recognize_multi_page_pres(pres, None)
}
pub fn recognize_multi_page_dynamic_streaming(
&self,
images: Vec<image::DynamicImage>,
on_page: &mut dyn FnMut(usize, &str),
) -> FocrResult<String> {
if images.is_empty() {
return Err(FocrError::Other(anyhow::anyhow!(
"recognize_multi_page_dynamic_streaming: images must be non-empty"
)));
}
if self.arch().id() != model_arch::default_arch().id() {
return Err(FocrError::NotImplemented(format!(
"multi-page cross-page parsing (infer_multi) is the Unlimited-OCR \
contract; model {:?} parses pages independently (use the plain \
multi-input/batch path)",
self.arch().id()
)));
}
Self::ensure_arch_implemented(self.arch())?;
let mut pres = Vec::with_capacity(images.len());
for img in images {
pres.push(preprocess::preprocess_dynamic_squash(
img,
MULTI_PAGE_BASE_SIZE,
)?);
}
self.recognize_multi_page_pres(pres, Some(on_page))
}
fn recognize_multi_page_pres(
&self,
pres: Vec<Preprocessed>,
on_page: Option<PageSink<'_>>,
) -> FocrResult<String> {
let _admission = forward_admission().acquire();
let tv = Instant::now();
let mut globals: Vec<Mat> = Vec::with_capacity(pres.len());
for pre in &pres {
let mut feats = self.vision_tower(pre)?;
if feats.len() != 1 {
return Err(FocrError::Other(anyhow::anyhow!(
"recognize_multi_page: Base-mode page produced {} vision blocks; expected 1",
feats.len()
)));
}
globals.push(feats.remove(0));
}
timing_log(&format!(
"multi_page.vision ({} pages) {:.2}s",
pres.len(),
tv.elapsed().as_secs_f64()
));
let (prompt_ids, images_seq_mask) = self.build_prompt_multi(&pres)?;
if prompt_ids.len() + 1 > MAX_POSITION_EMBEDDINGS {
return Err(FocrError::Other(anyhow::anyhow!(
"recognize_multi_page: the assembled {}-page prefix is {} tokens — over the \
{MAX_POSITION_EMBEDDINGS} position budget by {} with no room to generate; \
split the document into smaller multi-page passes",
pres.len(),
prompt_ids.len(),
(prompt_ids.len() + 1).saturating_sub(MAX_POSITION_EMBEDDINGS),
)));
}
let mut inputs_embeds = self.embed_prompt(&prompt_ids)?;
let image_newline = Self::image_newline(&self.weights)?;
let view_seperator = Self::view_seperator(&self.weights)?;
let grid = preprocess::num_queries(MULTI_PAGE_BASE_SIZE);
connector::fuse_no_crop(
&self.weights,
&mut inputs_embeds,
&globals,
grid,
grid,
&image_newline,
&view_seperator,
&images_seq_mask,
)?;
let params = sampler::DecodeParams {
ngram_window: sampler::NGRAM_WINDOW_MULTI,
max_length: self
.decode_params
.max_length
.min(MAX_POSITION_EMBEDDINGS - prompt_ids.len()),
..self.decode_params.clone()
};
let td = Instant::now();
let generated = match on_page {
None => self.generate_with(inputs_embeds, &prompt_ids, ¶ms)?,
Some(on_page) => {
const STREAM_CHECK_EVERY: usize = 8;
let tok = self.tokenizer()?;
let mut stream = postprocess::PageStream::new();
let mut ids: Vec<u32> = Vec::new();
let mut stream_broken = false;
let mut observer = |id: u32| {
ids.push(id);
if stream_broken || !ids.len().is_multiple_of(STREAM_CHECK_EVERY) {
return;
}
match tok.decode(&ids) {
Ok(text) => stream.feed(&text, &mut *on_page),
Err(_) => stream_broken = true,
}
};
let generated = self.generate_with_observer(
inputs_embeds,
&prompt_ids,
¶ms,
Some(&mut observer),
)?;
if !stream_broken && let Ok(text) = tok.decode(&generated) {
let text = postprocess::strip_eos(&text);
stream.feed(&text, &mut *on_page);
stream.finish(&text, &mut *on_page);
}
generated
}
};
timing_log(&format!(
"multi_page.decode {} tokens {:.2}s",
generated.len(),
td.elapsed().as_secs_f64()
));
let decoded = self.tokenizer()?.decode(&generated)?;
postprocess::finalize_multi(&decoded, pres.len())
}
#[must_use]
pub fn recognize_batch(&self, image_paths: &[&Path]) -> Vec<FocrResult<String>> {
let spine = batch_scheduler::spine_enabled()
&& (int8_decode_requested() || self.artifact_stores_gated_int8())
&& std::env::var_os(DECODE_STATELESS_ENV).is_none()
&& self.arch().id() == model_arch::default_arch().id();
if !spine
&& batch_scheduler::spine_enabled()
&& matches!(self.arch().id(), "got-ocr2" | "smolvlm2" | "onechart")
{
let _admission = forward_admission().acquire();
match self.recognize_batch_dense(image_paths) {
Ok(results) => return results,
Err(err) => {
let msg = err.to_string();
return image_paths
.iter()
.map(|_| Err(FocrError::Other(anyhow::anyhow!("{msg}"))))
.collect();
}
}
}
if !spine {
return image_paths.iter().map(|p| self.recognize(p)).collect();
}
let _admission = forward_admission().acquire();
match self.recognize_batch_spine(image_paths) {
Ok(results) => results,
Err(err) => {
let msg = err.to_string();
image_paths
.iter()
.map(|_| Err(FocrError::Other(anyhow::anyhow!("{msg}"))))
.collect()
}
}
}
fn recognize_batch_dense(&self, image_paths: &[&Path]) -> FocrResult<Vec<FocrResult<String>>> {
let arch_id = self.arch().id();
let max_new = match arch_id {
"got-ocr2" => self.decode_params.max_length.min(got::MAX_NEW_TOKENS),
_ => self.decode_params.max_length,
};
let t = Instant::now();
let mut out: Vec<Option<FocrResult<String>>> =
(0..image_paths.len()).map(|_| None).collect();
let mut live: Vec<usize> = Vec::new();
let mut imgs: Vec<image::DynamicImage> = Vec::new();
for (i, p) in image_paths.iter().enumerate() {
crate::cancel_checkpoint()?;
match preprocess::decode_path(p) {
Ok(img) => {
live.push(i);
imgs.push(img);
}
Err(e) => out[i] = Some(Err(e)),
}
}
if !imgs.is_empty() {
use image::GenericImageView;
let dims: Vec<(u32, u32)> = imgs.iter().map(|im| im.dimensions()).collect();
let refs: Vec<&image::DynamicImage> = imgs.iter().collect();
let texts: Vec<String> = match arch_id {
"got-ocr2" => got::recognize_batch(
&self.weights,
self.got_statics()?,
self.got_tokenizer()?,
&refs,
max_new,
got_format_requested(),
)?,
"smolvlm2" => smolvlm2::recognize_batch(
&self.weights,
self.smol_statics()?,
self.tokenizer()?,
&refs,
&smolvlm2_question(),
max_new,
)?,
"onechart" => onechart::recognize_batch(
&self.weights,
self.onechart_statics()?,
self.tokenizer()?,
&refs,
max_new,
)?
.into_iter()
.map(|r| r.json_text)
.collect(),
other => {
return Err(FocrError::Other(anyhow::anyhow!(
"dense batch spine: unrouted arch {other}"
)));
}
};
for ((slot, text), (iw, ih)) in live.into_iter().zip(texts).zip(dims) {
out[slot] = Some(postprocess::finalize(&text, iw, ih));
}
}
timing_log(&format!(
"{arch_id} forward(batch of {}) {:.2}s",
image_paths.len(),
t.elapsed().as_secs_f64()
));
Ok(out
.into_iter()
.map(|r| r.unwrap_or_else(|| Err(FocrError::Other(anyhow::anyhow!("page skipped")))))
.collect())
}
fn recognize_batch_spine(&self, image_paths: &[&Path]) -> FocrResult<Vec<FocrResult<String>>> {
let n = image_paths.len();
let wc = self.decoder_cache_i8()?;
let embed_table = self.embed_table()?;
let tokenizer = self.tokenizer()?;
let mut out: Vec<Option<FocrResult<String>>> = (0..n).map(|_| None).collect();
let mode = preprocess_mode();
let mut pres: Vec<Option<Preprocessed>> = (0..n).map(|_| None).collect();
for (gi, &path) in image_paths.iter().enumerate() {
match preprocess::preprocess_image(path, mode) {
Ok(p) => pres[gi] = Some(p),
Err(e) => out[gi] = Some(Err(e)),
}
}
let admitted: Vec<usize> = (0..n).filter(|&gi| pres[gi].is_some()).collect();
let mut feats: Vec<Option<Vec<Mat>>> = (0..n).map(|_| None).collect();
if !admitted.is_empty() {
let mut batched_ok = false;
if batch_vision_enabled() {
let prefs: Vec<&Preprocessed> = admitted
.iter()
.map(|&gi| pres[gi].as_ref().expect("admitted page has a preprocess"))
.collect();
match self.vision_tower_batched_pages(&prefs) {
Ok(all) => {
for (k, f) in all.into_iter().enumerate() {
feats[admitted[k]] = Some(f);
}
batched_ok = true;
}
Err(e) => {
crate::progress::stderr_message(format_args!(
"[focr] batched vision failed ({e}); using the per-page tower"
));
}
}
}
if !batched_ok {
for &gi in &admitted {
let pre = pres[gi].as_ref().expect("admitted page has a preprocess");
match self.vision_tower(pre) {
Ok(f) => feats[gi] = Some(f),
Err(e) => out[gi] = Some(Err(e)),
}
}
}
}
let mut streams: Vec<batch_scheduler::PageStream> = Vec::new();
let mut stream_caches: Vec<Vec<rswa::RingCache>> = Vec::new();
let mut scheduled: Vec<(usize, u32, u32)> = Vec::new();
for gi in 0..n {
let (Some(pre), Some(vf)) = (pres[gi].take(), feats[gi].take()) else {
continue;
};
match self.prefill_from_features(wc, &pre, &vf) {
Ok(p) => {
streams.push(batch_scheduler::PageStream::new(
gi,
p.prefill_len,
&p.prompt_ids,
p.last_hidden,
));
stream_caches.push(p.caches);
scheduled.push((gi, p.image_w, p.image_h));
}
Err(e) => out[gi] = Some(Err(e)),
}
}
if !streams.is_empty() {
let mut batched = rswa::BatchedRingCache::from_streams(stream_caches)?;
let mut step = batch_scheduler::DecoderBatchStep {
wc,
caches: &mut batched,
embed_table: &embed_table,
params: &self.decode_params,
};
let mut scheduler =
batch_scheduler::BatchScheduler::from_env(self.decode_params.max_length);
let token_lists = scheduler.run(streams, &mut step)?;
for (k, tokens) in token_lists.into_iter().enumerate() {
let (gi, w, h) = scheduled[k];
let finalized = tokenizer
.decode(&tokens)
.and_then(|decoded| postprocess::finalize(&decoded, w, h));
out[gi] = Some(finalized);
}
}
Ok(out
.into_iter()
.map(|slot| {
slot.unwrap_or_else(|| {
Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::recognize_batch: page produced no result"
)))
})
})
.collect())
}
fn decoder_cache_i8(&self) -> FocrResult<&decoder::DecoderWeightCacheI8> {
self.require_full_int8_authorization()?;
self.decoder_cache_i8
.get_or_try_init(|| decoder::DecoderWeightCacheI8::build(&self.weights))
}
fn artifact_stores_gated_int8(&self) -> bool {
self.weights.quant_recipe()
== Some(crate::quant::convert::UNLIMITED_OCR_WASM_INT4_RECIPE_ID)
}
fn require_full_int8_authorization(&self) -> FocrResult<()> {
if self.artifact_stores_gated_int8() {
return Ok(());
}
require_experimental_full_int8_recipe()
}
fn decoder_cache(&self) -> FocrResult<&decoder::DecoderWeightCache> {
self.decoder_cache
.get_or_try_init(|| decoder::DecoderWeightCache::build(&self.weights))
}
fn clip_weights(&self) -> FocrResult<&vision_clip::ClipWeights> {
self.clip_cache.get_or_try_init(|| {
let th = Instant::now();
let built = vision_clip::clip_weights_from(&self.weights)?;
timing_log(&format!(
" clip.hydrate(cached) {:.2}s",
th.elapsed().as_secs_f64()
));
Ok(built)
})
}
fn unlimited_vision_statics(&self) -> FocrResult<&UnlimitedVisionStatics> {
self.unlimited_vision.get_or_try_init(|| {
let started = Instant::now();
let built = UnlimitedVisionStatics {
sam: vision_sam::sam_weights_from(&self.weights, "model.sam_model")?,
projector: vision_bridge::projector_weights_from(&self.weights)?,
};
timing_log(&format!(
" unlimited_vision.hydrate(cached) {:.2}s",
started.elapsed().as_secs_f64()
));
Ok(built)
})
}
fn got_statics(&self) -> FocrResult<&got::GotStatics> {
self.got_statics.get_or_try_init(|| {
got::hydrate_statics(
&self.weights,
self.arch().vision_tower_prefix(),
stream_vision_tower(),
)
})
}
fn onechart_statics(&self) -> FocrResult<&onechart::OnechartStatics> {
self.onechart_statics.get_or_try_init(|| {
onechart::hydrate_statics(
&self.weights,
self.arch().vision_tower_prefix(),
stream_vision_tower(),
)
})
}
fn smol_statics(&self) -> FocrResult<&smolvlm2::SmolStatics> {
self.smol_statics
.get_or_try_init(|| smolvlm2::hydrate_statics(&self.weights, stream_vision_tower()))
}
fn prefill_from_features(
&self,
wc: &decoder::DecoderWeightCacheI8,
pre: &Preprocessed,
vision_features: &[Mat],
) -> FocrResult<PagePrefill> {
let (image_w, image_h) = Self::image_dims(pre);
let (inputs_embeds, prompt_ids) = self.build_inputs_embeds(pre, vision_features)?;
let prefill_len = inputs_embeds.rows;
let _fwd = enter_forward();
let (hidden, caches) = decoder::prefill_with_cache_i8(wc, &inputs_embeds)?;
let last_hidden = Self::last_hidden_row(&hidden)?;
Ok(PagePrefill {
prefill_len,
prompt_ids,
last_hidden,
caches,
image_w,
image_h,
})
}
fn tokenizer(&self) -> FocrResult<&crate::tokenizer::Tokenizer> {
self.tokenizer.get_or_try_init(|| {
if let Some(bytes) = &self.sidecars.tokenizer_json {
return crate::tokenizer::Tokenizer::from_json_bytes(bytes);
}
if let Some([vocab, merges, added]) = &self.sidecars.opt_triple {
return crate::tokenizer::Tokenizer::from_opt_files(vocab, merges, added);
}
let dir = self.path.parent().unwrap_or_else(|| Path::new("."));
if self.arch().id() == "onechart" {
crate::tokenizer::Tokenizer::from_opt_dir(dir)
} else {
crate::tokenizer::Tokenizer::load(&dir.join("tokenizer.json"))
}
})
}
fn vision_tower(&self, pre: &Preprocessed) -> FocrResult<Vec<Mat>> {
let _fwd = enter_forward();
let wasm_residency = self.artifact_stores_gated_int8();
let statics = if !wasm_residency && unlimited_vision_cache_enabled() {
Some(self.unlimited_vision_statics()?)
} else {
None
};
let mut features = Vec::new();
let views = Self::views(pre);
progress::vision_begin(
views.len() as u64
* (vision_sam::DEPTH as u64 + vision_clip::ClipConfig::default().num_layers as u64),
);
if wasm_residency {
return self.vision_tower_streamed_views(&views);
}
for view in views {
let ts = Instant::now();
let sam = if let Some(statics) = statics {
let side = (view.cols as f64).sqrt() as usize;
if side * side != view.cols {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::vision_tower: image.cols {} is not a perfect square",
view.cols
)));
}
vision_sam::forward_with(&statics.sam, &view, side, side)?
} else if wasm_residency {
vision_sam::forward_streamed(&self.weights, &view, "model.sam_model")?
} else {
vision_sam::forward(&self.weights, &view)?
};
timing_log(&format!(" vision.sam {:.2}s", ts.elapsed().as_secs_f64()));
let tc = Instant::now();
let clip = if wasm_residency {
vision_clip::forward_from_sam_streamed(
&vision_clip::ClipConfig::default(),
&self.weights,
&sam,
)?
} else {
vision_clip::forward_from_sam(
&vision_clip::ClipConfig::default(),
self.clip_weights()?,
&sam,
)?
};
timing_log(&format!(" vision.clip {:.2}s", tc.elapsed().as_secs_f64()));
let tb = Instant::now();
let projected = if let Some(statics) = statics {
vision_bridge::forward_with(&statics.projector, &clip, &sam)?
} else {
vision_bridge::forward(&self.weights, &clip, &sam)?
};
timing_log(&format!(
" vision.bridge {:.2}s",
tb.elapsed().as_secs_f64()
));
features.push(projected);
}
Ok(features)
}
fn vision_tower_streamed_views(&self, views: &[Mat]) -> FocrResult<Vec<Mat>> {
let clip_cfg = vision_clip::ClipConfig::default();
let mut features = Vec::with_capacity(views.len());
for chunk in views.chunks(STREAMED_VIEW_CHUNK) {
let view_refs: Vec<&Mat> = chunk.iter().collect();
let ts = Instant::now();
let sams =
vision_sam::forward_streamed_views(&self.weights, &view_refs, "model.sam_model")?;
timing_log(&format!(
" vision.sam(streamed x{}) {:.2}s",
view_refs.len(),
ts.elapsed().as_secs_f64()
));
let tc = Instant::now();
let sam_refs: Vec<&Mat> = sams.iter().collect();
let clips =
vision_clip::forward_from_sam_streamed_views(&clip_cfg, &self.weights, &sam_refs)?;
timing_log(&format!(
" vision.clip(streamed x{}) {:.2}s",
sam_refs.len(),
tc.elapsed().as_secs_f64()
));
let tb = Instant::now();
let projector = vision_bridge::projector_weights_from(&self.weights)?;
for (clip, sam) in clips.iter().zip(sams.iter()) {
features.push(vision_bridge::forward_with(&projector, clip, sam)?);
}
timing_log(&format!(
" vision.bridge(streamed x{}) {:.2}s",
clips.len(),
tb.elapsed().as_secs_f64()
));
}
Ok(features)
}
fn vision_tower_batched_pages(&self, pres: &[&Preprocessed]) -> FocrResult<Vec<Vec<Mat>>> {
let _fwd = enter_forward();
let mut page_views: Vec<Vec<Mat>> = pres.iter().map(|pre| Self::views(pre)).collect();
let th = Instant::now();
let retain_statics = unlimited_vision_cache_enabled();
let statics_owner = retained_or_owned(
retain_statics,
&self.unlimited_vision,
|| -> FocrResult<UnlimitedVisionStatics> {
let started = Instant::now();
let built = UnlimitedVisionStatics {
sam: vision_sam::sam_weights_from(&self.weights, "model.sam_model")?,
projector: vision_bridge::projector_weights_from(&self.weights)?,
};
timing_log(&format!(
" unlimited_vision.hydrate({}) {:.2}s",
if retain_statics {
"cached"
} else {
"batch-local"
},
started.elapsed().as_secs_f64()
));
Ok(built)
},
)?;
let statics = statics_owner.as_ref();
let clip_cfg = vision_clip::ClipConfig::default();
let clip_w = self.clip_weights()?;
timing_log(&format!(
" vision.hydrate(batch) {:.2}s",
th.elapsed().as_secs_f64()
));
let mut groups: std::collections::BTreeMap<usize, Vec<(usize, usize)>> =
std::collections::BTreeMap::new();
for (p, views) in page_views.iter().enumerate() {
for (v, view) in views.iter().enumerate() {
let side = (view.cols as f64).sqrt() as usize;
if side * side != view.cols {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::vision_tower_batched_pages: page {p} view {v} \
cols {} is not a perfect square",
view.cols
)));
}
groups.entry(side).or_default().push((p, v));
}
}
let mut features: Vec<Vec<Option<Mat>>> = page_views
.iter()
.map(|views| views.iter().map(|_| None).collect())
.collect();
for (side, slots) in &groups {
let view_refs: Vec<&Mat> = slots.iter().map(|&(p, v)| &page_views[p][v]).collect();
let ts = Instant::now();
let sams = vision_sam::forward_with_batched(&statics.sam, &view_refs, *side, *side)?;
timing_log(&format!(
" vision.sam(batch of {}, side {side}) {:.2}s",
slots.len(),
ts.elapsed().as_secs_f64()
));
let tc = Instant::now();
let sam_refs: Vec<&Mat> = sams.iter().collect();
let clips = vision_clip::forward_batched_from_sam(&clip_cfg, clip_w, &sam_refs)?;
timing_log(&format!(
" vision.clip(batch of {}) {:.2}s",
slots.len(),
tc.elapsed().as_secs_f64()
));
let tb = Instant::now();
for ((&(p, v), sam), clip) in slots.iter().zip(&sams).zip(&clips) {
let projected = vision_bridge::forward_with(&statics.projector, clip, sam)?;
features[p][v] = Some(projected);
}
timing_log(&format!(
" vision.bridge(batch of {}) {:.2}s",
slots.len(),
tb.elapsed().as_secs_f64()
));
}
page_views.clear();
features
.into_iter()
.enumerate()
.map(|(p, views)| {
views
.into_iter()
.enumerate()
.map(|(v, slot)| {
slot.ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::vision_tower_batched_pages: page {p} \
view {v} produced no feature (grouping bug)"
))
})
})
.collect()
})
.collect()
}
fn build_inputs_embeds(
&self,
pre: &Preprocessed,
vision_features: &[Mat],
) -> FocrResult<(Mat, Vec<u32>)> {
let (prompt_ids, images_seq_mask) = self.build_prompt(pre)?;
let mut inputs_embeds = self.embed_prompt(&prompt_ids)?;
let image_newline = Self::image_newline(&self.weights)?;
let view_seperator = Self::view_seperator(&self.weights)?;
if pre.crop_grid.is_tiled() {
let preprocess::PreprocessMode::Gundam { tile_size, .. } = pre.mode else {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::build_inputs_embeds: tiled crop grid requires Gundam preprocess mode"
)));
};
let local_count = pre.tiles.len();
let expected_feature_blocks = local_count.checked_add(1).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::build_inputs_embeds: local tile count overflow"
))
})?;
if vision_features.len() != expected_feature_blocks {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::build_inputs_embeds: {} vision feature blocks != {} local tiles + 1 global view",
vision_features.len(),
local_count
)));
}
let (locals, global_tail) = vision_features.split_at(local_count);
let q_local = preprocess::num_queries(tile_size);
connector::fuse_crop(
&self.weights,
&mut inputs_embeds,
locals,
pre.crop_grid.width_crop_num,
pre.crop_grid.height_crop_num,
q_local,
q_local,
&global_tail[0],
Self::global_grid_h(pre),
Self::global_grid_w(pre),
&image_newline,
&view_seperator,
&images_seq_mask,
)?;
} else {
connector::fuse_no_crop(
&self.weights,
&mut inputs_embeds,
vision_features,
Self::global_grid_h(pre),
Self::global_grid_w(pre),
&image_newline,
&view_seperator,
&images_seq_mask,
)?;
}
Ok((inputs_embeds, prompt_ids))
}
fn embed_table(&self) -> FocrResult<EmbedTable<'_>> {
embed_table_from(&self.weights)
}
fn embed_prompt(&self, prompt_ids: &[u32]) -> FocrResult<Mat> {
self.embed_table()?.embed_ids(prompt_ids)
}
fn generate(&self, inputs_embeds: Mat, prompt_ids: &[u32]) -> FocrResult<Vec<u32>> {
self.generate_with(inputs_embeds, prompt_ids, &self.decode_params)
}
fn generate_with(
&self,
inputs_embeds: Mat,
prompt_ids: &[u32],
params: &sampler::DecodeParams,
) -> FocrResult<Vec<u32>> {
self.generate_with_observer(inputs_embeds, prompt_ids, params, None)
}
fn generate_with_observer(
&self,
inputs_embeds: Mat,
prompt_ids: &[u32],
params: &sampler::DecodeParams,
observer: Option<&mut dyn FnMut(u32)>,
) -> FocrResult<Vec<u32>> {
let mut runaway_guard = sampler::RunawayGuard::from_env()?;
if std::env::var_os(DECODE_STATELESS_ENV).is_some() {
self.generate_stateless(
inputs_embeds,
prompt_ids,
params,
observer,
&mut runaway_guard,
)
} else if int8_decode_requested() || self.artifact_stores_gated_int8() {
self.generate_cached_i8(
inputs_embeds,
prompt_ids,
params,
observer,
&mut runaway_guard,
)
} else {
self.generate_cached(
inputs_embeds,
prompt_ids,
params,
observer,
&mut runaway_guard,
)
}
}
fn generate_cached(
&self,
inputs_embeds: Mat,
prompt_ids: &[u32],
params: &sampler::DecodeParams,
mut observer: Option<&mut dyn FnMut(u32)>,
runaway_guard: &mut sampler::RunawayGuard,
) -> FocrResult<Vec<u32>> {
timing_log("precision focr-mixed-ffn-int8");
let hidden_dim = inputs_embeds.cols;
let prefill_len = inputs_embeds.rows;
let table = self.embed_table()?;
let vocab = table.rows();
if table.cols() != hidden_dim {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate_cached: embed table hidden {} != inputs_embeds hidden {}",
table.cols(),
hidden_dim
)));
}
let tb = Instant::now();
let wc = self.decoder_cache()?;
timing_log(&format!(
"weight_cache_build {:.2}s",
tb.elapsed().as_secs_f64()
));
let tp = Instant::now();
progress::emit("prefill", 0, prefill_len as u64);
let (hidden, mut caches) = decoder::prefill_with_cache(wc, &inputs_embeds)?;
progress::emit("prefill", prefill_len as u64, prefill_len as u64);
let mut last_hidden = Self::last_hidden_row(&hidden)?;
timing_log(&format!(
"prefill {:.2}s ({} tokens)",
tp.elapsed().as_secs_f64(),
prefill_len
));
let td = Instant::now();
decoder::prof::reset();
let mut generated: Vec<u32> = prompt_ids.to_vec();
let mut emitted: Vec<u32> = Vec::new();
while emitted.len() < params.max_length {
crate::cancel_checkpoint()?;
let logits = decoder::lm_head_cached(wc, &last_hidden)?;
let step: DecodeOutput = sampler::decode_step(&logits, &generated, params)?;
generated.push(step.token_id);
emitted.push(step.token_id);
runaway_guard.check_after_emit(&emitted, step.is_eos)?;
progress::emit("decode", emitted.len() as u64, params.max_length as u64);
if let Some(f) = observer.as_deref_mut() {
f(step.token_id);
}
if step.is_eos {
break;
}
let next = step.token_id as usize;
if next >= vocab {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate_cached: decoded token id {next} outside embed vocab {vocab}"
)));
}
let row = table.row_f32(next);
let token_embed = Mat::from_vec(1, hidden_dim, row);
let position = prefill_len + (emitted.len() - 1);
let h = decoder::decode_step_with_cache(wc, &mut caches, &token_embed, position)?;
last_hidden = Self::last_hidden_row(&h)?;
}
timing_log(&format!(
"decode {:.2}s ({} tokens, {:.3}s/tok)",
td.elapsed().as_secs_f64(),
emitted.len(),
td.elapsed().as_secs_f64() / (emitted.len().max(1) as f64)
));
if decoder::prof::enabled() {
let (lmhead, attn, experts, route) = decoder::prof::snapshot_ms();
timing_log(&format!(
"decode phases (ms): lm_head {lmhead:.0} attn {attn:.0} experts {experts:.0} route {route:.0}"
));
}
if calib::enabled() {
calib::flush()?;
}
Ok(emitted)
}
fn generate_cached_i8(
&self,
inputs_embeds: Mat,
prompt_ids: &[u32],
params: &sampler::DecodeParams,
mut observer: Option<&mut dyn FnMut(u32)>,
runaway_guard: &mut sampler::RunawayGuard,
) -> FocrResult<Vec<u32>> {
self.require_full_int8_authorization()?;
timing_log("precision focr-full-int8");
let hidden_dim = inputs_embeds.cols;
let prefill_len = inputs_embeds.rows;
let table = self.embed_table()?;
let vocab = table.rows();
if table.cols() != hidden_dim {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate_cached_i8: embed table hidden {} != inputs_embeds hidden {}",
table.cols(),
hidden_dim
)));
}
let tb = Instant::now();
let wc = self.decoder_cache_i8()?;
timing_log(&format!(
"weight_cache_build_i8 {:.2}s",
tb.elapsed().as_secs_f64()
));
let tp = Instant::now();
progress::emit("prefill", 0, prefill_len as u64);
let (hidden, mut caches) = decoder::prefill_with_cache_i8(wc, &inputs_embeds)?;
progress::emit("prefill", prefill_len as u64, prefill_len as u64);
let mut last_hidden = Self::last_hidden_row(&hidden)?;
timing_log(&format!(
"prefill_i8 {:.2}s ({} tokens)",
tp.elapsed().as_secs_f64(),
prefill_len
));
let td = Instant::now();
decoder::prof::reset();
let mut generated: Vec<u32> = prompt_ids.to_vec();
let mut emitted: Vec<u32> = Vec::new();
let fuse_ngram_lmhead = decoder::fuse_ngram_lmhead_enabled();
let spec_decode =
spec_decode_enabled() && params.matches_frozen_spec_ban() && observer.is_none();
if spec_decode {
self.spec_decode_i8(
wc,
&mut caches,
&last_hidden,
&mut generated,
&mut emitted,
&table,
vocab,
hidden_dim,
prefill_len,
params,
runaway_guard,
)?;
}
while !spec_decode && emitted.len() < params.max_length {
crate::cancel_checkpoint()?;
let step: DecodeOutput = if fuse_ngram_lmhead
&& params.no_repeat_ngram_size > 0
&& generated.len() >= params.no_repeat_ngram_size
{
let banned = sampler::collect_sliding_window_ngram_bans(
&generated,
params.no_repeat_ngram_size,
params.ngram_window,
&[],
sampler::VOCAB_SIZE,
);
let logits = decoder::lm_head_cached_i8_ngram_masked(wc, &last_hidden, &banned)?;
sampler::decode_step_premasked(&logits, params)?
} else {
let logits = decoder::lm_head_cached_i8(wc, &last_hidden)?;
sampler::decode_step(&logits, &generated, params)?
};
generated.push(step.token_id);
emitted.push(step.token_id);
runaway_guard.check_after_emit(&emitted, step.is_eos)?;
progress::emit("decode", emitted.len() as u64, params.max_length as u64);
if let Some(f) = observer.as_deref_mut() {
f(step.token_id);
}
if step.is_eos {
break;
}
let next = step.token_id as usize;
if next >= vocab {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate_cached_i8: decoded token id {next} outside embed vocab {vocab}"
)));
}
let row = table.row_f32(next);
let token_embed = Mat::from_vec(1, hidden_dim, row);
let position = prefill_len + (emitted.len() - 1);
let h = decoder::decode_step_with_cache_i8(wc, &mut caches, &token_embed, position)?;
last_hidden = Self::last_hidden_row(&h)?;
}
timing_log(&format!(
"decode_i8 {:.2}s ({} tokens, {:.3}s/tok)",
td.elapsed().as_secs_f64(),
emitted.len(),
td.elapsed().as_secs_f64() / (emitted.len().max(1) as f64)
));
if decoder::prof::enabled() {
let (lmhead, attn, experts, route) = decoder::prof::snapshot_ms();
timing_log(&format!(
"decode_i8 phases (ms): lm_head {lmhead:.0} attn {attn:.0} experts {experts:.0} route {route:.0}"
));
}
if calib::enabled() {
calib::flush()?;
}
Ok(emitted)
}
#[allow(clippy::too_many_arguments)]
fn spec_decode_i8(
&self,
wc: &decoder::DecoderWeightCacheI8,
caches: &mut [rswa::RingCache],
last_hidden: &Mat,
generated: &mut Vec<u32>,
emitted: &mut Vec<u32>,
table: &EmbedTable<'_>,
vocab: usize,
hidden_dim: usize,
prefill_len: usize,
params: &sampler::DecodeParams,
runaway_guard: &mut sampler::RunawayGuard,
) -> FocrResult<()> {
let mut last_hidden = last_hidden.clone();
while emitted.len() < params.max_length {
crate::cancel_checkpoint()?;
progress::emit("decode", emitted.len() as u64, params.max_length as u64);
let draft = spec::draft_ngram(generated, spec::SPEC_DRAFT_MAX, spec::SPEC_DRAFT_NGRAM);
if draft.is_empty() {
let logits = decoder::lm_head_cached_i8(wc, &last_hidden)?;
let step = sampler::decode_step(&logits, generated, params)?;
generated.push(step.token_id);
emitted.push(step.token_id);
runaway_guard.check_after_emit(emitted, step.is_eos)?;
if step.is_eos {
break;
}
last_hidden = Self::commit_decode_token_i8(
wc,
caches,
table,
vocab,
hidden_dim,
prefill_len,
emitted.len(),
step.token_id,
)?;
continue;
}
let mut draft_embeds: Vec<Mat> = Vec::with_capacity(draft.len());
for &id in &draft {
draft_embeds.push(Self::embed_decode_token_i8(table, vocab, hidden_dim, id)?);
}
let base_position = prefill_len + emitted.len();
let verify_rows =
decoder::verify_forward_i8(wc, &*caches, &draft_embeds, base_position)?;
let mut verify_logits: Vec<Mat> = Vec::with_capacity(draft.len() + 1);
verify_logits.push(decoder::lm_head_cached_i8(wc, &last_hidden)?);
verify_logits.extend(verify_rows);
let emit = spec::resolve_round(generated, &draft, &verify_logits, params)?;
let mut stopped = false;
for &token in &draft[..emit.accepted] {
generated.push(token);
emitted.push(token);
let is_eos = params.eos_token_id == token;
runaway_guard.check_after_emit(emitted, is_eos)?;
if is_eos {
stopped = true;
break;
}
Self::commit_decode_token_i8(
wc,
caches,
table,
vocab,
hidden_dim,
prefill_len,
emitted.len(),
token,
)?;
if emitted.len() >= params.max_length {
stopped = true;
break;
}
}
if stopped {
break;
}
let Some(correction) = emit.correction else {
break;
};
generated.push(correction.token_id);
emitted.push(correction.token_id);
runaway_guard.check_after_emit(emitted, correction.is_eos)?;
if correction.is_eos {
break;
}
last_hidden = Self::commit_decode_token_i8(
wc,
caches,
table,
vocab,
hidden_dim,
prefill_len,
emitted.len(),
correction.token_id,
)?;
}
Ok(())
}
fn embed_decode_token_i8(
table: &EmbedTable<'_>,
vocab: usize,
hidden_dim: usize,
token: u32,
) -> FocrResult<Mat> {
let idx = token as usize;
if idx >= vocab {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate_cached_i8: decoded token id {idx} outside embed vocab {vocab}"
)));
}
let row = table.row_f32(idx);
Ok(Mat::from_vec(1, hidden_dim, row))
}
#[allow(clippy::too_many_arguments)]
fn commit_decode_token_i8(
wc: &decoder::DecoderWeightCacheI8,
caches: &mut [rswa::RingCache],
table: &EmbedTable<'_>,
vocab: usize,
hidden_dim: usize,
prefill_len: usize,
emitted_len: usize,
token: u32,
) -> FocrResult<Mat> {
let token_embed = Self::embed_decode_token_i8(table, vocab, hidden_dim, token)?;
let position = prefill_len + (emitted_len - 1);
let h = decoder::decode_step_with_cache_i8(wc, caches, &token_embed, position)?;
Self::last_hidden_row(&h)
}
fn generate_stateless(
&self,
mut inputs_embeds: Mat,
prompt_ids: &[u32],
params: &sampler::DecodeParams,
mut observer: Option<&mut dyn FnMut(u32)>,
runaway_guard: &mut sampler::RunawayGuard,
) -> FocrResult<Vec<u32>> {
let hidden_dim = inputs_embeds.cols;
let table = self.embed_table()?;
let vocab = table.rows();
if table.cols() != hidden_dim {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate_stateless: embed table hidden {} != inputs_embeds hidden {}",
table.cols(),
hidden_dim
)));
}
let mut generated: Vec<u32> = prompt_ids.to_vec();
let mut emitted: Vec<u32> = Vec::new();
while emitted.len() < params.max_length {
crate::cancel_checkpoint()?;
let hidden = decoder::forward(&self.weights, &inputs_embeds)?;
let last_hidden = Self::last_hidden_row(&hidden)?;
let logits = decoder::lm_head(&self.weights, &last_hidden)?;
let step: DecodeOutput = sampler::decode_step(&logits, &generated, params)?;
generated.push(step.token_id);
emitted.push(step.token_id);
runaway_guard.check_after_emit(&emitted, step.is_eos)?;
progress::emit("decode", emitted.len() as u64, params.max_length as u64);
if let Some(f) = observer.as_deref_mut() {
f(step.token_id);
}
if step.is_eos {
break;
}
let next = step.token_id as usize;
if next >= vocab {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate_stateless: decoded token id {next} outside embed vocab {vocab}"
)));
}
let new_rows = inputs_embeds.rows + 1;
let mut data = std::mem::take(&mut inputs_embeds.data);
data.extend_from_slice(&table.row_f32(next));
inputs_embeds = Mat::from_vec(new_rows, hidden_dim, data);
}
Ok(emitted)
}
fn last_hidden_row(hidden: &Mat) -> FocrResult<Mat> {
if hidden.rows == 0 {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate: decoder forward returned zero hidden rows"
)));
}
let expected_len = hidden.rows.checked_mul(hidden.cols).ok_or_else(|| {
FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate: decoder hidden shape product overflow for [{}, {}]",
hidden.rows,
hidden.cols
))
})?;
if hidden.data.len() != expected_len {
return Err(FocrError::Other(anyhow::anyhow!(
"native_engine::OcrModel::generate: decoder hidden data len {} != rows*cols {} for shape [{}, {}]",
hidden.data.len(),
expected_len,
hidden.rows,
hidden.cols
)));
}
Ok(Mat::from_vec(
1,
hidden.cols,
hidden.row(hidden.rows - 1).to_vec(),
))
}
fn image_dims(pre: &Preprocessed) -> (u32, u32) {
pre.original_size
}
fn views(pre: &Preprocessed) -> Vec<Mat> {
let mut views = Vec::with_capacity(pre.num_views());
views.extend(pre.tiles.iter().map(|tile| tile.pixels.clone()));
views.push(pre.global.pixels.clone());
views
}
fn build_prompt(&self, pre: &Preprocessed) -> FocrResult<(Vec<u32>, Vec<bool>)> {
let tok = self.tokenizer()?;
let n_image = pre.placeholder_token_count();
let text = tok.encode(BASE_PROMPT_TEXT)?;
let total = 1 + n_image + text.len();
let mut ids = Vec::with_capacity(total);
let mut mask = Vec::with_capacity(total);
ids.push(tok.bos_id());
mask.push(false);
for _ in 0..n_image {
ids.push(tok.image_id());
mask.push(true);
}
for id in text {
ids.push(id);
mask.push(false);
}
Ok((ids, mask))
}
fn build_prompt_multi(&self, pres: &[Preprocessed]) -> FocrResult<(Vec<u32>, Vec<bool>)> {
let tok = self.tokenizer()?;
let text = tok.encode(MULTI_PAGE_PROMPT_TEXT)?;
let n_image: usize = pres.iter().map(Preprocessed::placeholder_token_count).sum();
let total = 1 + n_image + text.len();
let mut ids = Vec::with_capacity(total);
let mut mask = Vec::with_capacity(total);
ids.push(tok.bos_id());
mask.push(false);
for pre in pres {
for _ in 0..pre.placeholder_token_count() {
ids.push(tok.image_id());
mask.push(true);
}
}
for id in text {
ids.push(id);
mask.push(false);
}
Ok((ids, mask))
}
fn global_grid_h(pre: &Preprocessed) -> usize {
preprocess::num_queries(pre.mode.base_size())
}
fn global_grid_w(pre: &Preprocessed) -> usize {
preprocess::num_queries(pre.mode.base_size())
}
fn image_newline(weights: &Weights) -> FocrResult<Vec<f32>> {
weights.vec("model.image_newline")
}
fn view_seperator(weights: &Weights) -> FocrResult<Vec<f32>> {
weights.vec("model.view_seperator")
}
}
#[cfg(test)]
mod tests {
use super::*;
fn synthetic_int8_embed_artifact(rows: usize, cols: usize) -> Weights {
let arch = model_arch::arch_by_id("got-ocr2").expect("test architecture is registered");
let payload: Vec<u8> = (0..rows * cols)
.map(|i| (((i * 37 + 11) % 255) as i64 - 127) as i8 as u8)
.collect();
let scales: Vec<u8> = (0..rows)
.flat_map(|r| (((r % 13) as f32) * 0.0007 + 0.0031).to_le_bytes())
.collect();
let mut builder = crate::quant::focrq::FocrqBuilder::new()
.with_model_id(arch.id())
.with_license_notice(arch.license_notice());
builder
.add_quantized(
EMBED_TOKENS,
crate::quant::focrq::WriteDType::QInt8PerChan,
vec![rows, cols],
payload,
scales,
0,
0,
)
.expect("valid synthetic QInt8 embedding record");
Weights::from_bytes(builder.build()).expect("synthetic artifact loads")
}
#[test]
fn int8_embed_row_lookup_matches_full_dequant() {
let (rows, cols) = (67usize, 16usize);
let weights = synthetic_int8_embed_artifact(rows, cols);
let full = weights.mat(EMBED_TOKENS).expect("whole-table widen");
assert_eq!(full.shape(), (rows, cols));
let table = embed_table_from(&weights).expect("int8 embed table resolves");
assert!(
matches!(table, EmbedTable::QInt8 { .. }),
"a QInt8PerChan table must take the zero-copy int8 arm"
);
assert_eq!((table.rows(), table.cols()), (rows, cols));
for r in 0..rows {
assert_eq!(
table.row_f32(r),
full.data[r * cols..(r + 1) * cols].to_vec(),
"per-row dequant of row {r} must be bit-identical to the full dequant"
);
}
let ids: Vec<u32> = vec![0, 5, 66, 5, 42];
let gathered = table.embed_ids(&ids).expect("gather over valid ids");
assert_eq!(gathered.shape(), (ids.len(), cols));
for (i, &id) in ids.iter().enumerate() {
let want = &full.data[id as usize * cols..(id as usize + 1) * cols];
assert_eq!(&gathered.data[i * cols..(i + 1) * cols], want, "id {id}");
}
assert!(
table.embed_ids(&[rows as u32]).is_err(),
"an out-of-vocab id must still error"
);
}
#[test]
fn forward_admission_serializes_concurrent_callers() {
use std::sync::Barrier;
use std::sync::atomic::{AtomicUsize, Ordering};
let start = Barrier::new(5);
let live = AtomicUsize::new(0);
let max_live = AtomicUsize::new(0);
std::thread::scope(|scope| {
let mut workers = Vec::new();
for _ in 0..4 {
workers.push(scope.spawn(|| {
start.wait();
let _permit = forward_admission().acquire();
let now = live.fetch_add(1, Ordering::SeqCst) + 1;
max_live.fetch_max(now, Ordering::SeqCst);
std::thread::sleep(std::time::Duration::from_millis(5));
live.fetch_sub(1, Ordering::SeqCst);
}));
}
start.wait();
for worker in workers {
worker.join().expect("gate contender did not panic");
}
});
assert_eq!(max_live.load(Ordering::SeqCst), 1);
assert_eq!(live.load(Ordering::SeqCst), 0);
}
#[test]
fn fallible_once_initializes_once_and_retries_after_error() {
use std::sync::Barrier;
use std::sync::atomic::{AtomicUsize, Ordering};
let cell = FallibleOnce::<usize>::new();
let start = Barrier::new(5);
let initializers = AtomicUsize::new(0);
std::thread::scope(|scope| {
let mut workers = Vec::new();
for _ in 0..4 {
workers.push(scope.spawn(|| {
start.wait();
let value = cell
.get_or_try_init(|| {
initializers.fetch_add(1, Ordering::SeqCst);
std::thread::sleep(std::time::Duration::from_millis(5));
Ok::<usize, ()>(42)
})
.expect("initializer succeeds");
assert_eq!(*value, 42);
}));
}
start.wait();
for worker in workers {
worker.join().expect("initializer contender did not panic");
}
});
assert_eq!(initializers.load(Ordering::SeqCst), 1);
let retry = FallibleOnce::<usize>::new();
assert_eq!(
retry
.get_or_try_init(|| Err::<usize, _>("first attempt"))
.expect_err("first initialization must fail"),
"first attempt"
);
assert_eq!(
*retry
.get_or_try_init(|| Ok::<usize, &str>(7))
.expect("failed initialization remains retryable"),
7
);
}
#[test]
fn resolve_model_rejects_missing_path() {
let missing = Path::new("/definitely/not/a/real/model/path.focrq");
let r = OcrModel::resolve_model(missing);
assert!(matches!(r, Err(FocrError::ModelNotFound(_))));
}
#[test]
fn load_missing_path_is_model_not_found() {
let missing = Path::new("/definitely/not/a/real/model/path.focrq");
let r = OcrModel::load(missing);
assert!(matches!(r, Err(FocrError::ModelNotFound(_))));
}
#[test]
fn load_existing_non_model_path_is_format_mismatch_not_panic() -> FocrResult<()> {
let mut tmp = std::env::temp_dir();
tmp.push(format!(
"franken_ocr_load_test_{}.focrq",
std::process::id()
));
std::fs::write(&tmp, b"not a real model blob").expect("write temp file");
let r = OcrModel::load(&tmp);
let _ = std::fs::remove_file(&tmp); let err = match r {
Ok(_) => {
return Err(FocrError::Other(anyhow::anyhow!(
"junk artifact must error, not load"
)));
}
Err(e) => e,
};
assert!(
matches!(err, FocrError::FormatMismatch(_)),
"expected FormatMismatch (exit 7) on a junk artifact, got {err:?}"
);
assert_eq!(err.exit_code(), 7, "FormatMismatch maps to exit 7");
Ok(())
}
fn temp_model_dir(label: &str) -> PathBuf {
let mut dir = std::env::temp_dir();
dir.push(format!("franken_ocr_{label}_{}", std::process::id()));
std::fs::create_dir_all(&dir).expect("create temp model dir");
dir
}
#[test]
fn resolve_model_searches_short_name_focrq_candidate() {
let dir = temp_model_dir("resolve_short_focrq");
let model = dir.join("unlimited-ocr.focrq");
std::fs::write(&model, b"only resolver existence is tested").expect("write model");
let resolved = resolve_model_from_search_dirs(Path::new("unlimited-ocr"), &[dir])
.expect("resolve short-name focrq candidate");
assert_eq!(resolved, model);
}
#[test]
fn resolve_model_searches_short_name_safetensors_directory() {
let root = temp_model_dir("resolve_short_safetensors_root");
let package = root.join("unlimited-ocr");
std::fs::create_dir_all(&package).expect("create safetensors package");
let shard = package.join(RAW_SAFETENSORS_SHARD_NAME);
std::fs::write(&shard, b"only resolver existence is tested").expect("write shard");
let resolved = resolve_model_from_search_dirs(Path::new("unlimited-ocr"), &[root])
.expect("resolve short-name safetensors directory");
assert_eq!(resolved, shard);
}
#[test]
fn resolve_model_quant_preference_picks_matching_focrq_candidate() {
let dir = temp_model_dir("resolve_quant_preference");
let generic = dir.join("unlimited-ocr.focrq");
let int4 = dir.join("unlimited-ocr.int4.focrq");
std::fs::write(&generic, weights::FOCRQ_MAGIC).expect("write generic model");
std::fs::write(&int4, weights::FOCRQ_MAGIC).expect("write int4 model");
let resolved = resolve_model_from_search_dirs_with_quant(
Path::new("unlimited-ocr"),
&[dir],
Some(ModelQuantPreference::Int4),
)
.expect("resolve quant-specific focrq candidate");
assert_eq!(resolved, int4);
}
#[test]
fn resolve_model_accepts_model_dir_direct_focrq_artifact() {
let dir = temp_model_dir("resolve_model_dir_direct_focrq");
let model = dir.join("custom.focrq");
std::fs::write(&model, weights::FOCRQ_MAGIC).expect("write model");
let resolved =
resolve_model_from_search_dirs(Path::new("anything"), std::slice::from_ref(&model))
.expect("resolve direct artifact from model dir entry");
assert_eq!(resolved, model);
}
#[test]
fn resolve_model_searches_relative_default_basename_in_model_dir() {
let dir = temp_model_dir("resolve_default_basename");
let model = dir.join("unlimited-ocr.focrq");
std::fs::write(&model, weights::FOCRQ_MAGIC).expect("write model");
let resolved =
resolve_model_from_search_dirs(Path::new("models/unlimited-ocr.focrq"), &[dir])
.expect("resolve default basename in model dir");
assert_eq!(resolved, model);
}
#[test]
fn resolve_model_default_spec_finds_pulled_int8_artifact() {
let dir = temp_model_dir("resolve_default_int8");
let int8 = dir.join("unlimited-ocr.int8.focrq");
std::fs::write(&int8, weights::FOCRQ_MAGIC).expect("write int8 model");
let resolved =
resolve_model_from_search_dirs(Path::new("models/unlimited-ocr.focrq"), &[dir])
.expect("default spec resolves the pulled int8 artifact");
assert_eq!(resolved, int8);
}
#[test]
fn resolve_model_prefers_manifest_versioned_pull_over_legacy_quant_name() {
let dir = temp_model_dir("resolve_manifest_versioned_int8");
let current = dir.join(format!(
"unlimited-ocr.v{}.int8.focrq",
crate::UNLIMITED_OCR_ARTIFACT_VERSION
));
let legacy = dir.join("unlimited-ocr.int8.focrq");
std::fs::write(¤t, weights::FOCRQ_MAGIC).expect("write current model");
std::fs::write(&legacy, weights::FOCRQ_MAGIC).expect("write legacy model");
let resolved =
resolve_model_from_search_dirs(Path::new("models/unlimited-ocr.focrq"), &[dir])
.expect("default spec resolves manifest-versioned pull");
assert_eq!(resolved, current);
}
#[test]
fn resolve_model_prefers_exact_focrq_over_quant_variant() {
let dir = temp_model_dir("resolve_exact_over_int8");
let generic = dir.join("unlimited-ocr.focrq");
let int8 = dir.join("unlimited-ocr.int8.focrq");
std::fs::write(&generic, weights::FOCRQ_MAGIC).expect("write generic model");
std::fs::write(&int8, weights::FOCRQ_MAGIC).expect("write int8 model");
let resolved =
resolve_model_from_search_dirs(Path::new("models/unlimited-ocr.focrq"), &[dir])
.expect("resolve exact generic focrq");
assert_eq!(resolved, generic);
}
#[test]
fn resolve_model_missing_short_name_lists_search_dirs() {
let dirs = [
PathBuf::from("/tmp/franken_ocr_missing_a"),
PathBuf::from("/tmp/franken_ocr_missing_b"),
];
let err = resolve_model_from_search_dirs(Path::new("missing-model"), &dirs)
.expect_err("missing short name should fail");
let text = err.to_string();
assert!(matches!(err, FocrError::ModelNotFound(_)));
assert!(text.contains("missing-model"));
assert!(text.contains("/tmp/franken_ocr_missing_a"));
assert!(text.contains("/tmp/franken_ocr_missing_b"));
assert!(text.contains(MODEL_DIR_ENV));
}
fn minimal_focrq_blob(version: u32, header_json: &str, payload: &[u8]) -> Vec<u8> {
let mut blob = Vec::new();
blob.extend_from_slice(weights::FOCRQ_MAGIC);
blob.extend_from_slice(&version.to_le_bytes());
blob.push(0);
blob.extend_from_slice(&[0u8; 32]);
blob.extend_from_slice(&(header_json.len() as u64).to_le_bytes());
blob.extend_from_slice(header_json.as_bytes());
blob.extend_from_slice(payload);
blob
}
#[test]
fn load_future_focrq_version_preserves_format_mismatch() {
let payload = [0u8, 0u8];
let header = "{\"t\":{\"dtype\":\"BF16\",\"shape\":[1],\"byte_offset\":0,\"byte_len\":2}}";
let path = temp_model_dir("future_focrq_version").join("future-version.focrq");
std::fs::write(
&path,
minimal_focrq_blob(weights::FOCRQ_FORMAT_VERSION + 1, header, &payload),
)
.expect("write future-version artifact");
let err = load_weights_from_resolved_model(&path).unwrap_err();
assert!(matches!(err, FocrError::FormatMismatch(_)));
assert_eq!(err.exit_code(), crate::error::EXIT_FORMAT_MISMATCH);
assert!(format!("{err}").contains("requires exactly"));
}
#[test]
fn load_malformed_safetensors_preserves_format_mismatch() {
let mut bytes = Vec::new();
bytes.extend_from_slice(&(1u64).to_le_bytes());
bytes.extend_from_slice(b"{");
let path = temp_model_dir("malformed_safetensors").join("bad.safetensors");
std::fs::write(&path, bytes).expect("write malformed safetensors");
let err = load_weights_from_resolved_model(&path).unwrap_err();
assert!(matches!(err, FocrError::FormatMismatch(_)));
assert_eq!(err.exit_code(), crate::error::EXIT_FORMAT_MISMATCH);
}
fn compatible_test_model_builder() -> crate::quant::focrq::FocrqBuilder {
let arch = model_arch::arch_by_id("got-ocr2").expect("test arch is registered");
let mut builder = crate::quant::focrq::FocrqBuilder::new()
.with_model_id(arch.id())
.with_license_notice(arch.license_notice());
builder
.add_tensor(
"test.base.weight",
crate::quant::focrq::WriteDType::Bf16,
vec![1, 1],
vec![0; 2],
)
.expect("add minimal compatible test-model tensor");
builder
}
fn incomplete_unlimited_builder() -> crate::quant::focrq::FocrqBuilder {
let mut builder =
crate::quant::focrq::FocrqBuilder::new().with_packing_manifest_json(format!(
r#"{{"quant_recipe":"{}"}}"#,
crate::quant::convert::UNLIMITED_OCR_INT8_RECIPE_ID,
));
builder
.add_tensor(
"model.embed_tokens.weight",
crate::quant::focrq::WriteDType::Bf16,
vec![1, 1],
vec![0; 2],
)
.expect("add minimal compatible high-precision tensor");
builder
}
#[test]
fn ocr_model_load_owns_weight_bytes_by_default() -> FocrResult<()> {
let path = temp_model_dir("ocr_model_mmap").join("minimal.focrq");
let blob = compatible_test_model_builder().build();
std::fs::write(&path, blob).expect("write valid minimal focrq");
let model = OcrModel::load(&path)?;
assert!(
!model.weights.is_mapped() || weights::mmap_requested(),
"production load may map only after explicit opt-in"
);
Ok(())
}
#[cfg(unix)]
fn loader_swap_fixture(label: &str, marker: &str) -> (PathBuf, Vec<u8>) {
static NEXT: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
let nonce = NEXT.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
let mut dir = std::env::temp_dir();
dir.push(format!(
"franken_ocr_loader_swap_{label}_{}_{nonce}",
std::process::id()
));
std::fs::create_dir_all(&dir).expect("create loader-swap fixture dir");
let mut builder = compatible_test_model_builder();
builder
.add_tensor(
marker,
crate::quant::focrq::WriteDType::Bf16,
vec![1],
vec![0; 2],
)
.expect("add loader-swap marker tensor");
(dir, builder.build())
}
#[cfg(unix)]
#[test]
fn production_loader_pins_validated_file_across_path_swap() {
let (dir, original_blob) = loader_swap_fixture("path", "test.original");
let model_path = dir.join("active.focrq");
let replacement_path = dir.join("replacement.focrq");
let displaced_path = dir.join("validated-original.focrq");
std::fs::write(&model_path, original_blob).expect("write original model");
let mut replacement = compatible_test_model_builder();
replacement
.add_tensor(
"test.replacement",
crate::quant::focrq::WriteDType::Bf16,
vec![1],
vec![0; 2],
)
.expect("add replacement marker tensor");
std::fs::write(&replacement_path, replacement.build()).expect("write replacement model");
let loaded = load_weights_from_resolved_model_after_validation(&model_path, || {
std::fs::rename(&model_path, &displaced_path).expect("preserve validated path entry");
std::fs::rename(&replacement_path, &model_path).expect("swap active path entry");
})
.expect("load the descriptor validated before the path swap");
assert!(loaded.record("test.original").is_some());
assert!(loaded.record("test.replacement").is_none());
let path_now = Weights::load(&model_path).expect("load replacement now at active path");
assert!(path_now.record("test.original").is_none());
assert!(path_now.record("test.replacement").is_some());
}
#[cfg(unix)]
#[test]
fn production_loader_pins_validated_file_across_symlink_swap() {
use std::os::unix::fs::symlink;
let (dir, original_blob) = loader_swap_fixture("symlink", "test.original_link");
let original_path = dir.join("original.focrq");
let replacement_path = dir.join("replacement.focrq");
let active_link = dir.join("active.focrq");
let next_link = dir.join("next.focrq");
let retired_link = dir.join("validated-link.focrq");
std::fs::write(&original_path, original_blob).expect("write original symlink target");
let mut replacement = compatible_test_model_builder();
replacement
.add_tensor(
"test.replacement_link",
crate::quant::focrq::WriteDType::Bf16,
vec![1],
vec![0; 2],
)
.expect("add replacement symlink marker tensor");
std::fs::write(&replacement_path, replacement.build())
.expect("write replacement symlink target");
symlink(&original_path, &active_link).expect("link active path to original");
symlink(&replacement_path, &next_link).expect("prepare replacement symlink");
let loaded = load_weights_from_resolved_model_after_validation(&active_link, || {
std::fs::rename(&active_link, &retired_link).expect("preserve validated symlink");
std::fs::rename(&next_link, &active_link).expect("swap active symlink");
})
.expect("load the descriptor validated before the symlink swap");
assert!(loaded.record("test.original_link").is_some());
assert!(loaded.record("test.replacement_link").is_none());
let link_now = Weights::load(&active_link).expect("load replacement symlink target");
assert!(link_now.record("test.original_link").is_none());
assert!(link_now.record("test.replacement_link").is_some());
}
fn add_unit_qint8(builder: &mut crate::quant::focrq::FocrqBuilder, name: &str) {
builder
.add_quantized(
name,
crate::quant::focrq::WriteDType::QInt8PerChan,
vec![1, 1],
vec![0],
1.0f32.to_le_bytes().to_vec(),
0,
0,
)
.expect("add synthetic qint8 tensor");
}
#[test]
fn unlimited_focrq_rejects_int8_attention_and_lm_head() {
for name in ["model.layers.0.self_attn.q_proj.weight", "lm_head.weight"] {
let mut builder = crate::quant::focrq::FocrqBuilder::new();
add_unit_qint8(&mut builder, name);
let weights = Weights::from_bytes(builder.build()).expect("parse synthetic artifact");
let err = validate_unlimited_ocr_quant_recipe(&weights)
.expect_err("gated tensor must stay high precision in the default recipe");
assert!(matches!(err, FocrError::FormatMismatch(_)));
assert!(err.to_string().contains(name));
assert!(
err.to_string()
.contains(crate::quant::convert::UNLIMITED_OCR_INT8_RECIPE_ID)
);
}
}
#[test]
fn unlimited_focrq_requires_int8_for_present_ffn_weights() {
let mut builder = crate::quant::focrq::FocrqBuilder::new();
builder
.add_tensor(
"model.layers.0.mlp.down_proj.weight",
crate::quant::focrq::WriteDType::Bf16,
vec![1, 1],
vec![0; 2],
)
.expect("add synthetic bf16 FFN");
let weights = Weights::from_bytes(builder.build()).expect("parse synthetic artifact");
let err = validate_unlimited_ocr_quant_recipe(&weights)
.expect_err("validated FFN tensor must use int8 storage");
assert!(err.to_string().contains("expected QInt8PerChan"));
}
fn add_unit_qint4(builder: &mut crate::quant::focrq::FocrqBuilder, name: &str) {
builder
.add_quantized(
name,
crate::quant::focrq::WriteDType::QInt4PerGroup,
vec![1, 16],
vec![0x21; 8],
1.0f32.to_le_bytes().to_vec(),
16,
0,
)
.expect("add synthetic qint4 tensor");
}
fn wasm_recipe_manifest_json() -> String {
format!(
r#"{{"quant_recipe":"{}"}}"#,
crate::quant::convert::UNLIMITED_OCR_WASM_INT4_RECIPE_ID,
)
}
#[test]
fn wasm_recipe_records_accept_int4_experts_and_int8_gated_sets() {
let wasm = Some(crate::quant::convert::UNLIMITED_OCR_WASM_INT4_RECIPE_ID);
validate_unlimited_ocr_quant_records(
true,
model_arch::default_arch().id(),
wasm,
[
(
"model.layers.3.mlp.experts.7.up_proj.weight",
DType::QInt4PerGroup,
),
("model.layers.0.mlp.down_proj.weight", DType::QInt4PerGroup),
(
"model.layers.11.mlp.shared_experts.gate_proj.weight",
DType::QInt4PerGroup,
),
(
"model.layers.0.self_attn.q_proj.weight",
DType::QInt8PerChan,
),
("lm_head.weight", DType::QInt8PerChan),
("model.embed_tokens.weight", DType::QInt8PerChan),
("model.norm.weight", DType::BF16),
("model.layers.5.mlp.gate.weight", DType::BF16), ],
)
.expect("the declared wasm recipe must accept its own storage dtypes");
}
#[test]
fn wasm_recipe_records_reject_int4_attention_and_wrong_expert_storage() {
let wasm = Some(crate::quant::convert::UNLIMITED_OCR_WASM_INT4_RECIPE_ID);
let err = validate_unlimited_ocr_quant_records(
true,
model_arch::default_arch().id(),
wasm,
[(
"model.layers.0.self_attn.q_proj.weight",
DType::QInt4PerGroup,
)],
)
.expect_err("int4 attention must be rejected under the wasm recipe");
assert!(matches!(err, FocrError::FormatMismatch(_)));
assert!(err.to_string().contains("expected QInt8PerChan"));
assert!(
err.to_string()
.contains(crate::quant::convert::UNLIMITED_OCR_WASM_INT4_RECIPE_ID)
);
for dtype in [DType::QInt8PerChan, DType::BF16] {
let err = validate_unlimited_ocr_quant_records(
true,
model_arch::default_arch().id(),
wasm,
[("model.layers.3.mlp.experts.7.up_proj.weight", dtype)],
)
.expect_err("non-int4 expert storage must be rejected under the wasm recipe");
assert!(err.to_string().contains("expected QInt4PerGroup"));
}
let err = validate_unlimited_ocr_quant_records(
true,
model_arch::default_arch().id(),
wasm,
[("model.layers.5.mlp.gate.weight", DType::QInt8PerChan)],
)
.expect_err("quantized router gate must be rejected");
assert!(err.to_string().contains("expected BF16 or F32"));
}
#[test]
fn conservative_recipe_still_rejects_int4_expert_storage() {
for declared in [
None,
Some(crate::quant::convert::UNLIMITED_OCR_INT8_RECIPE_ID),
] {
let err = validate_unlimited_ocr_quant_records(
true,
model_arch::default_arch().id(),
declared,
[(
"model.layers.3.mlp.experts.7.up_proj.weight",
DType::QInt4PerGroup,
)],
)
.expect_err("int4 experts must stay rejected outside the wasm recipe");
assert!(err.to_string().contains("expected QInt8PerChan"));
assert!(
err.to_string()
.contains(crate::quant::convert::UNLIMITED_OCR_INT8_RECIPE_ID)
);
}
}
#[test]
fn wasm_recipe_artifact_passes_dtype_validation_and_fails_only_the_census() {
let mut builder = crate::quant::focrq::FocrqBuilder::new()
.with_packing_manifest_json(wasm_recipe_manifest_json())
.with_source_sha256([
0x2b, 0xc4, 0x8a, 0x7a, 0x11, 0x00, 0x61, 0xea, 0x58, 0xff, 0xf6, 0x5d, 0x31, 0x69,
0x36, 0x7e, 0xeb, 0xe3, 0xae, 0xe3, 0x71, 0xca, 0x69, 0x68, 0xdc, 0x22, 0x19, 0xc1,
0xb2, 0x85, 0x5f, 0xc6,
]);
add_unit_qint4(&mut builder, "model.layers.3.mlp.experts.7.up_proj.weight");
add_unit_qint8(&mut builder, "model.layers.0.self_attn.q_proj.weight");
add_unit_qint8(&mut builder, "lm_head.weight");
let weights = Weights::from_bytes(builder.build()).expect("synthetic wasm artifact parses");
assert_eq!(
weights.quant_recipe(),
Some(crate::quant::convert::UNLIMITED_OCR_WASM_INT4_RECIPE_ID),
"the loader must retain the declared recipe id"
);
let error = validate_unlimited_ocr_quant_recipe(&weights)
.expect_err("a 3-tensor subset must still fail the completeness census");
let text = error.to_string();
assert!(
text.contains("expected 2710 tensors, found 3"),
"failure must be the census, not a dtype-recipe violation: {text}"
);
assert!(
!text.contains("violates quant recipe"),
"wasm dtypes must not be flagged as recipe violations: {text}"
);
}
#[test]
fn weights_without_manifest_report_no_quant_recipe() {
let mut builder = crate::quant::focrq::FocrqBuilder::new();
add_unit_qint8(&mut builder, "model.layers.0.mlp.down_proj.weight");
let weights = Weights::from_bytes(builder.build()).expect("synthetic artifact parses");
assert_eq!(weights.quant_recipe(), None);
}
#[test]
fn focrq_int4_invalid_group_size_is_rejected_with_format_mismatch() {
let mut builder = crate::quant::focrq::FocrqBuilder::new()
.with_packing_manifest_json(wasm_recipe_manifest_json());
add_unit_qint4(&mut builder, "model.layers.3.mlp.experts.7.up_proj.weight");
let blob = builder.build();
let needle = b"\"group_size\":16";
let position = blob
.windows(needle.len())
.position(|window| window == needle)
.expect("header must contain the int4 group_size field");
let mut corrupt = blob.clone();
corrupt[position + needle.len() - 2] = b'9';
let err = validate_model_header_from_reader(corrupt.as_slice(), corrupt.len() as u64)
.expect_err("group_size 96 must be rejected");
assert!(matches!(err, FocrError::FormatMismatch(_)));
assert!(
err.to_string().contains("group_size"),
"error must name the corrupt field: {err}"
);
let weights = Weights::from_bytes(corrupt).expect("directory framing is still intact");
let err = weights
.qint4("model.layers.3.mlp.experts.7.up_proj.weight")
.expect_err("qint4 accessor must reject group_size 96");
assert!(matches!(err, FocrError::FormatMismatch(_)));
}
#[test]
fn resolver_skips_wasm_recipe_artifacts_but_explicit_path_loads_them() {
let dir = temp_model_dir("resolver_skips_wasm_recipe");
let wasm_path = dir.join("unlimited-ocr.int8.focrq");
let mut wasm_builder = crate::quant::focrq::FocrqBuilder::new()
.with_packing_manifest_json(wasm_recipe_manifest_json());
add_unit_qint4(
&mut wasm_builder,
"model.layers.3.mlp.experts.7.up_proj.weight",
);
std::fs::write(&wasm_path, wasm_builder.build()).expect("write wasm-recipe artifact");
assert!(focrq_declares_wasm_recipe(&wasm_path));
let err =
resolve_model_from_search_dirs(Path::new("unlimited-ocr"), std::slice::from_ref(&dir))
.expect_err("the resolver must skip the wasm-recipe artifact");
assert!(matches!(err, FocrError::ModelNotFound(_)));
let resolved =
OcrModel::resolve_model(&wasm_path).expect("explicit path must resolve verbatim");
assert_eq!(resolved, wasm_path);
let conservative_path = dir.join("unlimited-ocr.int4.focrq");
let mut conservative_builder = crate::quant::focrq::FocrqBuilder::new()
.with_packing_manifest_json(format!(
r#"{{"quant_recipe":"{}"}}"#,
crate::quant::convert::UNLIMITED_OCR_INT8_RECIPE_ID,
));
add_unit_qint8(
&mut conservative_builder,
"model.layers.0.mlp.down_proj.weight",
);
std::fs::write(&conservative_path, conservative_builder.build())
.expect("write conservative artifact");
assert!(!focrq_declares_wasm_recipe(&conservative_path));
let resolved = resolve_model_from_search_dirs(Path::new("unlimited-ocr"), &[dir])
.expect("a non-wasm artifact must still resolve");
assert_eq!(resolved, conservative_path);
}
#[test]
fn unlimited_focrq_rejects_incomplete_conservative_subset() -> FocrResult<()> {
let mut builder = crate::quant::focrq::FocrqBuilder::new().with_source_sha256([
0x2b, 0xc4, 0x8a, 0x7a, 0x11, 0x00, 0x61, 0xea, 0x58, 0xff, 0xf6, 0x5d, 0x31, 0x69,
0x36, 0x7e, 0xeb, 0xe3, 0xae, 0xe3, 0x71, 0xca, 0x69, 0x68, 0xdc, 0x22, 0x19, 0xc1,
0xb2, 0x85, 0x5f, 0xc6,
]);
add_unit_qint8(&mut builder, "model.layers.0.mlp.down_proj.weight");
builder.add_tensor(
"model.layers.0.self_attn.q_proj.weight",
crate::quant::focrq::WriteDType::Bf16,
vec![1, 1],
vec![0; 2],
)?;
builder.add_tensor(
"lm_head.weight",
crate::quant::focrq::WriteDType::Bf16,
vec![1, 1],
vec![0; 2],
)?;
let weights = Weights::from_bytes(builder.build())?;
let error = validate_unlimited_ocr_quant_recipe(&weights)
.expect_err("a present-only subset must not pass the production census");
assert!(error.to_string().contains("expected 2710 tensors, found 3"));
assert!(error.to_string().contains("missing"));
Ok(())
}
#[test]
fn resolve_model_accepts_safetensors_directory_shard() {
let dir = temp_model_dir("resolve_safetensors_dir");
let shard = dir.join(RAW_SAFETENSORS_SHARD_NAME);
std::fs::write(&shard, b"not loaded by resolver").expect("write shard");
let resolved = OcrModel::resolve_model(&dir).expect("resolve safetensors directory");
assert_eq!(resolved, shard);
}
#[test]
fn load_safetensors_directory_preserves_format_mismatch() {
let dir = temp_model_dir("load_safetensors_dir");
let shard = dir.join(RAW_SAFETENSORS_SHARD_NAME);
let mut bytes = Vec::new();
bytes.extend_from_slice(&(1u64).to_le_bytes());
bytes.extend_from_slice(b"{");
std::fs::write(&shard, bytes).expect("write malformed shard");
let r = OcrModel::load(&dir);
assert!(matches!(&r, Err(FocrError::FormatMismatch(_))));
if let Err(err) = r {
assert_eq!(err.exit_code(), crate::error::EXIT_FORMAT_MISMATCH);
}
}
#[test]
fn native_model_available_requires_a_compatible_focrq_header() {
let mut tmp = std::env::temp_dir();
tmp.push(format!(
"franken_ocr_available_focrq_{}.focrq",
std::process::id()
));
std::fs::write(&tmp, compatible_test_model_builder().build())
.expect("write compatible focrq");
assert!(native_model_available(&tmp));
assert!(load_weights_from_resolved_model(&tmp).is_ok());
let mut magic_only = tmp.clone();
magic_only.set_file_name(format!(
"franken_ocr_available_magic_only_{}.focrq",
std::process::id()
));
std::fs::write(&magic_only, weights::FOCRQ_MAGIC).expect("write magic-only focrq");
assert!(!native_model_available(&magic_only));
assert!(load_weights_from_resolved_model(&magic_only).is_err());
}
#[test]
fn native_model_available_rejects_unsupported_versions_and_wrong_dtype() {
let mut future = compatible_test_model_builder().build();
future[6..10].copy_from_slice(&(weights::FOCRQ_FORMAT_VERSION + 1).to_le_bytes());
let future_path = temp_model_dir("available_future").join("future.focrq");
std::fs::write(&future_path, future).expect("write future focrq");
assert!(!native_model_available(&future_path));
assert!(load_weights_from_resolved_model(&future_path).is_err());
let mut builder = incomplete_unlimited_builder();
add_unit_qint8(&mut builder, "model.layers.0.self_attn.q_proj.weight");
let wrong_recipe_path = temp_model_dir("available_wrong_recipe").join("wrong-recipe.focrq");
std::fs::write(&wrong_recipe_path, builder.build()).expect("write wrong recipe focrq");
assert!(!native_model_available(&wrong_recipe_path));
assert!(load_weights_from_resolved_model(&wrong_recipe_path).is_err());
}
#[test]
fn native_model_available_rejects_version_zero_and_header_disagreement() {
let mut version_zero = compatible_test_model_builder().build();
version_zero[6..10].copy_from_slice(&0u32.to_le_bytes());
let version_zero_path = temp_model_dir("available_version_zero").join("zero.focrq");
std::fs::write(&version_zero_path, version_zero).expect("write v0 focrq");
assert!(!native_model_available(&version_zero_path));
assert!(load_weights_from_resolved_model(&version_zero_path).is_err());
let mut mismatch = compatible_test_model_builder().build();
let marker = b"\"format_version\":1";
let marker_start = mismatch
.windows(marker.len())
.position(|window| window == marker)
.expect("builder emits header format_version");
mismatch[marker_start + marker.len() - 1] = b'0';
let mismatch_path = temp_model_dir("available_version_mismatch").join("mismatch.focrq");
std::fs::write(&mismatch_path, mismatch).expect("write mismatched version focrq");
assert!(!native_model_available(&mismatch_path));
assert!(load_weights_from_resolved_model(&mismatch_path).is_err());
let missing_header_version = serde_json::json!({
"arch_target": 0,
"license_notice": crate::FOCR_MODEL_LICENSE_NOTICE,
"packing_manifest": {
"quant_recipe": crate::quant::convert::UNLIMITED_OCR_INT8_RECIPE_ID,
},
"source_sha256": "",
"tensors": {
"model.embed_tokens.weight": {
"dtype": "BF16",
"shape": [1, 1],
"byte_offset": 0,
"byte_len": 2,
},
},
})
.to_string();
let missing_header_version_path =
temp_model_dir("available_missing_header_version").join("missing-version.focrq");
std::fs::write(
&missing_header_version_path,
minimal_focrq_blob(
weights::FOCRQ_FORMAT_VERSION,
&missing_header_version,
&[0u8; 2],
),
)
.expect("write focrq without header format_version");
assert!(!native_model_available(&missing_header_version_path));
assert!(load_weights_from_resolved_model(&missing_header_version_path).is_err());
}
#[test]
fn native_model_available_rejects_preamble_header_provenance_disagreement() {
let mut arch_mismatch = compatible_test_model_builder().with_arch_target(2).build();
arch_mismatch[10] = 1;
let arch_path = temp_model_dir("available_arch_mismatch").join("arch-mismatch.focrq");
std::fs::write(&arch_path, arch_mismatch).expect("write arch-mismatched focrq");
let arch_error = validate_model_header_file(&arch_path)
.expect_err("header and preamble arch_target must agree");
assert!(arch_error.to_string().contains("header arch_target"));
assert!(!native_model_available(&arch_path));
let mut source_mismatch = compatible_test_model_builder()
.with_source_sha256([7u8; 32])
.build();
source_mismatch[11] = 6;
let source_path = temp_model_dir("available_source_mismatch").join("source-mismatch.focrq");
std::fs::write(&source_path, source_mismatch).expect("write source-mismatched focrq");
let source_error = validate_model_header_file(&source_path)
.expect_err("header and preamble source_sha256 must agree");
assert!(source_error.to_string().contains("header source_sha256"));
assert!(!native_model_available(&source_path));
let mut uppercase_source = compatible_test_model_builder()
.with_source_sha256([0xabu8; 32])
.build();
let marker = b"\"source_sha256\":\"ab";
let marker_start = uppercase_source
.windows(marker.len())
.position(|window| window == marker)
.expect("builder emits lowercase source sha256");
uppercase_source[marker_start + marker.len() - 2] = b'A';
let uppercase_path =
temp_model_dir("available_uppercase_source").join("uppercase-source.focrq");
std::fs::write(&uppercase_path, uppercase_source).expect("write uppercase-hash focrq");
let uppercase_error = validate_model_header_file(&uppercase_path)
.expect_err("header source_sha256 must be lowercase");
assert!(uppercase_error.to_string().contains("64 lowercase hex"));
assert!(!native_model_available(&uppercase_path));
}
#[test]
fn native_model_available_rejects_unknown_arch_target() {
let path = temp_model_dir("available_unknown_arch").join("unknown-arch.focrq");
let blob = compatible_test_model_builder().with_arch_target(4).build();
std::fs::write(&path, blob).expect("write unknown-arch focrq");
let error = validate_model_header_file(&path)
.expect_err("bounded production header validation must reject unknown arch targets");
assert!(error.to_string().contains("arch_target 4 is unsupported"));
assert!(!native_model_available(&path));
assert!(load_weights_from_resolved_model(&path).is_err());
}
#[test]
fn native_model_available_requires_exact_unlimited_quant_recipe_metadata() {
let empty_path = temp_model_dir("available_empty_focrq").join("empty.focrq");
std::fs::write(
&empty_path,
crate::quant::focrq::FocrqBuilder::new()
.with_packing_manifest_json(format!(
r#"{{"quant_recipe":"{}"}}"#,
crate::quant::convert::UNLIMITED_OCR_INT8_RECIPE_ID,
))
.build(),
)
.expect("write empty focrq");
assert!(!native_model_available(&empty_path));
assert!(load_weights_from_resolved_model(&empty_path).is_err());
let missing_path = temp_model_dir("available_missing_recipe").join("missing.focrq");
let mut missing_builder = crate::quant::focrq::FocrqBuilder::new();
missing_builder
.add_tensor(
"model.embed_tokens.weight",
crate::quant::focrq::WriteDType::Bf16,
vec![1, 1],
vec![0; 2],
)
.expect("add recipe-less tensor");
std::fs::write(&missing_path, missing_builder.build()).expect("write recipe-less focrq");
assert!(!native_model_available(&missing_path));
assert!(load_weights_from_resolved_model(&missing_path).is_err());
let wrong_path = temp_model_dir("available_bad_recipe").join("wrong.focrq");
let mut wrong_builder = crate::quant::focrq::FocrqBuilder::new()
.with_packing_manifest_json(r#"{"quant_recipe":"legacy-full-int8"}"#);
wrong_builder
.add_tensor(
"model.embed_tokens.weight",
crate::quant::focrq::WriteDType::Bf16,
vec![1, 1],
vec![0; 2],
)
.expect("add wrong-recipe tensor");
std::fs::write(&wrong_path, wrong_builder.build()).expect("write wrong-recipe focrq");
assert!(!native_model_available(&wrong_path));
assert!(load_weights_from_resolved_model(&wrong_path).is_err());
}
#[test]
fn native_model_available_rejects_malformed_quant_scale_metadata() {
let header = serde_json::json!({
"arch_target": 0,
"format_version": weights::FOCRQ_FORMAT_VERSION,
"license_notice": crate::FOCR_MODEL_LICENSE_NOTICE,
"packing_manifest": {
"quant_recipe": crate::quant::convert::UNLIMITED_OCR_INT8_RECIPE_ID,
},
"source_sha256": "00".repeat(32),
"tensors": {
"model.layers.0.mlp.down_proj.weight": {
"dtype": "QInt8PerChan",
"shape": [1, 1],
"byte_offset": 0,
"byte_len": 1,
"scales_offset": 1,
"scales_len": 0,
"group_size": 0,
"tier": 0,
},
},
})
.to_string();
let path = temp_model_dir("available_bad_quant_scales").join("bad-scales.focrq");
std::fs::write(
&path,
minimal_focrq_blob(weights::FOCRQ_FORMAT_VERSION, &header, &[0u8; 5]),
)
.expect("write malformed quant metadata focrq");
assert!(!native_model_available(&path));
assert!(load_weights_from_resolved_model(&path).is_err());
}
#[test]
fn native_model_available_rejects_overlapping_payload_ranges() {
let got = model_arch::arch_by_id("got-ocr2").expect("GOT test arch is registered");
let header = serde_json::json!({
"arch_target": 0,
"format_version": weights::FOCRQ_FORMAT_VERSION,
"license_notice": got.license_notice(),
"model_id": got.id(),
"source_sha256": "00".repeat(32),
"tensors": {
"test.a": {
"dtype": "BF16",
"shape": [2],
"byte_offset": 0,
"byte_len": 4,
},
"test.b": {
"dtype": "BF16",
"shape": [2],
"byte_offset": 0,
"byte_len": 4,
},
},
})
.to_string();
let path = temp_model_dir("available_overlapping_ranges").join("overlap.focrq");
std::fs::write(
&path,
minimal_focrq_blob(weights::FOCRQ_FORMAT_VERSION, &header, &[0u8; 4]),
)
.expect("write overlapping-range focrq");
let error = validate_model_header_file(&path)
.expect_err("bounded production header validation must reject aliases");
assert!(error.to_string().contains("payload ranges overlap"));
assert!(!native_model_available(&path));
assert!(load_weights_from_resolved_model(&path).is_err());
}
#[test]
fn native_model_available_rejects_incomplete_safetensors_directory_header() {
let dir = temp_model_dir("available_safetensors_dir");
let shard = dir.join(RAW_SAFETENSORS_SHARD_NAME);
let header = br#"{"x":{"dtype":"BF16","shape":[1],"data_offsets":[0,2]}}"#;
let mut bytes = Vec::new();
bytes.extend_from_slice(&(header.len() as u64).to_le_bytes());
bytes.extend_from_slice(header);
bytes.extend_from_slice(&[0u8; 2]);
std::fs::write(&shard, bytes).expect("write safetensors header");
assert!(!native_model_available(&dir));
assert!(load_weights_from_resolved_model(&shard).is_err());
}
#[test]
fn native_model_available_rejects_empty_safetensors_directory() {
let dir = temp_model_dir("available_empty_safetensors");
let shard = dir.join(RAW_SAFETENSORS_SHARD_NAME);
let mut bytes = Vec::new();
bytes.extend_from_slice(&(2u64).to_le_bytes());
bytes.extend_from_slice(b"{}");
std::fs::write(&shard, bytes).expect("write empty safetensors header");
assert!(!native_model_available(&dir));
assert!(load_weights_from_resolved_model(&shard).is_err());
}
#[test]
fn native_model_available_rejects_missing_and_garbage() {
let missing = Path::new("/definitely/not/a/real/model/path.focrq");
assert!(!native_model_available(missing));
let mut tmp = std::env::temp_dir();
tmp.push(format!(
"franken_ocr_available_garbage_{}.bin",
std::process::id()
));
std::fs::write(&tmp, b"not a model").expect("write garbage");
assert!(!native_model_available(&tmp));
}
#[test]
fn default_decode_params_are_single_image_greedy() {
let p = DecodeParams::single_image();
assert!(p.is_greedy());
assert_eq!(p.eos_token_id, sampler::DEFAULT_EOS_TOKEN_ID);
assert_eq!(
p.no_repeat_ngram_size,
sampler::DEFAULT_NO_REPEAT_NGRAM_SIZE
);
assert_eq!(p.ngram_window, sampler::NGRAM_WINDOW_SINGLE);
assert!(p.max_length > 0, "max_length must bound the decode loop");
}
#[test]
fn unlimited_vision_cache_kill_switch_parses_for_both_vision_schedules() {
for value in ["0", "off", "FALSE", " No "] {
assert!(
!unlimited_vision_cache_enabled_for(Some(value)),
"{value:?} must disable retained statics in sequential and batched vision"
);
}
for value in [None, Some("1"), Some("true"), Some("yes")] {
assert!(unlimited_vision_cache_enabled_for(value));
}
}
#[test]
fn batched_vision_cache_off_keeps_statics_batch_local() {
let cache = FallibleOnce::<usize>::new();
let initializations = std::cell::Cell::new(0usize);
let initialize = || {
initializations.set(initializations.get() + 1);
Ok::<usize, ()>(17)
};
let first = retained_or_owned(false, &cache, initialize).expect("batch-local init");
assert!(matches!(first, RetainedOrOwned::Owned(17)));
let second = retained_or_owned(false, &cache, initialize).expect("next batch-local init");
assert!(matches!(second, RetainedOrOwned::Owned(17)));
assert_eq!(initializations.get(), 2, "cache-off rebuilds per batch");
assert!(cache.value.get().is_none(), "cache-off retains no statics");
let retained = retained_or_owned(true, &cache, initialize).expect("retained init");
assert_eq!(*retained.as_ref(), 17);
let reused = retained_or_owned(true, &cache, initialize).expect("retained reuse");
assert_eq!(*reused.as_ref(), 17);
assert_eq!(
initializations.get(),
3,
"cache-on initializes exactly once"
);
}
#[test]
fn batched_vision_matches_per_page_tower_on_real_model() {
let (Ok(model_path), Ok(images)) = (
std::env::var("FOCR_MODEL_PATH"),
std::env::var("FOCR_PARITY_IMAGES"),
) else {
eprintln!("skip-with-SUCCESS: FOCR_MODEL_PATH / FOCR_PARITY_IMAGES unset");
return;
};
let model = OcrModel::load(Path::new(&model_path)).expect("model loads");
let paths: Vec<&str> = images.split(',').filter(|s| !s.is_empty()).collect();
assert!(
!paths.is_empty(),
"FOCR_PARITY_IMAGES must name at least one page"
);
let pres: Vec<Preprocessed> = paths
.iter()
.map(|p| {
preprocess::preprocess_image(
Path::new(p),
preprocess::PreprocessMode::Base { base_size: 1024 },
)
.expect("page preprocesses")
})
.collect();
let prefs: Vec<&Preprocessed> = pres.iter().collect();
let batched = model
.vision_tower_batched_pages(&prefs)
.expect("batched tower runs");
for (k, pre) in pres.iter().enumerate() {
let per_page = model.vision_tower(pre).expect("per-page tower runs");
assert_eq!(batched[k].len(), per_page.len(), "page {k} view count");
for (v, (b, s)) in batched[k].iter().zip(&per_page).enumerate() {
assert_eq!(b.shape(), s.shape(), "page {k} view {v} shape");
let bits_equal = b
.data
.iter()
.zip(&s.data)
.all(|(x, y)| x.to_bits() == y.to_bits());
assert!(
bits_equal,
"page {k} view {v}: batched vision != per-page tower (bit drift)"
);
}
}
}
#[test]
fn preprocess_overrides_resolve_modes() {
use preprocess::PreprocessMode as M;
assert_eq!(
resolve_preprocess_mode(PreprocessOverrides::default()),
M::Base { base_size: 1024 }
);
assert_eq!(
resolve_preprocess_mode(PreprocessOverrides {
base_size: Some(512),
..Default::default()
}),
M::Base { base_size: 512 }
);
assert_eq!(
resolve_preprocess_mode(PreprocessOverrides {
gundam: Some(true),
..Default::default()
}),
M::Gundam {
base_size: 1024,
tile_size: 640
}
);
assert_eq!(
resolve_preprocess_mode(PreprocessOverrides {
gundam: Some(false),
base_size: Some(768),
image_size: Some(512),
}),
M::Base { base_size: 768 }
);
}
#[test]
fn decode_overrides_apply_some_fields_and_keep_none() {
let base = DecodeParams::single_image();
let mut p = DecodeParams::single_image();
apply_decode_overrides(&mut p, DecodeOverrides::default());
assert_eq!(p.max_length, base.max_length);
assert_eq!(p.temperature.to_bits(), base.temperature.to_bits());
assert_eq!(p.no_repeat_ngram_size, base.no_repeat_ngram_size);
assert_eq!(p.ngram_window, base.ngram_window);
let mut p = DecodeParams::single_image();
apply_decode_overrides(
&mut p,
DecodeOverrides {
max_length: Some(700),
temperature: None,
no_repeat_ngram: Some(20),
ngram_window: None,
},
);
assert_eq!(p.max_length, 700);
assert_eq!(p.temperature.to_bits(), base.temperature.to_bits());
assert_eq!(p.no_repeat_ngram_size, 20);
assert_eq!(p.ngram_window, base.ngram_window);
}
#[test]
fn driver_uses_base_1024_global_grid() {
let pre = Preprocessed::default();
assert_eq!(OcrModel::global_grid_h(&pre), 16);
assert_eq!(OcrModel::global_grid_w(&pre), 16);
}
#[test]
fn structural_params_read_through_weights_accessor() {
let w = Weights::default();
assert!(matches!(
OcrModel::image_newline(&w),
Err(FocrError::FormatMismatch(_))
));
assert!(matches!(
OcrModel::view_seperator(&w),
Err(FocrError::FormatMismatch(_))
));
}
#[test]
fn image_dims_come_from_preprocessed_original_size() {
let pre = Preprocessed {
original_size: (123, 45),
..Preprocessed::default()
};
assert_eq!(OcrModel::image_dims(&pre), (123, 45));
}
#[test]
fn views_forward_local_tiles_then_global_thumbnail() {
let global = Mat::from_vec(3, 4, (0..12).map(|v| v as f32).collect());
let tile_a = Mat::from_vec(3, 1, vec![1.0, 2.0, 3.0]);
let tile_b = Mat::from_vec(3, 1, vec![4.0, 5.0, 6.0]);
let pre = Preprocessed {
mode: preprocess::PreprocessMode::Gundam {
base_size: 128,
tile_size: 64,
},
global: preprocess::ViewTensor {
pixels: global.clone(),
height: 2,
width: 2,
},
tiles: vec![
preprocess::ViewTensor {
pixels: tile_a.clone(),
height: 1,
width: 1,
},
preprocess::ViewTensor {
pixels: tile_b.clone(),
height: 1,
width: 1,
},
],
crop_grid: preprocess::CropGrid {
width_crop_num: 2,
height_crop_num: 1,
},
original_size: (640, 320),
};
assert_eq!(OcrModel::views(&pre), vec![tile_a, tile_b, global]);
assert_eq!(OcrModel::global_grid_h(&pre), 2);
assert_eq!(OcrModel::global_grid_w(&pre), 2);
}
#[test]
fn last_hidden_row_rejects_empty_decoder_output_without_panic() {
let hidden = Mat::zeros(0, 4);
let err = OcrModel::last_hidden_row(&hidden).expect_err("expected empty hidden error");
assert!(matches!(err, FocrError::Other(_)));
assert!(err.to_string().contains("zero hidden rows"));
}
#[test]
fn last_hidden_row_rejects_malformed_decoder_output_without_panic() {
let hidden = Mat {
rows: 2,
cols: 3,
data: vec![1.0, 2.0, 3.0, 4.0, 5.0],
};
let err = OcrModel::last_hidden_row(&hidden).expect_err("expected malformed hidden error");
assert!(matches!(err, FocrError::Other(_)));
assert!(err.to_string().contains("data len 5 != rows*cols 6"));
}
#[test]
fn last_hidden_row_extracts_final_row() {
let hidden = Mat::from_vec(3, 2, vec![1.0, 2.0, 3.0, 4.0, 5.0, 6.0]);
let last = OcrModel::last_hidden_row(&hidden).expect("last row");
assert_eq!(last.shape(), (1, 2));
assert_eq!(last.row(0), &[5.0, 6.0]);
}
#[test]
fn crop_figures_crops_only_image_spans_from_source() {
let source = image::DynamicImage::ImageRgb8(image::RgbImage::from_pixel(
120,
80,
image::Rgb([7, 8, 9]),
));
let decoded = concat!(
"<|ref|>title<|/ref|><|det|>[[0,0,500,500]]<|/det|>",
"<|ref|>image<|/ref|><|det|>[[0,0,999,999]]<|/det|>",
);
let figs = OcrModel::crop_figures(decoded, &source, 120, 80, "");
assert_eq!(figs.len(), 1, "only the image span is cropped");
assert_eq!(figs[0].index, 0);
assert_eq!(figs[0].label, "image");
assert_eq!(figs[0].markdown_ref, "");
assert_eq!(figs[0].bbox, [0, 0, 120, 80]);
assert_eq!(figs[0].image.width(), 120);
assert_eq!(figs[0].image.height(), 80);
}
#[test]
fn crop_figures_crops_the_right_subregion() {
let mut buf = image::RgbImage::new(100, 40);
for y in 0..40 {
for x in 0..100 {
let c = if x < 50 {
image::Rgb([255, 0, 0])
} else {
image::Rgb([0, 0, 255])
};
buf.put_pixel(x, y, c);
}
}
let source = image::DynamicImage::ImageRgb8(buf);
let decoded = "<|ref|>image<|/ref|><|det|>[[500,0,999,999]]<|/det|>";
let figs = OcrModel::crop_figures(decoded, &source, 100, 40, "");
assert_eq!(figs.len(), 1);
assert_eq!(figs[0].bbox, [50, 0, 100, 40]);
let crop = figs[0].image.to_rgb8();
assert_eq!(crop.dimensions(), (50, 40));
assert!(
crop.pixels().all(|p| p.0 == [0, 0, 255]),
"the right-half crop must be all blue"
);
}
#[test]
fn crop_figures_skips_degenerate_box() {
let source = image::DynamicImage::ImageRgb8(image::RgbImage::from_pixel(
50,
50,
image::Rgb([0, 0, 0]),
));
let decoded = "<|ref|>image<|/ref|><|det|>[[10,0,10,999]]<|/det|>";
assert!(OcrModel::crop_figures(decoded, &source, 50, 50, "").is_empty());
}
#[test]
fn forward_dispatch_guard_passes_unlimited_ocr_rejects_a_planned_arch() {
use crate::native_engine::model_arch::{
self, DecodeContract, Decoder, ModelArch, Task, TokenizerKind, VisionEncoder,
};
OcrModel::ensure_arch_implemented(model_arch::default_arch())
.expect("unlimited-ocr is implemented");
struct PlannedArch;
impl ModelArch for PlannedArch {
fn id(&self) -> &'static str {
"got-ocr2"
}
fn display_name(&self) -> &'static str {
"GOT-OCR2.0"
}
fn license_notice(&self) -> &'static str {
"Apache-2.0"
}
fn default_artifact_basename(&self) -> &'static str {
"got-ocr2.focrq"
}
fn vision_encoder(&self) -> VisionEncoder {
VisionEncoder::SamVit
}
fn decoder(&self) -> Decoder {
Decoder::Qwen2Dense
}
fn tokenizer(&self) -> TokenizerKind {
TokenizerKind::Qwen2Bpe
}
fn decode_contract(&self) -> DecodeContract {
DecodeContract {
temperature: 0.0,
eos_token_id: 0,
no_repeat_ngram_size: 0,
ngram_window: 0,
}
}
fn tasks(&self) -> &'static [Task] {
&[Task::Ocr]
}
fn implemented(&self) -> bool {
false
}
}
let err = OcrModel::ensure_arch_implemented(&PlannedArch)
.expect_err("a planned arch must be rejected");
assert!(matches!(err, FocrError::NotImplemented(_)), "got {err:?}");
assert_eq!(err.exit_code(), 1);
assert!(
err.to_string().contains("got-ocr2"),
"names the arch: {err}"
);
}
#[test]
fn arch_is_read_from_the_loaded_focrq_model_id_tag() {
use crate::native_engine::model_arch;
let got_notice = model_arch::arch_by_id("got-ocr2")
.expect("got-ocr2 registered")
.license_notice();
let mut builder = crate::quant::focrq::FocrqBuilder::new()
.with_model_id("got-ocr2")
.with_license_notice(got_notice);
builder
.add_tensor(
"model.embed_tokens.weight",
crate::quant::focrq::WriteDType::Bf16,
vec![1, 1],
vec![0; 2],
)
.expect("add minimal tagged tensor");
let blob = builder.build();
let path = std::env::temp_dir().join(format!(
"focr_arch_tag_{}_{}.focrq",
std::process::id(),
std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_nanos()
));
std::fs::write(&path, &blob).unwrap();
let model = OcrModel::load(&path).expect("a tagged .focrq loads");
assert_eq!(
model.arch().id(),
"got-ocr2",
"arch() reads the model_id tag"
);
OcrModel::ensure_arch_implemented(model.arch())
.expect("got-ocr2 is implemented, so its forward is admitted");
let _ = std::fs::remove_file(&path);
}
}