Skip to main content

rlx_vjepa2/
config.rs

1// RLX — versatile ML compiler + runtime.
2// Copyright (C) 2026 Eugene Hauptmann, Nataliya Kosmyna.
3//
4// This program is free software: you can redistribute it and/or modify
5// it under the terms of the GNU General Public License as published by
6// the Free Software Foundation, version 3.
7//
8// This program is distributed in the hope that it will be useful,
9// but WITHOUT ANY WARRANTY; without even the implied warranty of
10// MERCHANTABILITY or FITNESS FOR A PARTICULAR PURPOSE. See the
11// GNU General Public License for more details.
12//
13// You should have received a copy of the GNU General Public License
14// along with this program. If not, see <https://www.gnu.org/licenses/>.
15
16//! V-JEPA2 configuration — mirrors Meta / HuggingFace `config.json`.
17
18use serde::Deserialize;
19use std::path::Path;
20
21/// ImageNet-style mean/std (same as DINOv2 / HF VJEPA2VideoProcessor).
22pub 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    // HF configs carry both `crop_size` and `image_size` (equal); `image_size`
31    // is left as an ignored unknown field to avoid a serde duplicate-field error
32    // (do NOT add `alias = "image_size"` back — it collides with `crop_size`).
33    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    // Predictor
44    #[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    // Attentive pooler (finetuned checkpoints)
57    #[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    /// `facebook/vjepa2-vitg-fpc64-384` — ViT-G, 64 frames, 384².
101    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    /// Per-axis RoPE segment sizes (d, h, w). Matches Meta `RoPEAttention`.
157    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}