CoolFace
Apppublic

xndung/robustness_token_segmentation

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
feat.py90 linesDownload Raw Back to robustness
1import os
2import random
3import numpy as np
4import pandas as pd
5import torch
6from accelerate import Accelerator
7from torch.nn.functional import cosine_similarity
8from tqdm.auto import tqdm
9
10from attacks.pgd import pgd_attack
11from data.utils import get_loaders_fn
12from models.utils import get_model
13from utils import read_config
14
15
16def evaluate_robustness(surrogate, victim, loader, accelerator):
17    """Test loop that evaluates robustness of the model compared to the baseline."""
18    # Preparing accelerator
19    surrogate, victim, loader = accelerator.prepare(surrogate, victim, loader)
20
21    cossims = []
22    mses = []
23
24    def cossim(f1, f2):
25        out = cosine_similarity(f1, f2)
26        dim = list(range(1, out.ndim))
27        return out.mean(dim=dim).cpu().numpy()
28
29    def mse(f1, f2):
30        dim = list(range(1, f1.ndim))
31        return (f1 - f2).pow(2).mean(dim=dim).cpu().numpy()
32
33    for batch in tqdm(loader, desc="Evaluating robustness"):
34        batch = batch[0]
35        batch_adv = pgd_attack(surrogate, batch)
36
37        with torch.no_grad():
38            f1 = victim(batch)
39            f2 = victim(batch_adv)
40            cossims.extend(cossim(f1, f2))
41            mses.extend(mse(f1, f2))
42            print(f"Cosine Sim: {np.mean(cossims):.3f}  - MSE: {np.mean(mses):.3f}")
43
44    return {"Cosine Sim": cossims}, {"MSEs": mses}
45
46
47def main(args):
48    # Setting seed
49    random.seed(args["seed"])
50    np.random.seed(args["seed"])
51    torch.manual_seed(args["seed"])
52    torch.cuda.manual_seed_all(args["seed"])
53
54    # Accelerator
55    accelerator = Accelerator()
56
57    # Data
58    loaders_fn = get_loaders_fn(args["dataset"])
59    _, val_loader = loaders_fn(args["batch_size"], args["num_workers"])
60
61    # Surrogate
62    surrogate = get_model(**args["surrogate"])
63    if args.get("surrogate_state_dict", None) is not None:
64        surrogate.load_state_dict(
65            torch.load(args["surrogate_state_dict"], map_location=accelerator.device)
66        )
67
68    # Victim
69    victim = get_model(**args["victim"])
70    if args.get("victim_state_dict", None) is not None:
71        victim.load_state_dict(
72            torch.load(args["victim_state_dict"], map_location=accelerator.device)
73        )
74
75    # Attacking model
76    cossims, mses = evaluate_robustness(surrogate, victim, val_loader, accelerator)
77
78    # Saving metrics
79    rdir = args["results_dir"]
80    os.makedirs(rdir, exist_ok=True)
81    cossims = pd.DataFrame.from_dict(cossims)
82    mses = pd.DataFrame.from_dict(mses)
83    cossims.to_csv(os.path.join(rdir, "cossims.csv"))
84    mses.to_csv(os.path.join(rdir, "mses.csv"))
85    print(f"Robustness metrics saved in {rdir}")
86
87
88if __name__ == "__main__":
89    main(read_config())
90