dvarfe's picture
sync with github version
0705c62
Raw
History Blame Contribute Delete
20.9 kB
"""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()