Spaces:
Sleeping
Sleeping
Upload folder using huggingface_hub
Browse files
steerable_retrieval/steer/loading.py
CHANGED
|
@@ -111,7 +111,10 @@ def load_steerable_sae(
|
|
| 111 |
if not hasattr(dec, "b_dec"):
|
| 112 |
dec.b_dec = enc.b_dec
|
| 113 |
|
| 114 |
-
|
|
|
|
|
|
|
|
|
|
| 115 |
sd = state.get("state_dict", state)
|
| 116 |
enc_sd = {k[len("sae_encoder."):]: v for k, v in sd.items() if k.startswith("sae_encoder.")}
|
| 117 |
dec_sd = {k[len("sae_decoder."):]: v for k, v in sd.items() if k.startswith("sae_decoder.")}
|
|
|
|
| 111 |
if not hasattr(dec, "b_dec"):
|
| 112 |
dec.b_dec = enc.b_dec
|
| 113 |
|
| 114 |
+
# Lightning checkpoints from our training runs can contain OmegaConf metadata
|
| 115 |
+
# alongside tensors. PyTorch 2.6 defaults torch.load(weights_only=True), which
|
| 116 |
+
# rejects that metadata; this loader is for trusted project checkpoints.
|
| 117 |
+
state = torch.load(ckpt_path, map_location=device, weights_only=False)
|
| 118 |
sd = state.get("state_dict", state)
|
| 119 |
enc_sd = {k[len("sae_encoder."):]: v for k, v in sd.items() if k.startswith("sae_encoder.")}
|
| 120 |
dec_sd = {k[len("sae_decoder."):]: v for k, v in sd.items() if k.startswith("sae_decoder.")}
|