"""Training loop utilities and Trainer class.""" from __future__ import annotations import os import random from typing import Any, Callable, Dict, Optional, Tuple, Union import numpy as np import torch from accelerate import Accelerator from accelerate.utils.tqdm import tqdm # type:ignore from torch.utils.data import DataLoader from training.constants import HYPERPARAMS_LOG_FILENAME, SAE_CONFIG_FILENAME from training.losses.base import SaeForwardOutput from training.models.sae import get_wsd_scheduler from log_config import get_logger logger = get_logger(__name__) __all__ = [ "SAE_CONFIG_FILENAME", "HYPERPARAMS_LOG_FILENAME", "Trainer", "format_hyperparams_text", "write_hyperparams_log", "fix_seed", "get_latest_checkpoint_dir", "save_checkpoint", "setup_optimizer_and_scheduler", "setup_lambda_scaling", "compute_lambda", "extract_activations", "train_step", "train_epoch", "calc_loss_reconstr", "_extract_checkpoint_step", ] def format_hyperparams_text(fields: Dict[str, Any]) -> str: sep = "---" lines = [sep] for name, value in fields.items(): lines.append(f"{name} : {value}") lines.append(sep) return "\n".join(lines) def write_hyperparams_log( log_file_path: str, fields: Dict[str, Any], current_metrics: Optional[Dict[str, float]] = None, ) -> None: payload = dict(fields) if current_metrics is not None: for name, value in current_metrics.items(): payload[name] = float(value) log_dir = os.path.dirname(log_file_path) if log_dir: os.makedirs(log_dir, exist_ok=True) with open(log_file_path, "w", encoding="utf-8") as f: f.write(format_hyperparams_text(payload) + "\n") def fix_seed(seed: int) -> None: random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) torch.backends.cudnn.deterministic = True torch.backends.cudnn.benchmark = False def _extract_checkpoint_step(name: str) -> Optional[int]: prefix = "checkpoint-" if not name.startswith(prefix): return None try: return int(name[len(prefix) :]) except ValueError: return None def get_latest_checkpoint_dir(output_dir: str) -> Optional[str]: if output_dir is None or not os.path.isdir(output_dir): return None candidates = [] for entry in os.listdir(output_dir): step = _extract_checkpoint_step(entry) if step is None: continue candidates.append((step, os.path.join(output_dir, entry))) if not candidates: return None candidates.sort(key=lambda item: item[0]) return candidates[-1][1] def save_checkpoint( accelerator: Accelerator, autoenc_model: torch.nn.Module, optimizer: torch.optim.Optimizer, epoch: int, args: Any, global_step: int, ) -> None: save_path = os.path.join(args.output_dir, f"checkpoint-{global_step}") if accelerator.is_main_process: logger.info('Saving to %s', save_path) accelerator.save_state(save_path) def create_optimizer( model: torch.nn.Module, lr: float, weight_decay: float, betas: Tuple[float, float] = (0.9, 0.999), ) -> torch.optim.Optimizer: return torch.optim.Adam(model.parameters(), lr=lr, weight_decay=weight_decay, betas=betas) def create_step_scheduler( optimizer: torch.optim.Optimizer, step_size: int = 6, gamma: float = 0.5, ) -> torch.optim.lr_scheduler.StepLR: return torch.optim.lr_scheduler.StepLR(optimizer, step_size=step_size, gamma=gamma) def create_wsd_scheduler( optimizer: torch.optim.Optimizer, n_total_steps: int, wsd_warmup_steps: int = 100, wsd_cooldown_frac: float = 0.2, wsd_end_lr_factor: float = 0.1, ) -> torch.optim.lr_scheduler.LambdaLR: return get_wsd_scheduler( optimizer, n_steps=n_total_steps, n_warmup_steps=wsd_warmup_steps, percent_cooldown=wsd_cooldown_frac, end_lr_factor=wsd_end_lr_factor, ) _SCHEDULER_FACTORIES: Dict[str, Callable] = { "step": create_step_scheduler, "wsd": create_wsd_scheduler, } _SCHEDULER_OPTIMIZER_BETAS: Dict[str, Tuple[float, float]] = { "wsd": (0.5, 0.9375), "step": (0.9, 0.999), } def setup_optimizer_and_scheduler( model: torch.nn.Module, lr: float, weight_decay: float, scheduler_type: str = "step", scheduler_kwargs: Optional[Dict[str, Union[int, float]]] = None, ) -> Tuple[torch.optim.Optimizer, torch.optim.lr_scheduler.LambdaLR]: betas = _SCHEDULER_OPTIMIZER_BETAS.get(scheduler_type, (0.9, 0.999)) optimizer = create_optimizer(model, lr, weight_decay, betas=betas) factory = _SCHEDULER_FACTORIES[scheduler_type] scheduler = factory(optimizer, **(scheduler_kwargs or {})) return optimizer, scheduler def setup_lambda_scaling(dataloader_train: DataLoader, n_epochs: int) -> float: end_lambda_lin_scaling = 0.05 * len(dataloader_train) * n_epochs logger.info('linear increase of lambda for %s steps', end_lambda_lin_scaling) return end_lambda_lin_scaling def compute_lambda( epoch: int, batch_num: int, lambda_param: float, use_lambda_scaling: bool, end_lambda_lin_scaling: Optional[float], ) -> Tuple[float, int]: if not use_lambda_scaling: return lambda_param, batch_num if epoch == 0: return 0.0, batch_num cur_lambda = lambda_param * min(batch_num / end_lambda_lin_scaling, 1.0) return cur_lambda, batch_num + 1 def calc_loss_reconstr(x: torch.Tensor, x_rec: torch.Tensor) -> torch.Tensor: return torch.nn.functional.mse_loss(x_rec, x) def _set_loss_lambda(loss_module: Any, value: float) -> None: if hasattr(loss_module, "lambda_param"): loss_module.lambda_param = value def extract_activations( model: torch.nn.Module, activations: dict, layer_name: str, data_batch: dict, device: torch.device, weight_dtype: torch.dtype, scaling_factor: float, ) -> torch.Tensor: with torch.no_grad(): images = data_batch["images"] _ = model(images) data = activations[layer_name] data = data.to(device, dtype=weight_dtype) data *= scaling_factor return data def train_step( autoenc_model: torch.nn.Module, accelerator: Accelerator, data: torch.Tensor, optimizer: torch.optim.Optimizer, loss_fn: Callable, grad_clip_val: float, grad_accum_steps: int = 1, sample_frac: float = 0.5, ) -> float: real_batch_size = data.shape[0] optimizer.zero_grad() perm = torch.randperm(real_batch_size, device=data.device) data = data[perm[: int(len(perm) * sample_frac)]] chunk_size = max(1, real_batch_size // grad_accum_steps) chunks = data.split(chunk_size, dim=0) last_loss = 0.0 for chunk in chunks: chunk_weight = chunk.shape[0] / real_batch_size acts, x_rec, weights = autoenc_model.forward(chunk) out = SaeForwardOutput(acts=acts, x_rec=x_rec, decoder_weight=weights) loss = loss_fn(chunk, out) last_loss = float(loss.item()) scaled_loss = loss * chunk_weight accelerator.backward(scaled_loss) del acts, x_rec, out, loss, scaled_loss torch.nn.utils.clip_grad_norm_(autoenc_model.parameters(), grad_clip_val) optimizer.step() return last_loss def train_epoch( autoenc_model: torch.nn.Module, model: torch.nn.Module, activations: dict, layer_name: str, dataloader_train: DataLoader, accelerator: Accelerator, optimizer: torch.optim.Optimizer, scheduler: Any, loss_fn: Callable, args: Any, device: torch.device, weight_dtype: torch.dtype, scaling_factor: float, lambda_param: float, end_lambda_lin_scaling: Optional[float], grad_clip_val: float, sample_frac: float, epoch: int, global_step: int, batch_num: int, start_iter: int = 0, per_batch_scheduler: bool = False, ) -> Tuple[int, int]: autoenc_model.train() with tqdm(enumerate(dataloader_train), unit="batch", disable=not args.use_tqdm) as tepoch: for idx, data_batch in tepoch: if idx < start_iter: continue with accelerator.accumulate(autoenc_model): tepoch.set_description(f"Epoch {epoch}") data = extract_activations( model, activations, layer_name, data_batch, device, weight_dtype, scaling_factor, ) cur_lambda, batch_num = compute_lambda( epoch, batch_num, lambda_param, args.use_lambda_scaling, end_lambda_lin_scaling, ) _set_loss_lambda(loss_fn, cur_lambda) batch_loss = train_step( autoenc_model, accelerator, data, optimizer, loss_fn, grad_clip_val, args.grad_accum_steps, sample_frac, ) if accelerator.is_main_process: tepoch.set_postfix(loss=batch_loss) if (global_step + 1) % args.save_steps == 0: save_checkpoint(accelerator, autoenc_model, optimizer, epoch, args, global_step) if per_batch_scheduler: scheduler.step() global_step += 1 if not per_batch_scheduler: scheduler.step() return global_step, batch_num class Trainer: """High-level SAE training orchestrator.""" def __init__(self, backend, args) -> None: self.backend = backend self.args = args def save_sae_config(self, sae_input_dim: int, inner_dim: int) -> None: import json cfg = self.backend.sae_config_fields(self.args, sae_input_dim, inner_dim) os.makedirs(self.args.output_dir, exist_ok=True) path = os.path.join(self.args.output_dir, SAE_CONFIG_FILENAME) with open(path, "w", encoding="utf-8") as f: json.dump(cfg, f, indent=2) logger.info('SAE config saved to %s', path) def build_hyperparams_payload( self, sae_input_dim: int, inner_dim: int, start_epoch: int, start_iter: int, resume_checkpoint_dir: Optional[str], ) -> Dict[str, Any]: args = self.args hyperparams: Dict[str, Any] = { "resume_checkpoint": resume_checkpoint_dir if resume_checkpoint_dir else "none (training from scratch)", "start_epoch": start_epoch, "start_iter": start_iter, "experiment_num": args.experiment_num, "output_dir": args.output_dir, "dataset_name": args.dataset_name, "iqa_model": args.iqa_model, "n_epochs": args.n_epochs, "batch_size": args.batch_size, "learning_rate": args.learning_rate, "weight_decay": args.weight_decay, "grad_clip_val": args.grad_clip_val, "mixed_precision": args.mixed_precision, "sae_input_dim": sae_input_dim, "expansion_factor": args.expansion_factor if args.expansion_factor is not None else "not set", "inner_dim": inner_dim, "weight_norm_init": args.weight_norm_init, "scaling_factor": args.scaling_factor, "lambda_param": args.lambda_param, "use_lambda_scaling": args.use_lambda_scaling, "non_dist_prob": args.non_dist_prob, "train_frac": args.train_frac, "save_steps": args.save_steps, "grad_accum_steps": args.grad_accum_steps, "downscale_factor": args.downscale_factor, "crop_size": args.crop_size, "sample_frac": args.sample_frac, "sae_type": args.sae_type, "scheduler_type": args.scheduler_type, } hyperparams.update(self.backend.hyperparams_fields(args)) if args.sae_type == "mp_sae": hyperparams["mp_threshold"] = args.mp_threshold hyperparams["mp_normalize"] = bool(args.mp_normalize) hyperparams["aux_alpha"] = args.aux_alpha if args.scheduler_type == "wsd": hyperparams["wsd_warmup_steps"] = args.wsd_warmup_steps hyperparams["wsd_cooldown_frac"] = args.wsd_cooldown_frac hyperparams["wsd_end_lr_factor"] = args.wsd_end_lr_factor return hyperparams def run(self) -> None: from pathlib import Path from accelerate import Accelerator from accelerate.utils import ProjectConfiguration from training.data import TrainingDataConfig, create_training_dataloaders from training.losses.factory import create_loss from training.models.factory import create_sae fix_seed(42) args = self.args backend = self.backend hyperparams_log_path = os.path.join( args.output_dir, args.logging_dir, HYPERPARAMS_LOG_FILENAME ) resume_checkpoint_dir = get_latest_checkpoint_dir(args.output_dir) resume_global_step = 0 if resume_checkpoint_dir is not None: extracted_step = _extract_checkpoint_step(os.path.basename(resume_checkpoint_dir)) if extracted_step is not None: resume_global_step = extracted_step logging_dir = Path(args.output_dir, args.logging_dir) accelerator_project_config = ProjectConfiguration( project_dir=args.output_dir, logging_dir=logging_dir ) wandb_config = backend.tracker_config(args, layer_name="pending") wandb_config.update( { "sae_type": args.sae_type, "batch_size": args.batch_size, "inner_dim": args.inner_dim, "weight_norm_init": args.weight_norm_init, "lr": args.learning_rate, "n_epochs": args.n_epochs, "lambda_param": args.lambda_param, "scaling_factor": args.scaling_factor, "use_lambda_scaling": args.use_lambda_scaling, "grad_clip_val": args.grad_clip_val, "non_dist_prob": float(args.non_dist_prob), "grad_accum_steps": args.grad_accum_steps, "downscale_factor": args.downscale_factor, "crop_size": args.crop_size, "scheduler_type": args.scheduler_type, "mp_threshold": args.mp_threshold if args.sae_type == "mp_sae" else None, "aux_alpha": args.aux_alpha, } ) accelerator = Accelerator( mixed_precision=args.mixed_precision, project_config=accelerator_project_config, log_with="tensorboard", ) if args.job_name is not None: accelerator.init_trackers( project_name=os.path.basename(args.output_dir), config=wandb_config, init_kwargs={"tensorboard": {"name": args.job_name}}, ) else: accelerator.init_trackers( project_name=os.path.basename(args.output_dir), config=wandb_config, ) device = accelerator.device if accelerator.is_main_process: os.makedirs(args.output_dir, exist_ok=True) weight_dtype = torch.float32 if accelerator.mixed_precision == "fp16": weight_dtype = torch.float16 elif accelerator.mixed_precision == "bf16": weight_dtype = torch.bfloat16 iqa_weight_dtype = torch.float32 if args.iqa_model_weight_type == "fp16": iqa_weight_dtype = torch.float16 elif args.iqa_model_weight_type == "bf16": iqa_weight_dtype = torch.bfloat16 model = backend.create_model(device, iqa_weight_dtype, args) sae_input_dim = backend.infer_input_dim(model, device, iqa_weight_dtype, args) inner_dim = args.inner_dim if args.expansion_factor is not None: inner_dim = sae_input_dim * args.expansion_factor wandb_config["inner_dim"] = inner_dim if accelerator.is_main_process: self.save_sae_config(sae_input_dim, inner_dim) data_config = TrainingDataConfig( split="train", prestine_prob=float(args.non_dist_prob), preprocess=backend.training_preprocessing(args), train_frac=args.train_frac, batch_size=args.batch_size, num_workers=args.num_workers, token=args.hf_token, ) dataloader_train, _ = create_training_dataloaders( args.dataset_name, data_config, ) autoenc_model = create_sae( sae_input_dim, inner_dim, device, weight_dtype, sae_type=args.sae_type, weight_norm_init=args.weight_norm_init, mp_threshold=args.mp_threshold, mp_normalize=bool(args.mp_normalize), ) loss_fn = create_loss( args.sae_type, lambda_param=args.lambda_param, aux_alpha=args.aux_alpha, ) end_lambda_lin_scaling = None if args.use_lambda_scaling: end_lambda_lin_scaling = setup_lambda_scaling(dataloader_train, args.n_epochs) n_total_steps = len(dataloader_train) * args.n_epochs if args.scheduler_type == "wsd": scheduler_kwargs: Dict[str, Union[int, float]] = { "n_total_steps": n_total_steps, "wsd_warmup_steps": args.wsd_warmup_steps, "wsd_cooldown_frac": args.wsd_cooldown_frac, "wsd_end_lr_factor": args.wsd_end_lr_factor, } else: scheduler_kwargs = {"step_size": 6, "gamma": 0.5} optimizer, scheduler = setup_optimizer_and_scheduler( autoenc_model, args.learning_rate, args.weight_decay, scheduler_type=args.scheduler_type, scheduler_kwargs=scheduler_kwargs, ) global_step = resume_global_step start_epoch = resume_global_step // max(len(dataloader_train), 1) start_iter = resume_global_step % max(len(dataloader_train), 1) batch_num = resume_global_step autoenc_model, model, optimizer, dataloader_train = accelerator.prepare( autoenc_model, model, optimizer, dataloader_train ) if resume_checkpoint_dir is not None: accelerator.load_state(resume_checkpoint_dir) if accelerator.is_main_process: logger.info( 'Resumed from %s (epoch %s, step %s)', resume_checkpoint_dir, start_epoch, global_step, ) hook_ctx = backend.register_hook(accelerator.unwrap_model(model), args) layer_name = hook_ctx.layer_name activations = hook_ctx.activations base_hyperparams = self.build_hyperparams_payload( sae_input_dim, inner_dim, start_epoch, start_iter, resume_checkpoint_dir ) if accelerator.is_main_process: logger.info('Everything prepared') logger.info('%s', format_hyperparams_text(base_hyperparams)) write_hyperparams_log(hyperparams_log_path, base_hyperparams) last_epoch = start_epoch - 1 for epoch in range(start_epoch, args.n_epochs): cur_start_iter = start_iter if epoch == start_epoch else 0 global_step, batch_num = train_epoch( autoenc_model, model, activations, layer_name, dataloader_train, accelerator, optimizer, scheduler, loss_fn, args, device, weight_dtype, args.scaling_factor, args.lambda_param, end_lambda_lin_scaling, args.grad_clip_val, args.sample_frac, epoch, global_step, batch_num, cur_start_iter, per_batch_scheduler=(args.scheduler_type == "wsd"), ) last_epoch = epoch final_epoch = last_epoch if last_epoch >= 0 else 0 save_checkpoint(accelerator, autoenc_model, optimizer, final_epoch, args, global_step) accelerator.end_training()