CoolFace
Apppublic

xndung/robustness_token_segmentation

sourceHugging Facemitupdated 4mo agoView on Hugging Face
0likes
train.py139 linesDownload Raw Back to src
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