Pliploop commited on
Commit
b0af84d
·
verified ·
1 Parent(s): 6876324

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
- state = torch.load(ckpt_path, map_location=device)
 
 
 
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.")}