--- license: mit tags: - vjepa - vjepa2 - disentangled-representation - robot-manipulation - contrastive-learning --- # Disentangled encoders (vit-Large + LoRA, 4-quad dom pair, 20260514) Disentangle post-training of V-JEPA2 vit-Large with LoRA adapters, producing two encoders that separate **task content** from **domain (scene/cam/aug)** features. ## Architecture - Base: V-JEPA2 ViT-Large (`vit_large`, 24 layers, embed_dim=1024) - LoRA r=32 on `qkv`, `proj`, `fc1`, `fc2` (encoder backbone) - ProjectionHead: AttentivePooler (depth=2, 8 queries) + 2-layer MLP → 4096-d - **Two parallel encoders + heads**: `task_encoder` + `task_head`, `domain_encoder` + `domain_head` (each fully trained on its objective) ## Training config (summary) ```yaml model: model_name: vit_large freeze_encoder: false # LoRA mode lora_rank: 32 lora_alpha: 32 lora_target_modules: [qkv, proj, fc1, fc2] data: fpc: 8 frame_stride: 2 same_task_cam_pair: true # 4-quad dom pair scheme dom_pair_type_b_ratio: 0.5 # Type A intra-ep / Type B cross-task same-cam mix loss: disentangle: proj_dim: 4096 pooler_depth: 2 pooler_num_queries: 8 mlp_proj: true loss_mode: sigreg invariance_mode: infonce infonce_temperature: 0.1 task_inv_coeff: 5.5 alpha: 5.0 sigreg_lam: 0.11 ``` Full config in `config.yaml`. ## Datasets (mixed, weighted) - **Sim**: maniskill-franka-merged, maniskill-xarm (weight 4×), 16 robomimic & robosuite_d0 tasks - **Realworld**: oxe-bridge, oxe-fractal (30% mixing ratio) ## Disentangle pair structure (new in this version) - **Task pair**: 1 ep + 2 non-overlap windows of the same ep, each with a different domain augmentation - **Dom pair (4-quad)**: from one source (1 ep + 2 windows for Type A; or 2 eps from same `(ds, cam_id)` bucket with different tasks for Type B), produces **4 distinct-dtype pairs sharing the same scene/cam/temporal segments** but with different aug instances applied with shared seeds - Loss: InfoNCE invariance + SIGReg subspace decorrelation ## Files - `e4.pt` — checkpoint after **epoch 5** (e4 in saving convention, `save_every_freq=2`) - `config.yaml` — training params snapshot ## Loading ```python import torch ckpt = torch.load("e4.pt", map_location="cpu") # Keys ckpt.keys() # → dict_keys(['encoder', 'domain_encoder', 'task_head', 'domain_head', ...]) # encoder & domain_encoder are LoRA-wrapped — to load, apply peft LoRA wrap # BEFORE load_state_dict (see lerobot/policies/smolvlm_act/vjepa_target_encoder.py # in https://github.com/Kim-Eungseo/Spurious-Correlation-Bye-Bye for the # reference loader). ``` ## Validation (epoch 5 end, in-distribution eval) Clean-eval setup (same-content/diff-aug task pair vs. diff-content/same-3-aug dom pair): | Pair | encoder | mean cosine | |---|---|---| | Task pair (same content, diff aug) | task_enc | **0.78** ✓ HIGH | | Task pair (same content, diff aug) | dom_enc | 0.36 LOW ✓ | | Dom pair (diff content, same 3-aug) | task_enc | 0.13 LOW ✓ | | Dom pair (diff content, same 3-aug) | dom_enc | **0.74** ✓ HIGH | Two encoders cleanly separate task vs domain features.