Skip to main content

candle_transformers/models/gemma4/
config.rs

1//! Configuration for the Gemma 4 multimodal model.
2
3use candle_nn::Activation;
4
5// ── Text config defaults ────────────────────────────────────────────────────
6
7fn default_attention_bias() -> bool {
8    false
9}
10fn default_head_dim() -> usize {
11    256
12}
13fn default_hidden_activation() -> Activation {
14    Activation::GeluPytorchTanh
15}
16fn default_num_attention_heads() -> usize {
17    8
18}
19fn default_num_key_value_heads() -> usize {
20    4
21}
22fn default_rms_norm_eps() -> f64 {
23    1e-6
24}
25fn default_rope_theta() -> f64 {
26    1_000_000.
27}
28fn default_vocab_size() -> usize {
29    262144
30}
31fn default_query_pre_attn_scalar() -> usize {
32    256
33}
34fn default_max_position_embeddings() -> usize {
35    131072
36}
37fn default_tie_word_embeddings() -> bool {
38    true
39}
40fn default_sliding_window_pattern() -> usize {
41    6
42}
43fn default_global_head_dim() -> usize {
44    512
45}
46fn default_use_flash_attn() -> bool {
47    false
48}
49
50// ── Rope parameters ─────────────────────────────────────────────────────────
51
52#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
53pub struct Gemma4RopeLayerParams {
54    pub rope_theta: Option<f64>,
55    pub rope_type: Option<String>,
56    pub partial_rotary_factor: Option<f64>,
57}
58
59#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
60pub struct Gemma4RopeParameters {
61    pub full_attention: Option<Gemma4RopeLayerParams>,
62    pub sliding_attention: Option<Gemma4RopeLayerParams>,
63    pub rope_theta: Option<f64>,
64    pub rope_type: Option<String>,
65    pub partial_rotary_factor: Option<f64>,
66}
67
68// ── Gemma4TextConfig ────────────────────────────────────────────────────────
69
70#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
71pub struct Gemma4TextConfig {
72    #[serde(default = "default_attention_bias")]
73    pub attention_bias: bool,
74    #[serde(default = "default_head_dim")]
75    pub head_dim: usize,
76    #[serde(default = "default_hidden_activation")]
77    pub hidden_activation: Activation,
78    pub hidden_size: usize,
79    pub intermediate_size: usize,
80    #[serde(default = "default_num_attention_heads")]
81    pub num_attention_heads: usize,
82    pub num_hidden_layers: usize,
83    #[serde(default = "default_num_key_value_heads")]
84    pub num_key_value_heads: usize,
85    #[serde(default = "default_rms_norm_eps")]
86    pub rms_norm_eps: f64,
87    #[serde(default = "default_rope_theta")]
88    pub rope_theta: f64,
89    #[serde(default = "default_vocab_size")]
90    pub vocab_size: usize,
91    pub sliding_window: usize,
92    pub final_logit_softcapping: Option<f64>,
93    #[serde(default = "default_query_pre_attn_scalar")]
94    pub query_pre_attn_scalar: usize,
95    #[serde(default = "default_max_position_embeddings")]
96    pub max_position_embeddings: usize,
97    #[serde(default = "default_tie_word_embeddings")]
98    pub tie_word_embeddings: bool,
99    #[serde(
100        default = "default_sliding_window_pattern",
101        alias = "_sliding_window_pattern"
102    )]
103    pub sliding_window_pattern: usize,
104    pub layer_types: Vec<String>,
105    #[serde(default = "default_global_head_dim")]
106    pub global_head_dim: usize,
107    pub num_global_key_value_heads: Option<usize>,
108    pub rope_parameters: Option<Gemma4RopeParameters>,
109    pub use_bidirectional_attention: Option<String>,
110    #[serde(default = "default_use_flash_attn")]
111    pub use_flash_attn: bool,
112}
113
114impl Gemma4TextConfig {
115    pub fn effective_sliding_window(&self) -> usize {
116        if self.use_bidirectional_attention.as_deref() == Some("all") {
117            (self.sliding_window / 2) + 1
118        } else {
119            self.sliding_window
120        }
121    }
122
123    pub fn partial_rotary_factor(&self) -> f64 {
124        self.rope_parameters
125            .as_ref()
126            .and_then(|rp| rp.full_attention.as_ref())
127            .and_then(|fa| fa.partial_rotary_factor)
128            .unwrap_or(0.25)
129    }
130
131    pub fn rope_local_base_freq(&self) -> f64 {
132        self.rope_parameters
133            .as_ref()
134            .and_then(|rp| rp.sliding_attention.as_ref())
135            .and_then(|sa| sa.rope_theta)
136            .unwrap_or(10000.0)
137    }
138
139    pub fn is_sliding(&self, layer_idx: usize) -> bool {
140        self.layer_types
141            .get(layer_idx)
142            .map(|s| s == "sliding_attention")
143            .unwrap_or(false)
144    }
145}
146
147// ── Vision config defaults ──────────────────────────────────────────────────
148
149fn default_vision_hidden_size() -> usize {
150    768
151}
152fn default_vision_intermediate_size() -> usize {
153    3072
154}
155fn default_vision_num_hidden_layers() -> usize {
156    16
157}
158fn default_vision_num_attention_heads() -> usize {
159    12
160}
161fn default_vision_num_key_value_heads() -> usize {
162    12
163}
164fn default_vision_head_dim() -> usize {
165    64
166}
167fn default_vision_hidden_activation() -> Activation {
168    Activation::GeluPytorchTanh
169}
170fn default_vision_rms_norm_eps() -> f64 {
171    1e-6
172}
173fn default_vision_patch_size() -> usize {
174    16
175}
176fn default_vision_position_embedding_size() -> usize {
177    10240
178}
179fn default_vision_pooling_kernel_size() -> usize {
180    3
181}
182fn default_vision_default_output_length() -> usize {
183    280
184}
185fn default_vision_standardize() -> bool {
186    false
187}
188
189// ── Gemma4VisionConfig ──────────────────────────────────────────────────────
190
191#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
192pub struct Gemma4VisionConfig {
193    #[serde(default = "default_vision_hidden_size")]
194    pub hidden_size: usize,
195    #[serde(default = "default_vision_intermediate_size")]
196    pub intermediate_size: usize,
197    #[serde(default = "default_vision_num_hidden_layers")]
198    pub num_hidden_layers: usize,
199    #[serde(default = "default_vision_num_attention_heads")]
200    pub num_attention_heads: usize,
201    #[serde(default = "default_vision_num_key_value_heads")]
202    pub num_key_value_heads: usize,
203    #[serde(default = "default_vision_head_dim")]
204    pub head_dim: usize,
205    #[serde(default = "default_vision_hidden_activation")]
206    pub hidden_activation: Activation,
207    #[serde(default = "default_vision_rms_norm_eps")]
208    pub rms_norm_eps: f64,
209    #[serde(default = "default_vision_patch_size")]
210    pub patch_size: usize,
211    #[serde(default = "default_vision_position_embedding_size")]
212    pub position_embedding_size: usize,
213    #[serde(default = "default_vision_pooling_kernel_size")]
214    pub pooling_kernel_size: usize,
215    #[serde(default = "default_vision_default_output_length")]
216    pub default_output_length: usize,
217    #[serde(default = "default_vision_standardize")]
218    pub standardize: bool,
219    pub rope_parameters: Option<Gemma4RopeParameters>,
220}
221
222impl Gemma4VisionConfig {
223    pub fn rope_theta(&self) -> f64 {
224        self.rope_parameters
225            .as_ref()
226            .and_then(|rp| {
227                rp.full_attention
228                    .as_ref()
229                    .and_then(|fa| fa.rope_theta)
230                    .or(rp.rope_theta)
231            })
232            .unwrap_or(100.0)
233    }
234}
235
236// ── Audio config defaults ───────────────────────────────────────────────────
237
238fn default_audio_input_feat_size() -> usize {
239    128
240}
241fn default_audio_hidden_size() -> usize {
242    1024
243}
244fn default_conf_attention_chunk_size() -> usize {
245    12
246}
247fn default_conf_attention_context_left() -> usize {
248    13
249}
250fn default_conf_attention_context_right() -> usize {
251    0
252}
253fn default_conf_attention_invalid_logits_value() -> f64 {
254    -1e9
255}
256fn default_conf_attention_logit_cap() -> f64 {
257    50.0
258}
259fn default_conf_num_attention_heads() -> usize {
260    8
261}
262fn default_conf_num_hidden_layers() -> usize {
263    12
264}
265fn default_conf_conv_kernel_size() -> usize {
266    5
267}
268fn default_conf_reduction_factor() -> usize {
269    1
270}
271fn default_conf_residual_weight() -> f64 {
272    0.5
273}
274fn default_sscp_conv_channel_size() -> Vec<usize> {
275    vec![128, 32]
276}
277fn default_sscp_conv_kernel_size() -> Vec<Vec<usize>> {
278    vec![vec![3, 3], vec![3, 3]]
279}
280fn default_sscp_conv_stride_size() -> Vec<Vec<usize>> {
281    vec![vec![2, 2], vec![2, 2]]
282}
283fn default_audio_vocab_size() -> usize {
284    128
285}
286fn default_sscp_conv_group_norm_eps() -> f64 {
287    1e-6
288}
289fn default_sscp_conv_eps() -> f64 {
290    1e-3
291}
292fn default_audio_rms_norm_eps() -> f64 {
293    1e-6
294}
295fn default_gradient_clipping() -> f64 {
296    1e10
297}
298fn default_output_proj_dims() -> Option<usize> {
299    Some(1536)
300}
301
302// ── Gemma4AudioConfig ───────────────────────────────────────────────────────
303
304#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
305pub struct Gemma4AudioConfig {
306    #[serde(default = "default_audio_input_feat_size")]
307    pub input_feat_size: usize,
308    #[serde(default = "default_audio_hidden_size")]
309    pub hidden_size: usize,
310    #[serde(default = "default_output_proj_dims")]
311    pub output_proj_dims: Option<usize>,
312    #[serde(
313        default = "default_conf_attention_chunk_size",
314        alias = "attention_chunk_size"
315    )]
316    pub conf_attention_chunk_size: usize,
317    #[serde(
318        default = "default_conf_attention_context_left",
319        alias = "attention_context_left"
320    )]
321    pub conf_attention_context_left: usize,
322    #[serde(
323        default = "default_conf_attention_context_right",
324        alias = "attention_context_right"
325    )]
326    pub conf_attention_context_right: usize,
327    #[serde(
328        default = "default_conf_attention_invalid_logits_value",
329        alias = "attention_invalid_logits_value"
330    )]
331    pub conf_attention_invalid_logits_value: f64,
332    #[serde(
333        default = "default_conf_attention_logit_cap",
334        alias = "attention_logit_cap"
335    )]
336    pub conf_attention_logit_cap: f64,
337    #[serde(
338        default = "default_conf_num_attention_heads",
339        alias = "num_attention_heads"
340    )]
341    pub conf_num_attention_heads: usize,
342    #[serde(
343        default = "default_conf_num_hidden_layers",
344        alias = "num_hidden_layers"
345    )]
346    pub conf_num_hidden_layers: usize,
347    #[serde(default = "default_conf_conv_kernel_size", alias = "conv_kernel_size")]
348    pub conf_conv_kernel_size: usize,
349    #[serde(default = "default_conf_reduction_factor")]
350    pub conf_reduction_factor: usize,
351    #[serde(default = "default_conf_residual_weight", alias = "residual_weight")]
352    pub conf_residual_weight: f64,
353    #[serde(
354        default = "default_sscp_conv_channel_size",
355        alias = "subsampling_conv_channels"
356    )]
357    pub sscp_conv_channel_size: Vec<usize>,
358    #[serde(default = "default_sscp_conv_kernel_size")]
359    pub sscp_conv_kernel_size: Vec<Vec<usize>>,
360    #[serde(default = "default_sscp_conv_stride_size")]
361    pub sscp_conv_stride_size: Vec<Vec<usize>>,
362    #[serde(default = "default_audio_vocab_size")]
363    pub vocab_size: usize,
364    #[serde(default = "default_sscp_conv_group_norm_eps")]
365    pub sscp_conv_group_norm_eps: f64,
366    #[serde(default = "default_sscp_conv_eps")]
367    pub sscp_conv_eps: f64,
368    #[serde(default = "default_audio_rms_norm_eps")]
369    pub rms_norm_eps: f64,
370    #[serde(default = "default_gradient_clipping")]
371    pub gradient_clipping: f64,
372}
373
374// ── Top-level config defaults ───────────────────────────────────────────────
375
376fn default_image_token_id() -> usize {
377    258880
378}
379fn default_audio_token_id() -> usize {
380    258881
381}
382fn default_video_token_id() -> usize {
383    258884
384}
385
386// ── Gemma4Config ────────────────────────────────────────────────────────────
387
388#[derive(Debug, Clone, PartialEq, serde::Deserialize)]
389pub struct Gemma4Config {
390    pub text_config: Gemma4TextConfig,
391    pub vision_config: Gemma4VisionConfig,
392    pub audio_config: Option<Gemma4AudioConfig>,
393    #[serde(default = "default_image_token_id")]
394    pub image_token_id: usize,
395    #[serde(default = "default_audio_token_id")]
396    pub audio_token_id: usize,
397    #[serde(default = "default_video_token_id")]
398    pub video_token_id: usize,
399}