1use serde::Deserialize;
19use std::path::Path;
20
21pub const IMAGENET_MEAN: [f32; 3] = [0.485, 0.456, 0.406];
23pub const IMAGENET_STD: [f32; 3] = [0.229, 0.224, 0.225];
24
25#[derive(Debug, Clone, Deserialize)]
26pub struct Vjepa2Config {
27 pub hidden_size: usize,
28 pub num_hidden_layers: usize,
29 pub num_attention_heads: usize,
30 pub crop_size: usize,
34 pub patch_size: usize,
35 pub tubelet_size: usize,
36 pub frames_per_clip: usize,
37 #[serde(default = "default_mlp_ratio")]
38 pub mlp_ratio: f64,
39 #[serde(default = "default_ln_eps")]
40 pub layer_norm_eps: f64,
41 #[serde(default = "default_in_chans")]
42 pub in_chans: usize,
43 #[serde(default = "default_pred_hidden")]
45 pub pred_hidden_size: usize,
46 #[serde(default = "default_pred_heads")]
47 pub pred_num_attention_heads: usize,
48 #[serde(default = "default_pred_layers")]
49 pub pred_num_hidden_layers: usize,
50 #[serde(default = "default_pred_mlp_ratio")]
51 pub pred_mlp_ratio: f64,
52 #[serde(default = "default_pred_mask_tokens")]
53 pub pred_num_mask_tokens: usize,
54 #[serde(default = "default_true")]
55 pub pred_zero_init_mask_tokens: bool,
56 #[serde(default = "default_pooler_layers")]
58 pub num_pooler_layers: usize,
59 #[serde(default)]
60 pub num_classes: usize,
61}
62
63fn default_mlp_ratio() -> f64 {
64 48.0 / 11.0
65}
66fn default_ln_eps() -> f64 {
67 1e-6
68}
69fn default_in_chans() -> usize {
70 3
71}
72fn default_pred_hidden() -> usize {
73 384
74}
75fn default_pred_heads() -> usize {
76 12
77}
78fn default_pred_layers() -> usize {
79 12
80}
81fn default_pred_mlp_ratio() -> f64 {
82 4.0
83}
84fn default_pred_mask_tokens() -> usize {
85 10
86}
87fn default_true() -> bool {
88 true
89}
90fn default_pooler_layers() -> usize {
91 3
92}
93
94impl Vjepa2Config {
95 pub fn from_file(path: &Path) -> anyhow::Result<Self> {
96 let data = std::fs::read_to_string(path)?;
97 Ok(serde_json::from_str(&data)?)
98 }
99
100 pub fn vit_g_384() -> Self {
102 Self {
103 hidden_size: 1408,
104 num_hidden_layers: 40,
105 num_attention_heads: 22,
106 crop_size: 384,
107 patch_size: 16,
108 tubelet_size: 2,
109 frames_per_clip: 64,
110 mlp_ratio: 48.0 / 11.0,
111 layer_norm_eps: 1e-6,
112 in_chans: 3,
113 pred_hidden_size: 384,
114 pred_num_attention_heads: 12,
115 pred_num_hidden_layers: 12,
116 pred_mlp_ratio: 4.0,
117 pred_num_mask_tokens: 10,
118 pred_zero_init_mask_tokens: true,
119 num_pooler_layers: 3,
120 num_classes: 0,
121 }
122 }
123
124 pub fn head_dim(&self) -> usize {
125 self.hidden_size / self.num_attention_heads
126 }
127
128 pub fn pred_head_dim(&self) -> usize {
129 self.pred_hidden_size / self.pred_num_attention_heads
130 }
131
132 pub fn intermediate_size(&self) -> usize {
133 (self.hidden_size as f64 * self.mlp_ratio) as usize
134 }
135
136 pub fn pred_intermediate_size(&self) -> usize {
137 (self.pred_hidden_size as f64 * self.pred_mlp_ratio) as usize
138 }
139
140 pub fn pooler_intermediate_size(&self) -> usize {
141 (self.hidden_size as f64 * self.mlp_ratio) as usize
142 }
143
144 pub fn grid_spatial(&self) -> usize {
145 self.crop_size / self.patch_size
146 }
147
148 pub fn grid_temporal(&self) -> usize {
149 self.frames_per_clip / self.tubelet_size
150 }
151
152 pub fn num_patches(&self) -> usize {
153 self.grid_temporal() * self.grid_spatial() * self.grid_spatial()
154 }
155
156 pub fn rope_segment_dims(&self) -> (usize, usize, usize) {
158 rope_segment_dims(self.head_dim())
159 }
160
161 pub fn pred_rope_segment_dims(&self) -> (usize, usize, usize) {
162 rope_segment_dims(self.pred_head_dim())
163 }
164}
165
166pub fn rope_segment_dims(head_dim: usize) -> (usize, usize, usize) {
167 let third = head_dim / 3;
168 let seg = 2 * (third / 2);
169 (seg, seg, seg)
170}