Spaces:
Sleeping
Sleeping
| """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() | |