#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum VisionEncoder {
SamClip,
SamVit,
Siglip,
ResNetVit,
BeitVit,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum Decoder {
DeepSeekV2MoeRswa,
Qwen2Dense,
LlamaDense,
OptDense,
Seq2SeqDense,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum TokenizerKind {
DeepSeekBpe,
Qwen2Bpe,
SmolLm2Bpe,
Gpt2Bpe,
MusicVocab,
SentencePiece,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
pub enum Task {
Ocr,
Formula,
Tables,
Chart,
Molecular,
Geometry,
Music,
Describe,
Vqa,
Handwriting,
}
#[derive(Clone, Copy, Debug, PartialEq)]
pub struct DecodeContract {
pub temperature: f32,
pub eos_token_id: u32,
pub no_repeat_ngram_size: usize,
pub ngram_window: usize,
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct QuantPolicy {
pub decoder_gemms_int8: bool,
pub vision_high_precision: bool,
pub lm_head_int8_killswitched: bool,
}
impl QuantPolicy {
pub const DOCTRINE: Self = Self {
decoder_gemms_int8: true,
vision_high_precision: true,
lm_head_int8_killswitched: true,
};
}
pub trait ModelArch: Send + Sync {
fn id(&self) -> &'static str;
fn display_name(&self) -> &'static str;
fn license_notice(&self) -> &'static str;
fn default_artifact_basename(&self) -> &'static str;
fn vision_encoder(&self) -> VisionEncoder;
fn decoder(&self) -> Decoder;
fn tokenizer(&self) -> TokenizerKind;
fn quant_policy(&self) -> QuantPolicy {
QuantPolicy::DOCTRINE
}
fn decode_contract(&self) -> DecodeContract;
fn tasks(&self) -> &'static [Task];
fn implemented(&self) -> bool;
fn tie_word_embeddings(&self) -> bool {
false
}
fn vision_tower_prefix(&self) -> &'static str {
"model.sam_model"
}
fn decoder_layers_prefix(&self) -> &'static str {
"model.layers."
}
fn lm_head_stored_int8(&self) -> bool {
false
}
fn embed_tokens_name(&self) -> &'static str {
"model.embed_tokens.weight"
}
}
#[derive(Clone, Copy, Debug, Default)]
pub struct UnlimitedOcr;
impl ModelArch for UnlimitedOcr {
fn id(&self) -> &'static str {
"unlimited-ocr"
}
fn display_name(&self) -> &'static str {
"Baidu Unlimited-OCR"
}
fn license_notice(&self) -> &'static str {
crate::FOCR_MODEL_LICENSE_NOTICE
}
fn default_artifact_basename(&self) -> &'static str {
"unlimited-ocr.focrq"
}
fn vision_encoder(&self) -> VisionEncoder {
VisionEncoder::SamClip
}
fn decoder(&self) -> Decoder {
Decoder::DeepSeekV2MoeRswa
}
fn tokenizer(&self) -> TokenizerKind {
TokenizerKind::DeepSeekBpe
}
fn decode_contract(&self) -> DecodeContract {
DecodeContract {
temperature: 0.0,
eos_token_id: 1,
no_repeat_ngram_size: 35,
ngram_window: 128,
}
}
fn tasks(&self) -> &'static [Task] {
&[Task::Ocr]
}
fn implemented(&self) -> bool {
true
}
}
static UNLIMITED_OCR: UnlimitedOcr = UnlimitedOcr;
pub struct PlannedArch {
id: &'static str,
display_name: &'static str,
license_notice: &'static str,
default_artifact_basename: &'static str,
vision_encoder: VisionEncoder,
decoder: Decoder,
tokenizer: TokenizerKind,
decode_contract: DecodeContract,
tasks: &'static [Task],
tie_word_embeddings: bool,
vision_tower_prefix: &'static str,
decoder_layers_prefix: &'static str,
lm_head_stored_int8: bool,
embed_tokens_name: &'static str,
implemented: bool,
}
const PLACEHOLDER_CONTRACT: DecodeContract = DecodeContract {
temperature: 0.0,
eos_token_id: 0,
no_repeat_ngram_size: 0,
ngram_window: 0,
};
impl ModelArch for PlannedArch {
fn id(&self) -> &'static str {
self.id
}
fn display_name(&self) -> &'static str {
self.display_name
}
fn license_notice(&self) -> &'static str {
self.license_notice
}
fn default_artifact_basename(&self) -> &'static str {
self.default_artifact_basename
}
fn vision_encoder(&self) -> VisionEncoder {
self.vision_encoder
}
fn decoder(&self) -> Decoder {
self.decoder
}
fn tokenizer(&self) -> TokenizerKind {
self.tokenizer
}
fn decode_contract(&self) -> DecodeContract {
self.decode_contract
}
fn tasks(&self) -> &'static [Task] {
self.tasks
}
fn implemented(&self) -> bool {
self.implemented
}
fn tie_word_embeddings(&self) -> bool {
self.tie_word_embeddings
}
fn vision_tower_prefix(&self) -> &'static str {
self.vision_tower_prefix
}
fn decoder_layers_prefix(&self) -> &'static str {
self.decoder_layers_prefix
}
fn lm_head_stored_int8(&self) -> bool {
self.lm_head_stored_int8
}
fn embed_tokens_name(&self) -> &'static str {
self.embed_tokens_name
}
}
static GOT_OCR2: PlannedArch = PlannedArch {
id: "got-ocr2",
display_name: "GOT-OCR2.0",
license_notice: "GOT-OCR2.0 (StepFun) - Apache-2.0",
default_artifact_basename: "got-ocr2.focrq",
vision_encoder: VisionEncoder::SamVit,
decoder: Decoder::Qwen2Dense,
tokenizer: TokenizerKind::Qwen2Bpe,
decode_contract: DecodeContract {
temperature: 0.0,
eos_token_id: 151_643,
no_repeat_ngram_size: 20,
ngram_window: 0,
},
tasks: &[
Task::Ocr,
Task::Formula,
Task::Tables,
Task::Chart,
Task::Molecular,
Task::Geometry,
Task::Music,
],
tie_word_embeddings: true,
vision_tower_prefix: "model.vision_tower_high",
decoder_layers_prefix: "model.layers.", lm_head_stored_int8: true, embed_tokens_name: "model.embed_tokens.weight",
implemented: true, };
static SMOLVLM2: PlannedArch = PlannedArch {
id: "smolvlm2",
display_name: "SmolVLM2-500M",
license_notice: "SmolVLM2-500M-Video-Instruct (HuggingFaceTB) - Apache-2.0",
default_artifact_basename: "smolvlm2.focrq",
vision_encoder: VisionEncoder::Siglip,
decoder: Decoder::LlamaDense,
tokenizer: TokenizerKind::SmolLm2Bpe,
decode_contract: DecodeContract {
temperature: 0.0,
eos_token_id: 49_279,
no_repeat_ngram_size: 0,
ngram_window: 0,
},
tasks: &[Task::Describe, Task::Vqa],
tie_word_embeddings: false,
vision_tower_prefix: "model.vision_model", decoder_layers_prefix: "model.text_model.layers.", lm_head_stored_int8: false,
embed_tokens_name: "model.text_model.embed_tokens.weight",
implemented: true,
};
static ONECHART: PlannedArch = PlannedArch {
id: "onechart",
display_name: "OneChart",
license_notice: "OneChart (kppkkp) - Apache-2.0",
default_artifact_basename: "onechart.focrq",
vision_encoder: VisionEncoder::SamVit,
decoder: Decoder::OptDense,
tokenizer: TokenizerKind::Gpt2Bpe,
decode_contract: DecodeContract {
temperature: 0.0,
eos_token_id: 2,
no_repeat_ngram_size: 0,
ngram_window: 0,
},
tasks: &[Task::Chart],
tie_word_embeddings: true,
vision_tower_prefix: "model.vision_tower",
decoder_layers_prefix: "model.decoder.layers.",
lm_head_stored_int8: false,
embed_tokens_name: "model.decoder.embed_tokens.weight",
implemented: true,
};
static TROMR: PlannedArch = PlannedArch {
id: "tromr",
display_name: "Polyphonic-TrOMR",
license_notice: "Polyphonic-TrOMR (NetEase) - Apache-2.0",
default_artifact_basename: "tromr.focrq",
vision_encoder: VisionEncoder::ResNetVit,
decoder: Decoder::Seq2SeqDense,
tokenizer: TokenizerKind::MusicVocab,
decode_contract: DecodeContract {
temperature: 0.0,
eos_token_id: 2,
no_repeat_ngram_size: 0,
ngram_window: 0,
},
tasks: &[Task::Music],
tie_word_embeddings: false,
vision_tower_prefix: "encoder.",
decoder_layers_prefix: "decoder.net.attn_layers.layers.",
lm_head_stored_int8: false,
embed_tokens_name: "decoder.net.rhythm_emb.emb.weight",
implemented: true,
};
static TROCR: PlannedArch = PlannedArch {
id: "trocr",
display_name: "TrOCR",
license_notice: "TrOCR (Microsoft) - MIT",
default_artifact_basename: "trocr.focrq",
vision_encoder: VisionEncoder::BeitVit,
decoder: Decoder::Seq2SeqDense,
tokenizer: TokenizerKind::SentencePiece,
decode_contract: PLACEHOLDER_CONTRACT,
tasks: &[Task::Handwriting],
tie_word_embeddings: false, vision_tower_prefix: "model.sam_model", decoder_layers_prefix: "model.layers.", lm_head_stored_int8: true, embed_tokens_name: "model.embed_tokens.weight",
implemented: false,
};
static PIX2TEX: PlannedArch = PlannedArch {
id: "pix2tex",
display_name: "pix2tex (LaTeX-OCR)",
license_notice: "pix2tex / LaTeX-OCR - MIT",
default_artifact_basename: "pix2tex.focrq",
vision_encoder: VisionEncoder::ResNetVit,
decoder: Decoder::Seq2SeqDense,
tokenizer: TokenizerKind::SentencePiece,
decode_contract: PLACEHOLDER_CONTRACT,
tasks: &[Task::Formula],
tie_word_embeddings: false, vision_tower_prefix: "model.sam_model", decoder_layers_prefix: "model.layers.", lm_head_stored_int8: true, embed_tokens_name: "model.embed_tokens.weight",
implemented: false,
};
static REGISTRY: &[&dyn ModelArch] = &[
&UNLIMITED_OCR,
&GOT_OCR2,
&SMOLVLM2,
&ONECHART,
&TROMR,
&TROCR,
&PIX2TEX,
];
#[must_use]
pub fn registry() -> &'static [&'static dyn ModelArch] {
REGISTRY
}
#[must_use]
pub fn arch_by_id(id: &str) -> Option<&'static dyn ModelArch> {
registry().iter().copied().find(|a| a.id() == id)
}
#[must_use]
pub fn default_arch() -> &'static dyn ModelArch {
&UNLIMITED_OCR
}
#[cfg(test)]
mod tests {
use super::*;
use crate::native_engine::sampler;
use crate::{DEFAULT_MODEL_PATH, FOCR_MODEL_LICENSE_NOTICE};
#[test]
fn registry_lists_the_default_first_then_the_planned_zoo() {
let archs = registry();
assert_eq!(archs[0].id(), "unlimited-ocr");
assert!(archs[0].implemented());
let mut implemented: Vec<&str> = archs
.iter()
.filter(|a| a.implemented())
.map(|a| a.id())
.collect();
implemented.sort_unstable();
assert_eq!(
implemented,
["got-ocr2", "onechart", "smolvlm2", "tromr", "unlimited-ocr"]
);
for id in ["trocr", "pix2tex"] {
let a = arch_by_id(id).unwrap_or_else(|| unreachable!("planned arch {id} registered"));
assert!(!a.implemented(), "{id} is planned, not implemented");
}
let mut ids: Vec<&str> = archs.iter().map(|a| a.id()).collect();
ids.sort_unstable();
ids.dedup();
assert_eq!(ids.len(), archs.len(), "model ids must be unique");
}
#[test]
fn tromr_descriptor_is_the_censused_shape() {
let a = arch_by_id("tromr").expect("tromr registered");
assert!(a.implemented(), "E2-E9 shipped (model #5)");
assert_eq!(a.decoder(), Decoder::Seq2SeqDense);
assert_eq!(a.tokenizer(), TokenizerKind::MusicVocab);
assert!(
!a.tie_word_embeddings(),
"4 untied heads, no lm_head at all"
);
assert!(
!a.lm_head_stored_int8(),
"heads are tiny Linear+bias — always HP"
);
assert_eq!(a.vision_tower_prefix(), "encoder.");
assert_eq!(a.decoder_layers_prefix(), "decoder.net.attn_layers.layers.");
assert_eq!(a.embed_tokens_name(), "decoder.net.rhythm_emb.emb.weight");
let c = a.decode_contract();
assert_eq!(c.eos_token_id, 2, "rhythm [EOS]");
assert_eq!(c.temperature, 0.0);
assert_eq!(c.no_repeat_ngram_size, 0);
assert_eq!(
a.license_notice(),
"Polyphonic-TrOMR (NetEase) - Apache-2.0"
);
}
#[test]
fn lookup_and_default_resolve() {
assert_eq!(default_arch().id(), "unlimited-ocr");
assert!(default_arch().implemented());
assert!(
!default_arch().lm_head_stored_int8(),
"Unlimited-OCR keeps its untied lm_head high precision by default"
);
assert_eq!(
arch_by_id("unlimited-ocr").map(ModelArch::id),
Some("unlimited-ocr")
);
let got = arch_by_id("got-ocr2").expect("got-ocr2 is a registered arch");
assert_eq!(got.id(), "got-ocr2");
assert!(got.implemented());
assert_eq!(got.decoder(), Decoder::Qwen2Dense);
assert_eq!(got.tokenizer(), TokenizerKind::Qwen2Bpe);
assert!(arch_by_id("does-not-exist").is_none());
}
#[test]
fn got_ocr2_descriptor_matches_the_census() {
let got = arch_by_id("got-ocr2").expect("got-ocr2 registered");
assert!(got.implemented());
assert_eq!(got.vision_encoder(), VisionEncoder::SamVit);
assert_eq!(got.decoder(), Decoder::Qwen2Dense);
assert_eq!(got.tokenizer(), TokenizerKind::Qwen2Bpe);
let c = got.decode_contract();
assert_eq!(c.temperature, 0.0);
assert_eq!(c.eos_token_id, 151_643);
assert_eq!(c.no_repeat_ngram_size, 20);
assert!(got.quant_policy().vision_high_precision);
assert!(got.quant_policy().decoder_gemms_int8);
}
#[test]
fn smolvlm2_descriptor_matches_the_census() {
let a = arch_by_id("smolvlm2").expect("smolvlm2 registered");
assert!(a.implemented(), "sub-epic C shipped C1-C9 (describe/VQA)");
assert_eq!(a.vision_encoder(), VisionEncoder::Siglip);
assert_eq!(a.decoder(), Decoder::LlamaDense);
assert_eq!(a.tokenizer(), TokenizerKind::SmolLm2Bpe);
assert_eq!(a.tasks(), &[Task::Describe, Task::Vqa]);
assert!(a.license_notice().contains("Apache-2.0"));
assert!(a.license_notice().contains("HuggingFaceTB"));
let c = a.decode_contract();
assert_eq!(c.temperature, 0.0);
assert_eq!(c.eos_token_id, 49_279);
assert_eq!(c.no_repeat_ngram_size, 0);
assert_eq!(c.ngram_window, 0);
assert!(!a.tie_word_embeddings());
assert!(!a.lm_head_stored_int8());
assert_eq!(a.vision_tower_prefix(), "model.vision_model");
assert_eq!(a.decoder_layers_prefix(), "model.text_model.layers.");
assert_eq!(
a.embed_tokens_name(),
"model.text_model.embed_tokens.weight"
);
assert_eq!(a.quant_policy(), QuantPolicy::DOCTRINE);
}
#[test]
fn arch_namespace_defaults_are_the_historical_layout() {
for id in ["got-ocr2", "trocr", "pix2tex"] {
let a = arch_by_id(id).expect("registered");
assert_eq!(a.decoder_layers_prefix(), "model.layers.", "{id}");
assert!(a.lm_head_stored_int8(), "{id}");
assert_eq!(a.embed_tokens_name(), "model.embed_tokens.weight", "{id}");
}
}
#[test]
fn onechart_descriptor_matches_the_census() {
let a = arch_by_id("onechart").expect("onechart registered");
assert!(a.implemented(), "sub-epic D shipped D1-D9 (chart-data)");
assert_eq!(a.vision_encoder(), VisionEncoder::SamVit);
assert_eq!(a.decoder(), Decoder::OptDense);
assert_eq!(a.tokenizer(), TokenizerKind::Gpt2Bpe);
assert_eq!(a.tasks(), &[Task::Chart]);
assert!(a.license_notice().contains("Apache-2.0"));
let c = a.decode_contract();
assert_eq!(c.temperature, 0.0);
assert_eq!(c.eos_token_id, 2);
assert_eq!(c.no_repeat_ngram_size, 0);
assert!(a.tie_word_embeddings());
assert!(!a.lm_head_stored_int8());
assert_eq!(a.vision_tower_prefix(), "model.vision_tower");
assert_eq!(a.decoder_layers_prefix(), "model.decoder.layers.");
assert_eq!(a.embed_tokens_name(), "model.decoder.embed_tokens.weight");
}
#[test]
fn unlimited_ocr_descriptor_matches_the_live_engine() {
let a = UnlimitedOcr;
assert_eq!(a.license_notice(), FOCR_MODEL_LICENSE_NOTICE);
let want_basename = std::path::Path::new(DEFAULT_MODEL_PATH)
.file_name()
.and_then(|s| s.to_str())
.unwrap();
assert_eq!(a.default_artifact_basename(), want_basename);
assert_eq!(a.vision_encoder(), VisionEncoder::SamClip);
assert_eq!(a.decoder(), Decoder::DeepSeekV2MoeRswa);
assert_eq!(a.tokenizer(), TokenizerKind::DeepSeekBpe);
assert_eq!(a.tasks(), &[Task::Ocr]);
assert_eq!(a.quant_policy(), QuantPolicy::DOCTRINE);
assert!(a.quant_policy().vision_high_precision);
assert!(a.quant_policy().decoder_gemms_int8);
}
#[test]
fn unlimited_ocr_decode_contract_matches_sampler_default() {
let c = UnlimitedOcr.decode_contract();
let d = sampler::DecodeParams::default();
assert_eq!(c.temperature, d.temperature);
assert_eq!(c.eos_token_id, d.eos_token_id);
assert_eq!(c.no_repeat_ngram_size, d.no_repeat_ngram_size);
assert_eq!(c.ngram_window, d.ngram_window);
assert_eq!(c.temperature, 0.0);
assert_eq!(c.eos_token_id, 1);
assert_eq!(c.no_repeat_ngram_size, 35);
assert_eq!(c.ngram_window, 128);
}
}