candle_transformers/models/gemma4/
config.rs1use candle_nn::Activation;
4
5fn 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#[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#[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
147fn 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#[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
236fn 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#[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
374fn 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#[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}