Skip to main content

candle_transformers/models/paddleocr_vl/
config.rs

1//! PaddleOCR-VL configuration structures.
2//!
3//! Defines the configuration for the vision encoder, text decoder, and combined model.
4
5use candle_nn::Activation;
6use serde::Deserialize;
7
8fn default_vision_hidden_size() -> usize {
9    1152
10}
11
12fn default_vision_intermediate_size() -> usize {
13    4304
14}
15
16fn default_vision_num_hidden_layers() -> usize {
17    27
18}
19
20fn default_vision_num_attention_heads() -> usize {
21    16
22}
23
24fn default_vision_num_channels() -> usize {
25    3
26}
27
28fn default_vision_image_size() -> usize {
29    384
30}
31
32fn default_vision_patch_size() -> usize {
33    14
34}
35
36fn default_vision_hidden_act() -> Activation {
37    Activation::GeluPytorchTanh
38}
39
40fn default_vision_layer_norm_eps() -> f64 {
41    1e-6
42}
43
44fn default_vision_attention_dropout() -> f64 {
45    0.0
46}
47
48fn default_vision_spatial_merge_size() -> usize {
49    2
50}
51
52/// Vision encoder configuration for PaddleOCR-VL.
53///
54/// Uses a NaViT-style dynamic resolution visual encoder with 2D rotary position embeddings.
55#[derive(Debug, Clone, Deserialize)]
56pub struct VisionConfig {
57    #[serde(default = "default_vision_hidden_size")]
58    pub hidden_size: usize,
59
60    #[serde(default = "default_vision_intermediate_size")]
61    pub intermediate_size: usize,
62
63    #[serde(default = "default_vision_num_hidden_layers")]
64    pub num_hidden_layers: usize,
65
66    #[serde(default = "default_vision_num_attention_heads")]
67    pub num_attention_heads: usize,
68
69    #[serde(default = "default_vision_num_channels")]
70    pub num_channels: usize,
71
72    #[serde(default = "default_vision_image_size")]
73    pub image_size: usize,
74
75    #[serde(default = "default_vision_patch_size")]
76    pub patch_size: usize,
77
78    #[serde(default = "default_vision_hidden_act")]
79    pub hidden_act: Activation,
80
81    #[serde(default = "default_vision_layer_norm_eps")]
82    pub layer_norm_eps: f64,
83
84    #[serde(default = "default_vision_attention_dropout")]
85    pub attention_dropout: f64,
86
87    #[serde(default = "default_vision_spatial_merge_size")]
88    pub spatial_merge_size: usize,
89}
90
91impl Default for VisionConfig {
92    fn default() -> Self {
93        Self {
94            hidden_size: default_vision_hidden_size(),
95            intermediate_size: default_vision_intermediate_size(),
96            num_hidden_layers: default_vision_num_hidden_layers(),
97            num_attention_heads: default_vision_num_attention_heads(),
98            num_channels: default_vision_num_channels(),
99            image_size: default_vision_image_size(),
100            patch_size: default_vision_patch_size(),
101            hidden_act: default_vision_hidden_act(),
102            layer_norm_eps: default_vision_layer_norm_eps(),
103            attention_dropout: default_vision_attention_dropout(),
104            spatial_merge_size: default_vision_spatial_merge_size(),
105        }
106    }
107}
108
109impl VisionConfig {
110    pub fn head_dim(&self) -> usize {
111        self.hidden_size / self.num_attention_heads
112    }
113}
114
115fn default_vocab_size() -> usize {
116    103424
117}
118
119fn default_hidden_size() -> usize {
120    1024
121}
122
123fn default_intermediate_size() -> usize {
124    3072
125}
126
127fn default_num_hidden_layers() -> usize {
128    18
129}
130
131fn default_num_attention_heads() -> usize {
132    16
133}
134
135fn default_num_key_value_heads() -> usize {
136    2
137}
138
139fn default_hidden_act() -> Activation {
140    Activation::Silu
141}
142
143fn default_max_position_embeddings() -> usize {
144    131072
145}
146
147fn default_rms_norm_eps() -> f64 {
148    1e-5
149}
150
151fn default_rope_theta() -> f64 {
152    500000.0
153}
154
155fn default_head_dim() -> usize {
156    128
157}
158
159fn default_use_bias() -> bool {
160    false
161}
162
163fn default_tie_word_embeddings() -> bool {
164    false
165}
166
167fn default_image_token_id() -> u32 {
168    100295
169}
170
171fn default_video_token_id() -> u32 {
172    101307
173}
174
175fn default_vision_start_token_id() -> u32 {
176    101305
177}
178
179fn default_vision_end_token_id() -> u32 {
180    101306
181}
182
183fn default_tokens_per_second() -> usize {
184    25
185}
186
187/// RoPE scaling configuration for multimodal position embeddings.
188#[derive(Debug, Clone, Deserialize)]
189pub struct RopeScaling {
190    /// Sections for multimodal RoPE: [temporal, height, width].
191    /// Splits head_dim/2 into 3 parts for 3D position encoding.
192    /// Default: [16, 24, 24] (total = 64 = head_dim/2 for head_dim=128)
193    #[serde(default = "default_mrope_section")]
194    pub mrope_section: Vec<usize>,
195
196    #[serde(default)]
197    pub rope_type: Option<String>,
198}
199
200fn default_mrope_section() -> Vec<usize> {
201    vec![16, 24, 24]
202}
203
204impl Default for RopeScaling {
205    fn default() -> Self {
206        Self {
207            mrope_section: default_mrope_section(),
208            rope_type: Some("default".to_string()),
209        }
210    }
211}
212
213/// Combined configuration for PaddleOCR-VL model.
214///
215/// The text model parameters are at the top level (not nested in `text_config`),
216/// following the HuggingFace format where the main model config contains LLM params directly.
217#[derive(Debug, Clone, Deserialize)]
218pub struct Config {
219    // Vision config (nested)
220    #[serde(default)]
221    pub vision_config: VisionConfig,
222
223    // Text model parameters (at top level)
224    #[serde(default = "default_vocab_size")]
225    pub vocab_size: usize,
226
227    #[serde(default = "default_hidden_size")]
228    pub hidden_size: usize,
229
230    #[serde(default = "default_intermediate_size")]
231    pub intermediate_size: usize,
232
233    #[serde(default = "default_num_hidden_layers")]
234    pub num_hidden_layers: usize,
235
236    #[serde(default = "default_num_attention_heads")]
237    pub num_attention_heads: usize,
238
239    #[serde(default = "default_num_key_value_heads")]
240    pub num_key_value_heads: usize,
241
242    #[serde(default = "default_hidden_act")]
243    pub hidden_act: Activation,
244
245    #[serde(default = "default_max_position_embeddings")]
246    pub max_position_embeddings: usize,
247
248    #[serde(default = "default_rms_norm_eps", alias = "rms_norm_eps")]
249    pub layer_norm_eps: f64,
250
251    #[serde(default = "default_rope_theta")]
252    pub rope_theta: f64,
253
254    #[serde(default = "default_head_dim")]
255    pub head_dim: usize,
256
257    #[serde(default = "default_use_bias")]
258    pub use_bias: bool,
259
260    #[serde(default = "default_tie_word_embeddings")]
261    pub tie_word_embeddings: bool,
262
263    // Special token IDs
264    #[serde(default = "default_image_token_id")]
265    pub image_token_id: u32,
266
267    #[serde(default = "default_video_token_id")]
268    pub video_token_id: u32,
269
270    #[serde(default = "default_vision_start_token_id")]
271    pub vision_start_token_id: u32,
272
273    #[serde(default = "default_vision_end_token_id")]
274    pub vision_end_token_id: u32,
275
276    /// RoPE scaling configuration for multimodal position embeddings.
277    #[serde(default)]
278    pub rope_scaling: Option<RopeScaling>,
279
280    /// Tokens per second for video temporal position encoding.
281    #[serde(default = "default_tokens_per_second")]
282    pub tokens_per_second: usize,
283}
284
285impl Default for Config {
286    fn default() -> Self {
287        Self {
288            vision_config: VisionConfig::default(),
289            vocab_size: default_vocab_size(),
290            hidden_size: default_hidden_size(),
291            intermediate_size: default_intermediate_size(),
292            num_hidden_layers: default_num_hidden_layers(),
293            num_attention_heads: default_num_attention_heads(),
294            num_key_value_heads: default_num_key_value_heads(),
295            hidden_act: default_hidden_act(),
296            max_position_embeddings: default_max_position_embeddings(),
297            layer_norm_eps: default_rms_norm_eps(),
298            rope_theta: default_rope_theta(),
299            head_dim: default_head_dim(),
300            use_bias: default_use_bias(),
301            tie_word_embeddings: default_tie_word_embeddings(),
302            image_token_id: default_image_token_id(),
303            video_token_id: default_video_token_id(),
304            vision_start_token_id: default_vision_start_token_id(),
305            vision_end_token_id: default_vision_end_token_id(),
306            rope_scaling: Some(RopeScaling::default()),
307            tokens_per_second: default_tokens_per_second(),
308        }
309    }
310}
311
312/// Helper struct for text config (used internally).
313/// This provides a view of the text-related config fields.
314#[derive(Debug, Clone)]
315pub struct TextConfig {
316    pub vocab_size: usize,
317    pub hidden_size: usize,
318    pub intermediate_size: usize,
319    pub num_hidden_layers: usize,
320    pub num_attention_heads: usize,
321    pub num_key_value_heads: usize,
322    pub hidden_act: Activation,
323    pub max_position_embeddings: usize,
324    pub rms_norm_eps: f64,
325    pub rope_theta: f64,
326    pub head_dim: usize,
327    pub use_bias: bool,
328    pub tie_word_embeddings: bool,
329    /// Multimodal RoPE sections: [temporal, height, width].
330    pub mrope_section: Vec<usize>,
331}
332
333impl From<&Config> for TextConfig {
334    fn from(cfg: &Config) -> Self {
335        let mrope_section = cfg
336            .rope_scaling
337            .as_ref()
338            .map(|rs| rs.mrope_section.clone())
339            .unwrap_or_else(default_mrope_section);
340        Self {
341            vocab_size: cfg.vocab_size,
342            hidden_size: cfg.hidden_size,
343            intermediate_size: cfg.intermediate_size,
344            num_hidden_layers: cfg.num_hidden_layers,
345            num_attention_heads: cfg.num_attention_heads,
346            num_key_value_heads: cfg.num_key_value_heads,
347            hidden_act: cfg.hidden_act,
348            max_position_embeddings: cfg.max_position_embeddings,
349            rms_norm_eps: cfg.layer_norm_eps,
350            rope_theta: cfg.rope_theta,
351            head_dim: cfg.head_dim,
352            use_bias: cfg.use_bias,
353            tie_word_embeddings: cfg.tie_word_embeddings,
354            mrope_section,
355        }
356    }
357}