liar-detector / examples /train.py
zeechimp's picture
Create examples/train.py
3bc5798 verified
Raw History Blame Contribute Delete
5.36 kB
"""Train LiarDetectorForLieDetection and save a Hub-ready folder."""
from __future__ import annotations
import argparse, os, time
import numpy as np
import torch
from torch.utils.data import DataLoader, TensorDataset
from liar_detector import (
LiarDetectorConfig,
LiarDetectorForLieDetection,
LiarDetectorFeatureExtractor,
extract_features,
synth_true, apply_lie,
ID_FAMILIES, OOD_FAMILIES, T_LEN,
)
from liar_detector.signals import FAMILY_INDEX
SCENARIO_WEIGHTS = (("A_lies", 0.33), ("B_lies", 0.33),
("both_honest", 0.17), ("both_lie", 0.17))
def pick_scenario(rng):
r = rng.random(); acc = 0.0
for name, w in SCENARIO_WEIGHTS:
acc += w
if r < acc:
return name
return "A_lies"
def make_pair(rng, families, scenario):
x = synth_true(rng)
sx = float(np.std(x)) + 1e-9
noise_std = rng.uniform(0.05, 0.15) * sx
def honest():
return x + noise_std * rng.standard_normal(T_LEN)
def lied():
fam = families[int(rng.integers(len(families)))]
return apply_lie(x, rng, fam) + noise_std * rng.standard_normal(T_LEN), fam
if scenario == "A_lies":
A, fam = lied(); B = honest()
return A, B, 0, FAMILY_INDEX.get(fam, 0), 1, fam
if scenario == "B_lies":
A = honest(); B, fam = lied()
return A, B, 1, FAMILY_INDEX.get(fam, 0), 1, fam
if scenario == "both_honest":
return honest(), honest(), 0, 0, 0, "none"
if scenario == "both_lie":
A, fA = lied(); B, fB = lied()
return A, B, 0, 0, 0, f"{fA}+{fB}"
raise ValueError(scenario)
def build_arrays(n, seed, families):
rng = np.random.default_rng(seed)
X = np.zeros((n, 44), dtype=np.float32)
by = np.zeros(n, dtype=np.int64)
fy = np.zeros(n, dtype=np.int64)
py = np.zeros(n, dtype=np.int64)
fams = []
for i in range(n):
s = pick_scenario(rng)
A, B, b, f, p, fam = make_pair(rng, families, s)
X[i] = extract_features(A, B)
by[i], fy[i], py[i] = b, f, p
fams.append(fam)
return X, by, fy, py, fams
def fit_temperature(logits, y, grid=None):
if grid is None:
grid = np.linspace(0.5, 100.0, 200)
best_T, best_nll = 1.0, float("inf")
l = torch.as_tensor(logits, dtype=torch.float32)
t = torch.as_tensor(y, dtype=torch.long)
for T in grid:
p = torch.softmax(l / float(T), dim=-1)
nll = -torch.log(p[torch.arange(len(t)), t] + 1e-12).mean().item()
if nll < best_nll:
best_nll, best_T = nll, float(T)
return best_T
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--out", default="liar-detector-v4")
ap.add_argument("--n-train", type=int, default=6000)
ap.add_argument("--epochs", type=int, default=200)
ap.add_argument("--batch", type=int, default=64)
ap.add_argument("--lr", type=float, default=3e-3)
ap.add_argument("--seed", type=int, default=0)
args = ap.parse_args()
torch.manual_seed(args.seed)
device = torch.device("cuda" if torch.cuda.is_available() else "cpu")
print("Building training data...")
t0 = time.time()
X, by, fy, py, _ = build_arrays(args.n_train, args.seed, ID_FAMILIES)
print(f" {X.shape} in {time.time() - t0:.1f}s")
fe = LiarDetectorFeatureExtractor().fit(X)
Xn = fe.transform(X)
rng = np.random.default_rng(args.seed + 100)
perm = rng.permutation(len(Xn))
Xn, by, fy, py = Xn[perm], by[perm], fy[perm], py[perm]
n_val = len(Xn) // 5
tr, va = slice(None, -n_val), slice(-n_val,)
ds = TensorDataset(
torch.from_numpy(Xn[tr]),
torch.from_numpy(by[tr]),
torch.from_numpy(fy[tr]),
torch.from_numpy(py[tr]),
)
dl = DataLoader(ds, batch_size=args.batch, shuffle=True)
cfg = LiarDetectorConfig()
model = LiarDetectorForLieDetection(cfg).to(device)
n_par = sum(p.numel() for p in model.parameters())
print(f" parameters: {n_par}")
opt = torch.optim.AdamW(model.parameters(), lr=args.lr)
for epoch in range(args.epochs):
model.train()
total, nb = 0.0, 0
for xb, bb, fb, pb in dl:
xb, bb, fb, pb = xb.to(device), bb.to(device), fb.to(device), pb.to(device)
out = model(features=xb,
binary_labels=bb, family_labels=fb, presence_labels=pb)
opt.zero_grad()
out.loss.backward()
opt.step()
total += out.loss.item(); nb += 1
if (epoch + 1) % 25 == 0 or epoch == 0:
print(f" epoch {epoch+1:>4} loss {total/nb:.4f}")
# temperature calibration on validation split
model.eval()
with torch.no_grad():
xv = torch.from_numpy(Xn[va]).to(device)
out = model(features=xv)
cfg.temp_binary = fit_temperature(out.binary_logits.cpu().numpy(), by[va])
cfg.temp_family = fit_temperature(out.family_logits.cpu().numpy(), fy[va])
cfg.temp_presence = fit_temperature(out.presence_logits.cpu().numpy(), py[va])
print(f" temperatures: bin={cfg.temp_binary:.2f} "
f"fam={cfg.temp_family:.2f} pres={cfg.temp_presence:.2f}")
os.makedirs(args.out, exist_ok=True)
model.save_pretrained(args.out)
fe.save_pretrained(args.out)
print(f"Saved to {args.out}/")
if __name__ == "__main__":
main()