Skip to content

Commit 3c8ea60

Browse files
committed
ruff fmt
1 parent 947ad78 commit 3c8ea60

19 files changed

Lines changed: 92 additions & 131 deletions

contrib/birdsong/notebooks/clips.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -92,8 +92,7 @@ def add_target(obs: pl.DataFrame, fields: list[str]) -> pl.DataFrame:
9292
obs = obs.with_columns([pl.col(field).fill_null("unknown") for field in fields])
9393

9494
combos = (
95-
obs
96-
.select(fields)
95+
obs.select(fields)
9796
.unique(maintain_order=True) # first-seen ordering
9897
.with_columns(pl.arange(0, pl.len(), dtype=pl.Int32).alias("target"))
9998
)
@@ -102,8 +101,7 @@ def add_target(obs: pl.DataFrame, fields: list[str]) -> pl.DataFrame:
102101

103102
target2fields = {
104103
target: tuple(rest)
105-
for target, *rest in obs
106-
.unique(pl.col("target"))
104+
for target, *rest in obs.unique(pl.col("target"))
107105
.select("target", *fields)
108106
.iter_rows()
109107
}

contrib/trait_discovery/notebooks/001_actfn.py

Lines changed: 6 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -261,8 +261,7 @@ def _finalize_df(rows: list[dict[str, object]]):
261261
)
262262

263263
df = df.with_columns(
264-
pl
265-
.when(pl.col("config/sae/activation").struct.field("top_k").is_not_null())
264+
pl.when(pl.col("config/sae/activation").struct.field("top_k").is_not_null())
266265
.then(pl.lit("topk"))
267266
.otherwise(pl.lit("relu"))
268267
.alias("config/sae/activation_kind"),
@@ -351,8 +350,7 @@ def _(df, np, pl, plt):
351350
def _():
352351
fig, ax = plt.subplots(figsize=(4.5, 3), dpi=300, layout="constrained")
353352
ks, ys, ids = (
354-
df
355-
.filter(pl.col("config/sae/activation_kind") == "topk")
353+
df.filter(pl.col("config/sae/activation_kind") == "topk")
356354
.group_by(pl.col("config/sae/activation").struct.field("top_k"))
357355
.agg(pl.col("summary/eval/l0"), pl.col("id"))
358356
.sort(by="top_k")
@@ -707,8 +705,7 @@ def _(df: pl.DataFrame):
707705
(pl.col("config/sae/activation_kind") == "relu")
708706
& (pl.col("config/val_data/layer") == layer)
709707
& (
710-
pl
711-
.col("config/sae/activation")
708+
pl.col("config/sae/activation")
712709
.struct.field("sparsity")
713710
.struct.field("coeff")
714711
== lam
@@ -806,8 +803,7 @@ def _(df: pl.DataFrame):
806803
(pl.col("config/sae/activation_kind") == "relu")
807804
& (pl.col("config/val_data/layer") == layer)
808805
& (
809-
pl
810-
.col("config/sae/activation")
806+
pl.col("config/sae/activation")
811807
.struct.field("sparsity")
812808
.struct.field("coeff")
813809
== lam
@@ -1002,8 +998,7 @@ def _(df):
1002998
), # .select('best_train_probe_r').unique(),
1003999
)
10041000
group = (
1005-
df
1006-
.filter(
1001+
df.filter(
10071002
pl.col(col).is_not_null()
10081003
& pl.col("is_pareto")
10091004
& (pl.col("data_key") == data_key)
@@ -1057,8 +1052,7 @@ def _(df):
10571052

10581053
for kind in ["relu", "topk"]:
10591054
group = (
1060-
df
1061-
.filter(
1055+
df.filter(
10621056
(pl.col("config/sae/activation_kind") == kind)
10631057
& pl.col(col).is_not_null()
10641058
# & pl.col("is_pareto")

contrib/trait_discovery/notebooks/002_optim.py

Lines changed: 4 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -261,8 +261,7 @@ def _finalize_df(rows: list[dict[str, object]]):
261261
)
262262

263263
df = df.with_columns(
264-
pl
265-
.when(pl.col("config/sae/activation").struct.field("top_k").is_not_null())
264+
pl.when(pl.col("config/sae/activation").struct.field("top_k").is_not_null())
266265
.then(pl.lit("topk"))
267266
.otherwise(pl.lit("relu"))
268267
.alias("config/sae/activation_kind"),
@@ -699,8 +698,7 @@ def _(df):
699698

700699
for optim in ["adam", "muon"]:
701700
group = (
702-
df
703-
.filter(
701+
df.filter(
704702
(pl.col("config/optim") == optim)
705703
# & pl.col(col).is_not_null()
706704
# & pl.col("is_pareto")
@@ -712,12 +710,10 @@ def _(df):
712710
)
713711
.agg(
714712
pl.len().alias("n_trials"),
715-
pl
716-
.col("summary/metrics/dead_unit_pct")
713+
pl.col("summary/metrics/dead_unit_pct")
717714
.mean()
718715
.alias("train_mean_pct"),
719-
pl
720-
.col("summary/metrics/dead_unit_pct")
716+
pl.col("summary/metrics/dead_unit_pct")
721717
.std()
722718
.alias("train_std_pct"),
723719
(pl.col("summary/eval/n_dead") / (1024 * 16) * 100)

contrib/trait_discovery/notebooks/003_auxk.py

Lines changed: 4 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -270,8 +270,7 @@ def _finalize_df(rows: list[dict[str, object]]):
270270
)
271271

272272
df = (
273-
df
274-
.unnest("config/sae", "config/train_data/metadata", separator="/")
273+
df.unnest("config/sae", "config/train_data/metadata", separator="/")
275274
.unnest("config/sae/activation", separator="/")
276275
.unnest(
277276
"config/sae/activation/aux",
@@ -607,8 +606,7 @@ def _(df):
607606
# col = "summary/eval/n_almost_dead"
608607

609608
group = (
610-
df
611-
.filter(
609+
df.filter(
612610
True
613611
# & pl.col(col).is_not_null()
614612
& pl.col("is_pareto")
@@ -654,8 +652,7 @@ def _(df):
654652
def _(df, mo, pl):
655653
def _(df):
656654
group = (
657-
df
658-
.filter(
655+
df.filter(
659656
pl.col("downstream/train/probe_r").is_not_null() & pl.col("is_pareto")
660657
# & pl.col("config/val_data/layer") == 21
661658
)
@@ -699,7 +696,6 @@ def _(df):
699696

700697
@app.cell
701698
def _(df, pl):
702-
703699
df.filter(
704700
pl.col("downstream/train/probe_r").is_not_null() & pl.col("is_pareto")
705701
).group_by(
@@ -817,8 +813,7 @@ def _(df: pl.DataFrame):
817813
def _(df, mo, pl):
818814
def _(df):
819815
group = (
820-
df
821-
.filter(pl.col("downstream/train/probe_r").is_not_null())
816+
df.filter(pl.col("downstream/train/probe_r").is_not_null())
822817
.select(
823818
"id",
824819
"data_key",

contrib/trait_discovery/notebooks/004_fishbase.py

Lines changed: 6 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -162,8 +162,7 @@ def _finalize_df(rows: list[dict[str, object]]):
162162
)
163163

164164
df = (
165-
df
166-
.unnest("config/sae", "config/train_data/metadata", separator="/")
165+
df.unnest("config/sae", "config/train_data/metadata", separator="/")
167166
.unnest("config/sae/activation", separator="/")
168167
.unnest(
169168
"config/sae/activation/aux",
@@ -533,8 +532,7 @@ def _(pl):
533532
MigrationEnum = pl.Enum(migration_cols)
534533

535534
fishbase_df = (
536-
pl
537-
.read_csv(
535+
pl.read_csv(
538536
"contrib/trait_discovery/data/fishvista_fishbase.csv",
539537
null_values=["?"],
540538
# Order matters
@@ -575,17 +573,15 @@ def _(pl):
575573
}),
576574
)
577575
.with_columns(
578-
pl
579-
.coalesce([
576+
pl.coalesce([
580577
pl.when(pl.col(col) == 1.0).then(pl.lit(col)) for col in habitat_cols
581578
])
582579
.cast(HabitatEnum)
583580
.alias("habitat")
584581
)
585582
.drop(habitat_cols)
586583
.with_columns(
587-
pl
588-
.coalesce([
584+
pl.coalesce([
589585
pl.when(pl.col(col) == 1.0).then(pl.lit(col)) for col in migration_cols
590586
])
591587
.cast(MigrationEnum)
@@ -666,15 +662,13 @@ def _():
666662
)
667663
print(species_df.get_column("habitat").dtype.categories.to_list())
668664
habitats = (
669-
species_df
670-
.get_column("habitat")
665+
species_df.get_column("habitat")
671666
.to_physical()
672667
.to_numpy()
673668
.repeat(md.content_tokens_per_example)
674669
)
675670
migration = (
676-
species_df
677-
.get_column("migration")
671+
species_df.get_column("migration")
678672
.to_physical()
679673
.to_numpy()
680674
.repeat(md.content_tokens_per_example)

contrib/trait_discovery/notebooks/004_fishbase_cls.py

Lines changed: 2 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -163,8 +163,7 @@ def _finalize_df(rows: list[dict[str, object]]):
163163
)
164164

165165
df = (
166-
df
167-
.unnest("config/sae", "config/train_data/metadata", separator="/")
166+
df.unnest("config/sae", "config/train_data/metadata", separator="/")
168167
.unnest("config/sae/activation", separator="/")
169168
.unnest(
170169
"config/sae/activation/aux",
@@ -611,8 +610,7 @@ def _(df: pl.DataFrame):
611610
& (pl.col(x_col).is_not_null())
612611
)
613612
group = group.sort(by=x_col).with_columns(
614-
pl
615-
.col("cls/classifier")
613+
pl.col("cls/classifier")
616614
.map_elements(
617615
lambda clf: len(np.nonzero(clf.coef_)[0]), return_dtype=pl.Int32
618616
)

contrib/trait_discovery/notebooks/005_bufferflies.py

Lines changed: 3 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -259,8 +259,7 @@ def _finalize_sae_df(rows: list[dict[str, object]]):
259259
)
260260

261261
df = (
262-
df
263-
.unnest("config/sae", "config/train_data/metadata", separator="/")
262+
df.unnest("config/sae", "config/train_data/metadata", separator="/")
264263
.unnest("config/sae/activation", separator="/")
265264
.unnest(
266265
"config/sae/activation/aux",
@@ -298,8 +297,7 @@ def _finalize_sae_df(rows: list[dict[str, object]]):
298297
def _finalize_clf_df(rows: list[dict[str, object]]):
299298
df = pl.DataFrame(rows, infer_schema_length=None)
300299
df = (
301-
df
302-
.unnest("config/sae", "config/train_data/metadata", separator="/")
300+
df.unnest("config/sae", "config/train_data/metadata", separator="/")
303301
.unnest("config/sae/activation", separator="/")
304302
.unnest(
305303
"config/sae/activation/aux",
@@ -518,8 +516,7 @@ def _(df: pl.DataFrame):
518516
continue
519517

520518
group = group.with_columns(
521-
pl
522-
.col("cls/classifier")
519+
pl.col("cls/classifier")
523520
.map_elements(
524521
lambda clf: len(np.nonzero(clf.coef_)[0]),
525522
return_dtype=pl.Int32,

contrib/trait_discovery/notebooks/006_proposal_audit.py

Lines changed: 8 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -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

Comments
 (0)