xndung/robustness_token_segmentation
0
1import os
2import random
3
4import numpy as np
5import torch
6import yaml
7from accelerate import Accelerator
8from torch.nn.functional import mse_loss
9from tqdm.auto import tqdm
10
11import wandb
12from attacks.pgd import pgd_attack
13from data.transforms import unnormalize
14from data.utils import get_loaders_fn
15from models.utils import get_model
16from utils import read_config
17
18
19def train_rtokens(
20 model,
21 loader,
22 criterion,
23 optim,
24 accelerator,
25 max_steps,
26 checkpoint_freq,
27 store_path,
28):
29 """Training loop to optimize robustness tokens."""
30 # Preparing model, optimizer and data loader
31 model, optim, loader = accelerator.prepare(model, optim, loader)
32 model_path = os.path.join(store_path, "last.ckpt")
33 tokens_path = os.path.join(store_path, "rtokens.pt")
34
35 # Loop
36 steps = 0
37 with tqdm(total=max_steps, desc="Training") as pbar:
38 while steps < max_steps:
39 for batch in loader:
40 batch = batch[0]
41
42 # Getting adversarial batch w.r.t. standard model
43 model.enable_robust = False
44 batch_adv = pgd_attack(model, batch)
45
46 # Getting target features and loss w.r.t. standard model
47 with torch.no_grad():
48 target = model(batch)
49 mse = mse_loss(unnormalize(batch_adv), unnormalize(batch))
50 baseline = criterion(model(batch_adv), target).mean().item()
51
52 # Training robustness tokens
53 model.enable_robust = True
54 loss_inv = criterion(model(batch), target).mean()
55 loss_adv = criterion(model(batch_adv), target).mean()
56 loss = loss_inv + loss_adv
57 optim.zero_grad()
58 accelerator.backward(loss)
59 optim.step()
60
61 wandb.log(
62 {
63 "Train Loss": loss.item(),
64 "Train Image MSE": mse.item(),
65 "Train Loss Invariance": loss_inv.item(),
66 "Train Loss Adversarial": loss_adv.item(),
67 "Without robustness": baseline,
68 },
69 step=steps,
70 )
71
72 steps += 1
73 pbar.update(1)
74
75 if steps % checkpoint_freq == 0:
76 torch.save(accelerator.get_state_dict(model), model_path)
77 model.store_rtokens(tokens_path)
78
79 if steps >= max_steps:
80 break
81 torch.save(accelerator.get_state_dict(model), model_path)
82 model.store_rtokens(tokens_path)
83
84
85def main(args):
86 # Setting seed
87 random.seed(args["seed"])
88 np.random.seed(args["seed"])
89 torch.manual_seed(args["seed"])
90 torch.cuda.manual_seed_all(args["seed"])
91
92 # Initializing wandb
93 wandb.init(project="Robustness Tokens", name=args["run_name"], config=args)
94
95 # Creating result directory and copying config file
96 os.makedirs(args["results_dir"], exist_ok=True)
97 yaml.dump(args, open(os.path.join(args["results_dir"], "config.yaml"), "w"))
98
99 # Initializing model
100 model = get_model(**args["model"])
101
102 # Preparing data loaders
103 loaders_fn = get_loaders_fn(args["dataset"])
104 train_loader, _ = loaders_fn(
105 args["train"]["batch_size"], num_workers=args["train"]["num_workers"]
106 )
107
108 # Training hyper-parameters
109 max_steps = args["train"]["max_steps"]
110 checkpoint_freq = args["train"]["checkpoint_freq"]
111 store_path = args["results_dir"]
112 criterion = getattr(torch.nn, args["train"]["criterion"])()
113 optim = getattr(torch.optim, args["train"]["optimizer"])(
114 model.get_trainable_parameters(),
115 lr=args["train"]["lr"],
116 maximize=(args["train"]["mode"] == "max"),
117 )
118
119 # Training loop
120 accelerator = Accelerator()
121 train_rtokens(
122 model,
123 train_loader,
124 criterion,
125 optim,
126 accelerator,
127 max_steps,
128 checkpoint_freq,
129 store_path,
130 )
131
132 # Finishing wandb
133 wandb.finish()
134 print("\n\n\nProgram completed successfully.")
135
136
137if __name__ == "__main__":
138 main(read_config())
139 