hallucination / probe /probing.py
ToiTenBao's picture
Upload hallucination folder
a2ffd07 verified
Raw
History Blame Contribute Delete
9.84 kB
import os
import random
from dataclasses import dataclass
from math import exp, log
from pathlib import Path
from typing import Callable, Literal
import numpy as np
import pandas as pd
import torch
from torch import nn, optim
from jaxtyping import Int, Float
from torch.utils.data import DataLoader, TensorDataset
# Autocheck if the instance is a notebook or not (fixes weird bugs in colab)
from tqdm.autonotebook import tqdm
from .util import DEFAULT_DEVICE
class LinearProbe(nn.Module):
"""
Based on by https://github.com/jbloomAus/alphabetical_probe/blob/main/src/probes.py
"""
def __init__(self, input_dim, num_outputs: int = 1):
super().__init__()
self.fc = nn.Linear(input_dim, num_outputs)
def forward(self, x):
return self.fc(x)
@property
def weights(self):
return self.fc.weight
@property
def biases(self):
return self.fc.bias
def _calc_pos_weights(y: torch.Tensor) -> torch.Tensor:
num_pos_samples = y.sum(dim=0)
num_neg_samples = len(y) - num_pos_samples
return num_neg_samples / num_pos_samples
def train_multi_probe(
x_train: torch.Tensor, # tensor of shape (num_samples, input_dim)
y_train: torch.Tensor, # tensor of shape (num_samples, num_probes), with values in [0, 1]
num_probes: int | None = None, # inferred from y_train if None
batch_size: int = 4096,
num_epochs: int = 100,
lr: float = 0.01,
end_lr: float = 1e-5,
weight_decay: float = 1e-6,
show_progress: bool = True,
optimizer: Literal["Adam", "SGD", "AdamW"] = "Adam",
extra_loss_fn: (
Callable[[LinearProbe, torch.Tensor, torch.Tensor], torch.Tensor] | None
) = None,
verbose: bool = False,
device: torch.device = DEFAULT_DEVICE,
map_acts: Callable[[torch.Tensor], torch.Tensor] | None = None,
probe_dim: int | None = None,
) -> LinearProbe:
"""
Train a multi-class one-vs-rest logistic regression probe on the given data.
This is equivalent to training num_probes separate binary logistic regression probes.
Args:
x_train: tensor of shape (num_samples, input_dim)
y_train: one_hot (or multi-hot) tensor of shape (num_samples, num_probes), with values in [0, 1]
num_probes: number of probes to train simultaneously
batch_size: batch size for training
num_epochs: number of epochs to train for
lr: learning rate
weight_decay: weight decay
show_progress: whether to show a progress bar
device: device to train on
"""
dtype = x_train.dtype
num_probes = num_probes or y_train.shape[-1]
dataset = TensorDataset(x_train, y_train.to(dtype=dtype))
loader = DataLoader(dataset, batch_size=batch_size, shuffle=True)
if probe_dim is None:
probe_dim = x_train.shape[-1]
probe = LinearProbe(probe_dim, num_outputs=num_probes).to(device, dtype=dtype)
_run_probe_training(
probe,
loader,
loss_fn=nn.BCEWithLogitsLoss(pos_weight=_calc_pos_weights(y_train).to(device)),
num_epochs=num_epochs,
lr=lr,
end_lr=end_lr,
weight_decay=weight_decay,
show_progress=show_progress,
optimizer_name=optimizer,
extra_loss_fn=extra_loss_fn,
verbose=verbose,
device=device,
map_acts=map_acts,
)
return probe
def train_binary_probe(
x_train: torch.Tensor, # tensor of shape (num_samples, input_dim)
y_train: torch.Tensor, # tensor of shape (num_samples,), with values in [0, 1]
batch_size: int = 256,
num_epochs: int = 100,
lr: float = 0.01,
end_lr: float = 1e-5,
weight_decay: float = 1e-6,
show_progress: bool = True,
optimizer: Literal["Adam", "SGD", "AdamW"] = "Adam",
extra_loss_fn: (
Callable[[LinearProbe, torch.Tensor, torch.Tensor], torch.Tensor] | None
) = None,
verbose: bool = False,
device: torch.device = DEFAULT_DEVICE,
) -> LinearProbe:
"""
Train a logistic regression probe on the given data. This is a thin wrapped around train_multi_probe.
Args:
x_train: tensor of shape (num_samples, input_dim)
y_train: tensor of shape (num_samples,), with values in [0, 1]
batch_size: batch size for training
num_epochs: number of epochs to train for
lr: learning rate
weight_decay: weight decay
show_progress: whether to show a progress bar
device: device to train on
"""
return train_multi_probe(
x_train,
y_train.unsqueeze(1),
num_probes=1,
batch_size=batch_size,
num_epochs=num_epochs,
lr=lr,
end_lr=end_lr,
weight_decay=weight_decay,
show_progress=show_progress,
optimizer=optimizer,
extra_loss_fn=extra_loss_fn,
verbose=verbose,
device=device,
)
def _get_exponential_decay_scheduler(
optimizer: optim.Optimizer, # type: ignore
start_lr: float,
end_lr: float,
num_steps: int,
) -> optim.lr_scheduler.ExponentialLR:
gamma = exp(log(end_lr / start_lr) / num_steps)
return optim.lr_scheduler.ExponentialLR(optimizer, gamma=gamma)
def _run_probe_training(
probe: LinearProbe,
loader: DataLoader,
loss_fn: Callable[[torch.Tensor, torch.Tensor], torch.Tensor],
num_epochs: int,
lr: float,
end_lr: float,
weight_decay: float,
show_progress: bool,
optimizer_name: Literal["Adam", "SGD", "AdamW"],
extra_loss_fn: (
Callable[[LinearProbe, torch.Tensor, torch.Tensor], torch.Tensor] | None
),
verbose: bool,
device: torch.device,
map_acts: Callable[[torch.Tensor], torch.Tensor] | None = None,
) -> None:
probe.train()
if optimizer_name == "Adam":
optimizer = optim.Adam(probe.parameters(), lr=lr, weight_decay=weight_decay) # type: ignore
elif optimizer_name == "SGD":
optimizer = optim.SGD(probe.parameters(), lr=lr, weight_decay=weight_decay) # type: ignore
elif optimizer_name == "AdamW":
optimizer = optim.AdamW(probe.parameters(), lr=lr, weight_decay=weight_decay) # type: ignore
else:
raise ValueError(f"Unknown optimizer: {optimizer_name}")
scheduler = _get_exponential_decay_scheduler(
optimizer, start_lr=lr, end_lr=end_lr, num_steps=num_epochs
)
epoch_pbar = tqdm(range(num_epochs), disable=not show_progress, desc="Epochs")
for epoch in epoch_pbar:
epoch_sum_loss = 0
batch_pbar = tqdm(
loader,
disable=not show_progress,
leave=False,
desc=f"Epoch {epoch + 1}/{num_epochs}",
)
for batch_embeddings, batch_labels in batch_pbar:
batch_embeddings = batch_embeddings.to(device)
if map_acts is not None:
batch_embeddings = map_acts(batch_embeddings)
batch_labels = batch_labels.to(device)
optimizer.zero_grad()
logits = probe(batch_embeddings)
loss = loss_fn(logits, batch_labels)
if extra_loss_fn is not None:
loss += extra_loss_fn(probe, batch_embeddings, batch_labels)
loss.backward()
optimizer.step()
batch_loss = loss.item()
epoch_sum_loss += batch_loss
batch_pbar.set_postfix({"Loss": f"{batch_loss:.8f}"})
epoch_mean_loss = epoch_sum_loss / len(loader)
current_lr = scheduler.get_last_lr()[0]
epoch_pbar.set_postfix(
{"Mean Loss": f"{epoch_mean_loss:.8f}", "LR": f"{current_lr:.2e}"}
)
if verbose:
print(
f"Epoch {epoch + 1}: Mean Loss: {epoch_mean_loss:.8f}, LR: {current_lr:.2e}"
)
scheduler.step()
probe.eval()
def select_k_features(
l1_probe: LinearProbe,
k: int,
label: int,
acts: torch.Tensor | None = None, # (n_sample, n_feats)
) -> Int[torch.Tensor, "k"]:
if acts is None:
return l1_probe.weights[label].topk(k).indices
else:
return (torch.sum(acts, dim=0) * l1_probe.weights[label].cpu()).topk(k).indices
def load_probe(
save_name: str,
d_model: int,
num_outputs: int,
device: torch.device | str = DEFAULT_DEVICE,
) -> LinearProbe:
probe = LinearProbe(d_model, num_outputs)
probe.load_state_dict(torch.load(save_name, map_location=device))
probe.to(device)
return probe
def train_or_load_lr(
x_train: torch.Tensor, # tensor of shape (num_samples, input_dim)
y_train: torch.Tensor, # tensor of shape (num_samples, num_probes), with values in [0, 1]
save_name: str,
batch_size_lr: int = 4096,
num_epochs: int = 100,
lr: float = 0.01,
end_lr: float = 1e-5,
weight_decay: float = 1e-6,
show_progress: bool = True,
optimizer: Literal["Adam", "SGD", "AdamW"] = "Adam",
verbose: bool = False,
device: torch.device | str = DEFAULT_DEVICE,
map_acts: Callable[[torch.Tensor], torch.Tensor] | None = None,
probe_dim: int | None = None,
) -> LinearProbe:
if not os.path.exists(save_name):
probe = train_multi_probe(
x_train,
y_train,
batch_size=batch_size_lr,
num_epochs=num_epochs,
lr=lr,
end_lr=end_lr,
weight_decay=weight_decay,
show_progress=show_progress,
optimizer=optimizer,
verbose=verbose,
device=device,
map_acts=map_acts,
probe_dim=probe_dim
).to("cpu")
if verbose:
print(f"Probe trained with {probe.num_probes} probes")
torch.save(probe.state_dict(), save_name)
else:
probe = load_probe(
save_name, x_train.shape[-1], y_train.shape[-1], "cpu",
)
return probe