villekuosmanen commited on
Commit
b9dff20
·
verified ·
1 Parent(s): 85bee19

Upload SAE model weights, config, and training state

Browse files
Files changed (4) hide show
  1. README.md +96 -0
  2. config.json +36 -0
  3. model.safetensors +3 -0
  4. training_state.pt +3 -0
README.md ADDED
@@ -0,0 +1,96 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ ---
2
+ license: mit
3
+ tags:
4
+ - physical-ai-interpretability-sae
5
+ - LeRobot
6
+ - Robotics
7
+ datasets:
8
+ - villekuosmanen/build_block_tower
9
+ - villekuosmanen/dAgger_build_block_tower_1.0.0
10
+ - villekuosmanen/dAgger_build_block_tower_1.1.0
11
+ - villekuosmanen/dAgger_build_block_tower_1.2.0
12
+ - villekuosmanen/dAgger_build_block_tower_1.3.0
13
+ - villekuosmanen/dAgger_build_block_tower_1.4.0
14
+ - villekuosmanen/fail_build_block_tower_stationary
15
+ - villekuosmanen/fail_build_block_tower_autonomous_interaction
16
+ library_name: physical-ai-interpretability
17
+ ---
18
+
19
+ # Sparse Autoencoder (SAE) Model
20
+
21
+ This model is a Sparse Autoencoder trained for interpretability analysis of robotics policies using the LeRobot framework.
22
+
23
+ ## Model Details
24
+
25
+ - **Architecture**: Multi-modal Sparse Autoencoder
26
+ - **Training Dataset**: `villekuosmanen/build_block_tower`, `villekuosmanen/dAgger_build_block_tower_1.0.0`, `villekuosmanen/dAgger_build_block_tower_1.1.0`, `villekuosmanen/dAgger_build_block_tower_1.2.0`, `villekuosmanen/dAgger_build_block_tower_1.3.0`, `villekuosmanen/dAgger_build_block_tower_1.4.0`, `villekuosmanen/fail_build_block_tower_stationary`, `villekuosmanen/fail_build_block_tower_autonomous_interaction`
27
+ - **Base Policy**: LeRobot ACT policy
28
+ - **Layer Target**: `model.encoder.layers.3.norm2`
29
+ - **Tokens**: 77
30
+ - **Token Dimension**: 128
31
+ - **Feature Dimension**: 12320
32
+ - **Expansion Factor**: 1.25
33
+
34
+ ## Training Configuration
35
+
36
+ - **Learning Rate**: 0.0001
37
+ - **Batch Size**: 16
38
+ - **L1 Penalty**: 0.3
39
+ - **Epochs**: 20
40
+ - **Optimizer**: adam
41
+
42
+ ## Usage
43
+
44
+ ```python
45
+ from physical_ai_interpretability.sae.trainer import load_sae_from_hub
46
+
47
+ # Load model from Hub
48
+ model = load_sae_from_hub("villekuosmanen/build_block_tower_all_small_sae")
49
+
50
+ # Or load using builder
51
+ from physical_ai_interpretability.sae.builder import SAEBuilder
52
+ builder = SAEBuilder(device='cuda')
53
+ model = builder.load_from_hub("villekuosmanen/build_block_tower_all_small_sae")
54
+ ```
55
+
56
+ ## Out-of-Distribution Detection
57
+
58
+ This SAE model can be used for OOD detection with LeRobot policies:
59
+
60
+ ```python
61
+ from physical_ai_interpretability.ood import OODDetector
62
+
63
+ # Create OOD detector with Hub-loaded SAE
64
+ ood_detector = OODDetector(
65
+ policy=your_policy,
66
+ sae_hub_repo_id="villekuosmanen/build_block_tower_all_small_sae"
67
+ )
68
+
69
+ # Fit threshold and use for detection
70
+ ood_detector.fit_ood_threshold_to_validation_dataset(validation_dataset)
71
+ is_ood, error = ood_detector.is_out_of_distribution(observation)
72
+ ```
73
+
74
+ ## Files
75
+
76
+ - `model.safetensors`: The trained SAE model weights
77
+ - `config.json`: Training and model configuration
78
+ - `training_state.pt`: Complete training state (optimizer, scheduler, metrics)
79
+ - `ood_params.json`: OOD detection parameters (if fitted)
80
+
81
+ ## Citation
82
+
83
+ If you use this model in your research, please cite:
84
+
85
+ ```bibtex
86
+ @misc{sae_model,
87
+ title={Sparse Autoencoder for Build Block Tower},
88
+ author={Your Name},
89
+ year={2024},
90
+ url={https://huggingface.co/villekuosmanen/build_block_tower_all_small_sae}
91
+ }
92
+ ```
93
+
94
+ ## Framework
95
+
96
+ This model was trained using the [physical-ai-interpretability](https://github.com/your-repo/physical-ai-interpretability) framework with [LeRobot](https://github.com/huggingface/lerobot).
config.json ADDED
@@ -0,0 +1,36 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ {
2
+ "num_tokens": 77,
3
+ "token_dim": 128,
4
+ "expansion_factor": 1.25,
5
+ "activation_fn": "relu",
6
+ "use_token_sampling": true,
7
+ "fixed_tokens": [
8
+ 0,
9
+ 1
10
+ ],
11
+ "sampling_strategy": "block_average",
12
+ "sampling_stride": 8,
13
+ "max_sampled_tokens": 200,
14
+ "block_size": 8,
15
+ "batch_size": 16,
16
+ "learning_rate": 0.0001,
17
+ "num_epochs": 20,
18
+ "validation_split": 0.1,
19
+ "l1_penalty": 0.3,
20
+ "optimizer": "adam",
21
+ "weight_decay": 1e-05,
22
+ "lr_schedule": "constant",
23
+ "warmup_epochs": 2,
24
+ "gradient_clip_norm": 1.0,
25
+ "early_stopping_patience": 10,
26
+ "early_stopping_min_delta": 1e-05,
27
+ "log_every": 5,
28
+ "save_every": 1000,
29
+ "validate_every": 500,
30
+ "device": "cuda",
31
+ "repo_id": "[villekuosmanen/build_block_tower, villekuosmanen/dAgger_build_block_tower_1.0.0, villekuosmanen/dAgger_build_block_tower_1.1.0, villekuosmanen/dAgger_build_block_tower_1.2.0, villekuosmanen/dAgger_build_block_tower_1.3.0, villekuosmanen/dAgger_build_block_tower_1.4.0, villekuosmanen/fail_build_block_tower_stationary, villekuosmanen/fail_build_block_tower_autonomous_interaction]",
32
+ "repo_hash": "881292d6",
33
+ "layer_name": "model.encoder.layers.3.norm2",
34
+ "activation_cache_path": "/home/ville/.cache/physical_ai_interpretability/sae_activations",
35
+ "experiment_name": "sae_fail_build_block_tow_881292d6"
36
+ }
model.safetensors ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:0211c1fe33e69f654dd269531f5aee5e04b26a56a7552ab4e5a7f14ea2449906
3
+ size 971496408
training_state.pt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:d142098a493441e4ab7adbb34caebce779e5659f678299accf1253b46743bd21
3
+ size 1942998303