@@ -28,6 +28,7 @@ def _():
2828
2929 import saev .colors
3030 import saev .data .datasets
31+
3132 return (
3233 Float ,
3334 base64 ,
@@ -177,8 +178,7 @@ def _finalize_sae_df(rows: list[dict[str, object]]):
177178 )
178179
179180 df = (
180- df
181- .unnest ("config/sae" , "config/train_data/metadata" , separator = "/" )
181+ df .unnest ("config/sae" , "config/train_data/metadata" , separator = "/" )
182182 .unnest ("config/sae/activation" , separator = "/" )
183183 .unnest (
184184 "config/sae/activation/aux" ,
@@ -216,8 +216,7 @@ def _finalize_sae_df(rows: list[dict[str, object]]):
216216 def _finalize_clf_df (rows : list [dict [str , object ]]):
217217 df = pl .DataFrame (rows , infer_schema_length = None )
218218 df = (
219- df
220- .unnest ("config/sae" , "config/train_data/metadata" , separator = "/" )
219+ df .unnest ("config/sae" , "config/train_data/metadata" , separator = "/" )
221220 .unnest ("config/sae/activation" , separator = "/" )
222221 .unnest (
223222 "config/sae/activation/aux" ,
@@ -334,8 +333,7 @@ def sigmoid(x, L, k, x0, b):
334333
335334 # Table of individual points
336335 table = (
337- filtered
338- .select (
336+ filtered .select (
339337 "config/val_data/layer" ,
340338 "cls/cfg/cls/key" ,
341339 k_col ,
@@ -479,8 +477,7 @@ def sigmoid(x, L, k, x0, b):
479477
480478 # Table of individual points
481479 table = (
482- filtered
483- .select (
480+ filtered .select (
484481 "config/val_data/layer" ,
485482 "cls/cfg/cls/key" ,
486483 k_col ,
@@ -636,6 +633,7 @@ def load_mean_values(run) -> Float[np.ndarray, " d_sae"]:
636633 raise RuntimeError (f"Wandb sucks: { err } " ) from err
637634
638635 raise ValueError (f"mean_values not found in run '{ run .id } '" )
636+
639637 return load_freqs , load_mean_values
640638
641639
@@ -701,6 +699,7 @@ def get_data_key(metadata: dict[str, object]) -> str | None:
701699
702700 print (f"Unknown data: { data_cfg } " )
703701 return None
702+
704703 return get_data_key , get_model_key
705704
706705
@@ -809,6 +808,7 @@ def get_cls_results(run: saev.disk.Run) -> list[dict[str, object]]:
809808 continue
810809
811810 return results
811+
812812 return (get_cls_results ,)
813813
814814
0 commit comments