Skip to content

Commit d6ecfac

Browse files
committed
feat(trainer): support param groups
- Updated the `get_parameters` method in `AbstractSparseAutoEncoder` to separate parameters into groups for "others" and "jumprelu" with respective learning rates. - Modified the `_initialize_optimizer` method in `Trainer` to apply different learning rates based on parameter group names and log detailed parameter information. - Added a new configuration field `jumprelu_lr_factor` in `TrainerConfig` to control the learning rate for JumpReLU parameters.
1 parent 26020f2 commit d6ecfac

4 files changed

Lines changed: 32 additions & 6 deletions

File tree

src/lm_saes/abstract_sae.py

Lines changed: 5 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -409,7 +409,11 @@ def init_encoder_with_decoder_transpose(self, factor: float = 1.0):
409409

410410
def get_parameters(self) -> list[dict[str, Any]]:
411411
"""Get the parameters of the model for optimization."""
412-
return [{"params": self.parameters()}]
412+
jumprelu_params = (
413+
list(self.activation_function.parameters()) if isinstance(self.activation_function, JumpReLU) else []
414+
)
415+
other_params = [p for p in self.parameters() if not any(p is param for param in jumprelu_params)]
416+
return [{"params": other_params, "name": "others"}, {"params": jumprelu_params, "name": "jumprelu"}]
413417

414418
def load_full_state_dict(self, state_dict: dict[str, torch.Tensor], device_mesh: DeviceMesh | None = None) -> None:
415419
# Extract and set dataset_average_activation_norm if present

src/lm_saes/config.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -201,6 +201,7 @@ class TrainerConfig(BaseConfig):
201201
lr_end_ratio: float = 1 / 32
202202
lr_warm_up_steps: int | float = 5000
203203
lr_cool_down_steps: int | float = 0.2
204+
jumprelu_lr_factor: float = 1.0
204205
clip_grad_norm: float = 0.0
205206
feature_sampling_window: int = 1000
206207
total_training_tokens: int = 300_000_000

src/lm_saes/sae.py

Lines changed: 0 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -424,6 +424,3 @@ def prepare_label(self, batch: dict[str, torch.Tensor], **kwargs) -> torch.Tenso
424424
else:
425425
label = batch[self.cfg.hook_point_out]
426426
return label
427-
428-
def get_parameters(self) -> list[dict[str, Any]]:
429-
return [{"params": self.parameters()}]

src/lm_saes/trainer.py

Lines changed: 26 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,6 @@
11
import math
22
import os
3-
from typing import Callable, Iterable
3+
from typing import Any, Callable, Iterable
44

55
import torch
66
import torch.optim.lr_scheduler as lr_scheduler
@@ -72,7 +72,31 @@ def calculate_warmup_steps(warmup_steps: float | int) -> int:
7272
@timer.time("initialize_optimizer")
7373
def _initialize_optimizer(self, sae: AbstractSparseAutoEncoder):
7474
assert isinstance(self.cfg.lr, float)
75-
optimizer = Adam(sae.get_parameters(), lr=self.cfg.lr, betas=self.cfg.betas)
75+
76+
def _apply_lr(parameters: dict[str, Any]):
77+
assert isinstance(self.cfg.lr, float)
78+
if parameters["name"] == "jumprelu":
79+
return {**parameters, "lr": self.cfg.jumprelu_lr_factor * self.cfg.lr}
80+
return parameters
81+
82+
params = [_apply_lr(parameters) for parameters in sae.get_parameters()]
83+
84+
def _format_parameters(parameters: dict[str, Any]) -> str:
85+
param_info = f"{parameters['name']}:"
86+
for i, param in enumerate(parameters["params"]):
87+
param_info += f"\n [{i}] shape={list(param.shape)}, dtype={param.dtype}"
88+
if param.requires_grad:
89+
param_info += ", trainable"
90+
else:
91+
param_info += ", frozen"
92+
if "lr" in parameters:
93+
param_info += f"\n lr={parameters['lr']}"
94+
return param_info
95+
96+
param_str = "\n".join([_format_parameters(p) for p in params])
97+
logger.info(f"\nParameter Groups: \n{param_str}\n")
98+
99+
optimizer = Adam(params, lr=self.cfg.lr, betas=self.cfg.betas)
76100
scheduler = get_scheduler(
77101
scheduler_name=self.cfg.lr_scheduler_name,
78102
optimizer=optimizer,

0 commit comments

Comments
 (0)